From 0d1a4fdac34c5e649be7a450cfa742fa38dd4429 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Sun, 29 Mar 2026 18:53:35 +0000 Subject: [PATCH] feat: real DNS lookups, manual DKIM selectors, Cloudflare-ready DNS provider architecture Agent-Logs-Url: https://github.com/christianlouis/dmarq/sessions/19d17518-732d-4644-889b-cc63256e19b1 Co-authored-by: christianlouis <361235+christianlouis@users.noreply.github.com> --- backend/app/api/api_v1/endpoints/domains.py | 210 +++++++++++++-- backend/app/services/dns_resolver.py | 263 +++++++++++++++++++ backend/app/templates/domain_details.html | 125 ++++++++- backend/app/templates/domains.html | 16 +- backend/app/tests/test_dns_endpoints.py | 252 ++++++++++++++++++ backend/app/tests/test_dns_resolver.py | 269 ++++++++++++++++++++ 6 files changed, 1096 insertions(+), 39 deletions(-) create mode 100644 backend/app/services/dns_resolver.py create mode 100644 backend/app/tests/test_dns_endpoints.py create mode 100644 backend/app/tests/test_dns_resolver.py diff --git a/backend/app/api/api_v1/endpoints/domains.py b/backend/app/api/api_v1/endpoints/domains.py index 457e727..cf3cdd2 100644 --- a/backend/app/api/api_v1/endpoints/domains.py +++ b/backend/app/api/api_v1/endpoints/domains.py @@ -1,11 +1,23 @@ +import asyncio +import logging from datetime import datetime, timezone from typing import Any, Dict, List, Optional -from fastapi import APIRouter, HTTPException, Path, Query, status -from pydantic import BaseModel +from fastapi import APIRouter, Depends, HTTPException, Path, Query, status +from pydantic import BaseModel, Field +from sqlalchemy.orm import Session +from app.core.database import get_db +from app.models.domain import Domain +from app.services.dns_resolver import ( + DomainDNSResult, + extract_dmarc_policy, + get_default_provider, +) from app.services.report_store import ReportStore +logger = logging.getLogger(__name__) + router = APIRouter() @@ -98,15 +110,70 @@ class DomainSummaryResponse(BaseModel): domains: List[Dict[str, Any]] +class SelectorRequest(BaseModel): + """Request body for adding a DKIM selector""" + + selector: str = Field(..., min_length=1, description="DKIM selector name") + + +def _get_selectors_from_reports(store: "ReportStore", domain: str) -> List[str]: + """Extract DKIM selectors seen in stored DMARC reports for *domain*. + + DMARC aggregate report records include DKIM auth results that carry the + selector used by the sending server. Collecting these gives us a set of + real-world selectors to verify against live DNS, in addition to any + manually configured selectors. + """ + selectors: List[str] = [] + for report in store.get_domain_reports(domain): + for record in report.get("records", []): + for dkim_entry in record.get("dkim", []): + sel = dkim_entry.get("selector", "").strip() + if sel and sel not in selectors: + selectors.append(sel) + return selectors + + +def _get_domain_selectors_from_db(db: Session, domain_name: str) -> List[str]: + """Return the manually configured DKIM selectors for *domain_name* from the DB.""" + domain_db = db.query(Domain).filter(Domain.name == domain_name).first() + if domain_db and domain_db.dkim_selectors: + return [s.strip() for s in domain_db.dkim_selectors.split(",") if s.strip()] + return [] + + @router.get("/summary", response_model=DomainSummaryResponse) -async def get_domains_summary(): +async def get_domains_summary(db: Session = Depends(get_db)): """ Get summary statistics for all domains, formatted for the dashboard. + + Performs live DNS lookups for each domain concurrently and includes the + results (DMARC/SPF/DKIM status and live DMARC policy) in the per-domain + entries. A per-domain timeout of 10 s prevents slow DNS responses from + blocking the page load. """ store = ReportStore.get_instance() domains = store.get_domains() summaries = store.get_all_domain_summaries() + # Perform DNS checks concurrently for all domains + provider = get_default_provider() + + async def _dns_for_domain(domain_name: str) -> DomainDNSResult: + manual_selectors = _get_domain_selectors_from_db(db, domain_name) + report_selectors = _get_selectors_from_reports(store, domain_name) + combined = list(dict.fromkeys(manual_selectors + report_selectors)) + try: + return await asyncio.wait_for( + provider.check_domain(domain_name, selectors=combined), + timeout=10.0, + ) + except (asyncio.TimeoutError, Exception) as exc: # pylint: disable=broad-exception-caught + logger.warning("DNS check failed for %s: %s", domain_name, exc) + return DomainDNSResult() + + dns_results = await asyncio.gather(*[_dns_for_domain(d) for d in domains]) + # Calculate overall statistics total_domains = len(domains) total_emails = 0 @@ -115,22 +182,34 @@ async def get_domains_summary(): domains_list = [] - for domain_name in domains: + for domain_name, dns in zip(domains, dns_results): summary = summaries.get(domain_name, {}) total_emails += summary.get("total_count", 0) total_passed += summary.get("passed_count", 0) total_reports += summary.get("reports_processed", 0) + # Prefer live DNS policy; fall back to policy seen in reports + live_policy = extract_dmarc_policy(dns.dmarc_record) + reported_policy = summary.get("policy", {}) + if isinstance(reported_policy, dict): + reported_policy = reported_policy.get("p") + dmarc_policy = live_policy or reported_policy or "none" + # Format domain data for frontend domains_list.append( { - "id": domain_name, # Using the domain name as ID for now + "id": domain_name, "domain_name": domain_name, "total_emails": summary.get("total_count", 0), "passed_count": summary.get("passed_count", 0), "failed_count": summary.get("failed_count", 0), "pass_rate": summary.get("compliance_rate", 0), "report_count": summary.get("reports_processed", 0), + # Real DNS status + "dmarc_status": dns.dmarc, + "dmarc_policy": dmarc_policy, + "spf_status": dns.spf, + "dkim_status": dns.dkim, } ) @@ -232,10 +311,16 @@ async def get_domain_stats(domain_id: str = Path(..., title="The domain ID or na @router.get("/{domain_id}/dns", response_model=DNSRecordResponse) -async def get_domain_dns_records(domain_id: str = Path(..., title="The domain ID or name")): +async def get_domain_dns_records( + domain_id: str = Path(..., title="The domain ID or name"), + db: Session = Depends(get_db), +): """ - Get DNS records for a specific domain. For Milestone 1, - this returns mock data since DNS integration is part of a future milestone. + Get DNS records for a specific domain using live DNS lookups. + + Manual selectors (stored in the database) are checked first, followed by + selectors observed in stored DMARC reports, with common well-known + selectors used as a final fallback. """ store = ReportStore.get_instance() domains = store.get_domains() @@ -246,19 +331,20 @@ async def get_domain_dns_records(domain_id: str = Path(..., title="The domain ID detail="Domain not found", ) - # For Milestone 1, return mock DNS record data - # In a future milestone, this will be replaced with actual DNS lookups - mock_dmarc_record = ( - "v=DMARC1; p=none; rua=mailto:dmarc@example.com;" - " ruf=mailto:forensic@example.com; pct=100" - ) + manual_selectors = _get_domain_selectors_from_db(db, domain_id) + report_selectors = _get_selectors_from_reports(store, domain_id) + combined_selectors = list(dict.fromkeys(manual_selectors + report_selectors)) + + provider = get_default_provider() + result = await provider.check_domain(domain_id, selectors=combined_selectors) + return DNSRecordResponse( - dmarc=True, - dmarcRecord=mock_dmarc_record, - spf=True, - spfRecord="v=spf1 include:_spf.google.com include:spf.protection.outlook.com -all", - dkim=True, - dkimSelectors="selector1, selector2", + dmarc=result.dmarc, + dmarcRecord=result.dmarc_record, + spf=result.spf, + spfRecord=result.spf_record, + dkim=result.dkim, + dkimSelectors=result.dkim_selector, ) @@ -385,6 +471,90 @@ async def get_domain_sources( return DomainSourcesResponse(sources=source_entries) +@router.get("/{domain_id}/selectors") +async def get_domain_selectors( + domain_id: str = Path(..., title="The domain ID or name"), + db: Session = Depends(get_db), +): + """Return the manually configured DKIM selectors for a domain.""" + store = ReportStore.get_instance() + if domain_id not in store.get_domains(): + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail="Domain not found", + ) + selectors = _get_domain_selectors_from_db(db, domain_id) + return {"selectors": selectors} + + +@router.post("/{domain_id}/selectors", status_code=status.HTTP_201_CREATED) +async def add_domain_selector( + selector_data: SelectorRequest, + domain_id: str = Path(..., title="The domain ID or name"), + db: Session = Depends(get_db), +): + """Add a DKIM selector to the manual list for a domain. + + The selector is persisted in the ``Domain`` database row so that it will + be used in all subsequent DNS checks, even if it has not yet appeared in + any received DMARC report. + """ + store = ReportStore.get_instance() + if domain_id not in store.get_domains(): + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail="Domain not found", + ) + + selector = selector_data.selector.strip() + if not selector: + raise HTTPException( + status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, + detail="Selector must not be empty", + ) + + domain_db = db.query(Domain).filter(Domain.name == domain_id).first() + if not domain_db: + domain_db = Domain(name=domain_id) + db.add(domain_db) + + existing = [s.strip() for s in (domain_db.dkim_selectors or "").split(",") if s.strip()] + if selector not in existing: + existing.append(selector) + domain_db.dkim_selectors = ",".join(existing) + db.commit() + + return {"selectors": existing} + + +@router.delete("/{domain_id}/selectors/{selector}", status_code=status.HTTP_200_OK) +async def delete_domain_selector( + domain_id: str = Path(..., title="The domain ID or name"), + selector: str = Path(..., title="The DKIM selector to remove"), + db: Session = Depends(get_db), +): + """Remove a manually configured DKIM selector from a domain.""" + domain_db = db.query(Domain).filter(Domain.name == domain_id).first() + if not domain_db: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail="Domain not found", + ) + + existing = [s.strip() for s in (domain_db.dkim_selectors or "").split(",") if s.strip()] + if selector not in existing: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=f"Selector '{selector}' not found", + ) + + existing.remove(selector) + domain_db.dkim_selectors = ",".join(existing) + db.commit() + + return {"selectors": existing} + + @router.delete("/{domain_id}", status_code=status.HTTP_204_NO_CONTENT) async def delete_domain(domain_id: str = Path(..., title="The domain ID or name")): """ diff --git a/backend/app/services/dns_resolver.py b/backend/app/services/dns_resolver.py new file mode 100644 index 0000000..502deb7 --- /dev/null +++ b/backend/app/services/dns_resolver.py @@ -0,0 +1,263 @@ +""" +DNS resolver service for DMARC, SPF, and DKIM record lookups. + +Provides an extensible provider architecture so that DNS data can be fetched +either via the system resolver (dnspython) or via the Cloudflare DNS API for +future Cloudflare integration. +""" + +import asyncio +import logging +from abc import ABC, abstractmethod +from dataclasses import dataclass, field +from typing import List, Optional, Tuple + +logger = logging.getLogger(__name__) + +# Well-known DKIM selectors tried when no selectors are configured +COMMON_DKIM_SELECTORS: List[str] = [ + "default", + "google", + "mail", + "selector1", + "selector2", + "dkim", + "k1", + "key1", + "mta", + "email", + "smtp", + "s1", + "s2", + "pm", + "mandrill", + "sendgrid", +] + +# Seconds to wait for a single DNS query before giving up +DNS_TIMEOUT: float = 5.0 + + +@dataclass +class DomainDNSResult: + """Aggregated DNS authentication record results for one domain.""" + + dmarc: bool = False + dmarc_record: Optional[str] = None + spf: bool = False + spf_record: Optional[str] = None + dkim: bool = False + dkim_selector: Optional[str] = None + dkim_record: Optional[str] = None + # Track which selectors were tried so callers can surface this information + selectors_checked: List[str] = field(default_factory=list) + + +class BaseDNSProvider(ABC): + """ + Abstract base class for DNS providers. + + Subclasses implement ``lookup_txt`` and inherit the higher-level helper + methods for DMARC, SPF, and DKIM checks so that provider-specific + differences stay confined to a single method. + """ + + @abstractmethod + async def lookup_txt(self, name: str) -> List[str]: + """Return TXT record strings for *name*. + + Raises ``LookupError`` on failure (NXDOMAIN, timeout, network error + etc.). Returns an empty list when the name exists but has no TXT + records. + """ + + # ------------------------------------------------------------------ + # High-level record checks built on top of lookup_txt + # ------------------------------------------------------------------ + + async def check_dmarc(self, domain: str) -> Tuple[bool, Optional[str]]: + """Return *(found, record_string)* for the domain's DMARC TXT record.""" + try: + records = await self.lookup_txt(f"_dmarc.{domain}") + for record in records: + if record.lower().startswith("v=dmarc1"): + return True, record + except LookupError as exc: + logger.debug("DMARC lookup failed for %s: %s", domain, exc) + return False, None + + async def check_spf(self, domain: str) -> Tuple[bool, Optional[str]]: + """Return *(found, record_string)* for the domain's SPF TXT record.""" + try: + records = await self.lookup_txt(domain) + for record in records: + if record.lower().startswith("v=spf1"): + return True, record + except LookupError as exc: + logger.debug("SPF lookup failed for %s: %s", domain, exc) + return False, None + + async def check_dkim( + self, domain: str, selectors: List[str] + ) -> Tuple[bool, Optional[str], Optional[str]]: + """Return *(found, selector, record_string)* for the first working DKIM selector.""" + for selector in selectors: + try: + records = await self.lookup_txt(f"{selector}._domainkey.{domain}") + for record in records: + if "v=dkim1" in record.lower() or "p=" in record.lower(): + return True, selector, record + except LookupError as exc: + logger.debug( + "DKIM lookup failed for selector=%s domain=%s: %s", selector, domain, exc + ) + return False, None, None + + async def check_domain( + self, domain: str, selectors: Optional[List[str]] = None + ) -> DomainDNSResult: + """Run DMARC, SPF, and DKIM checks concurrently for *domain*. + + *selectors* are tried first; common well-known selectors are appended + as a fallback so that a domain with no explicitly configured selectors + can still be verified. + """ + # Deduplicate while preserving priority order (manual selectors first) + all_selectors: List[str] = list(selectors or []) + for s in COMMON_DKIM_SELECTORS: + if s not in all_selectors: + all_selectors.append(s) + + dmarc_coro = self.check_dmarc(domain) + spf_coro = self.check_spf(domain) + dkim_coro = self.check_dkim(domain, all_selectors) + + (dmarc_ok, dmarc_record), (spf_ok, spf_record), (dkim_ok, dkim_sel, dkim_record) = ( + await asyncio.gather(dmarc_coro, spf_coro, dkim_coro) + ) + + return DomainDNSResult( + dmarc=dmarc_ok, + dmarc_record=dmarc_record, + spf=spf_ok, + spf_record=spf_record, + dkim=dkim_ok, + dkim_selector=dkim_sel, + dkim_record=dkim_record, + selectors_checked=all_selectors, + ) + + +class SystemDNSProvider(BaseDNSProvider): + """DNS provider that resolves records via the system resolver using dnspython.""" + + async def lookup_txt(self, name: str) -> List[str]: + """Resolve TXT records using dnspython's async resolver.""" + # Import here so the module can be imported even if dnspython is absent + # (tests can mock this method directly without needing the library). + import dns.asyncresolver # type: ignore[import] + import dns.exception # type: ignore[import] + + try: + answers = await dns.asyncresolver.resolve( + name, "TXT", lifetime=DNS_TIMEOUT, raise_on_no_answer=False + ) + result: List[str] = [] + if answers: + for rdata in answers: + for string in rdata.strings: + result.append(string.decode("utf-8", errors="replace")) + return result + except dns.exception.DNSException as exc: + raise LookupError(f"TXT lookup failed for {name}: {exc}") from exc + + +class CloudflareDNSProvider(BaseDNSProvider): + """DNS provider using Cloudflare's DNS-over-HTTPS (DoH) endpoint. + + This provider resolves DNS queries via Cloudflare's public DoH API + (``1.1.1.1`` / ``cloudflare-dns.com``). When *api_token* and *zone_id* + are supplied, future versions will also support reading and writing DNS + records directly through the Cloudflare REST API, enabling automated DNS + synchronisation. + + Current status + -------------- + * DoH-based lookups are fully functional. + * Direct Cloudflare API integration (zone management, record sync) is + reserved for a future release. + """ + + #: Cloudflare DNS-over-HTTPS endpoint (JSON wire format) + CLOUDFLARE_DOH_URL: str = "https://cloudflare-dns.com/dns-query" + #: Cloudflare REST API base URL (for future zone-management support) + CLOUDFLARE_API_BASE: str = "https://api.cloudflare.com/client/v4" + + def __init__( + self, + api_token: Optional[str] = None, + zone_id: Optional[str] = None, + ) -> None: + """ + Parameters + ---------- + api_token: + Cloudflare API token. Required for future DNS record management; + not needed for read-only DoH lookups. + zone_id: + Cloudflare zone identifier. Required for future DNS record + management. + """ + self.api_token = api_token + self.zone_id = zone_id + + async def lookup_txt(self, name: str) -> List[str]: + """Resolve TXT records via Cloudflare's DoH endpoint (JSON format).""" + import httpx # type: ignore[import] + + params = {"name": name, "type": "TXT"} + headers = {"Accept": "application/dns-json"} + try: + async with httpx.AsyncClient() as client: + response = await client.get( + self.CLOUDFLARE_DOH_URL, + params=params, + headers=headers, + timeout=DNS_TIMEOUT, + ) + response.raise_for_status() + data = response.json() + records: List[str] = [] + for answer in data.get("Answer", []): + if answer.get("type") == 16: # TXT record type + # Cloudflare wraps TXT values in double-quotes + txt = answer.get("data", "").strip('"') + records.append(txt) + return records + except Exception as exc: + raise LookupError(f"Cloudflare DoH lookup failed for {name}: {exc}") from exc + + +def get_default_provider() -> BaseDNSProvider: + """Return the default DNS provider (system resolver). + + In a future release this function will inspect application settings and + return a ``CloudflareDNSProvider`` when Cloudflare credentials are + configured. + """ + return SystemDNSProvider() + + +def extract_dmarc_policy(dmarc_record: Optional[str]) -> Optional[str]: + """Parse the *p=* tag from a DMARC TXT record string. + + Returns the policy value (e.g. ``"none"``, ``"quarantine"``, + ``"reject"``) or ``None`` if the record is absent or unparsable. + """ + if not dmarc_record: + return None + for part in dmarc_record.split(";"): + part = part.strip() + if part.lower().startswith("p="): + return part[2:].strip().lower() + return None diff --git a/backend/app/templates/domain_details.html b/backend/app/templates/domain_details.html index 227787d..c3d5fed 100644 --- a/backend/app/templates/domain_details.html +++ b/backend/app/templates/domain_details.html @@ -152,13 +152,55 @@
-
+

- DKIM Selectors - - + DKIM (live check) + +

-
-
+
-
+
+ +
+

DKIM Selectors

+

+ Manually configure selectors to check. Selectors seen in received DMARC + reports and common well-known selectors are checked automatically. +

+ +
+ + +
+ +
+ + +
+

{% endcall %} @@ -347,6 +389,9 @@ function domainDetailsApp(domainId) { dkim: false, dkimSelectors: '' }, + selectors: [], + newSelector: '', + selectorError: '', reports: [], sources: [], complianceChart: null, @@ -354,27 +399,28 @@ function domainDetailsApp(domainId) { dateRange: '30', sourceFilter: '' }, - + init() { this.fetchDomainStats(); this.fetchDNSRecords(); + this.fetchSelectors(); this.fetchReports(); this.fetchSources(); - + this.$watch('filters.dateRange', () => { this.fetchSources(); }); }, - + get filteredSources() { if (!this.sources) return []; - + return this.sources.filter(source => { if (!this.filters.sourceFilter) return true; return source.ip.toLowerCase().includes(this.filters.sourceFilter.toLowerCase()); }); }, - + async fetchDomainStats() { try { const response = await fetch(`/api/v1/domains/${this.domainId}/stats`); @@ -386,7 +432,7 @@ function domainDetailsApp(domainId) { console.error('Error fetching domain stats:', error); } }, - + async fetchDNSRecords() { try { const response = await fetch(`/api/v1/domains/${this.domainId}/dns`); @@ -398,7 +444,64 @@ function domainDetailsApp(domainId) { console.error('Error fetching DNS records:', error); } }, - + + async fetchSelectors() { + try { + const response = await fetch(`/api/v1/domains/${this.domainId}/selectors`); + if (response.ok) { + const data = await response.json(); + this.selectors = data.selectors || []; + } + } catch (error) { + console.error('Error fetching selectors:', error); + } + }, + + async addSelector() { + this.selectorError = ''; + const sel = this.newSelector.trim(); + if (!sel) return; + try { + const response = await fetch(`/api/v1/domains/${this.domainId}/selectors`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ selector: sel }) + }); + if (response.ok) { + const data = await response.json(); + this.selectors = data.selectors || []; + this.newSelector = ''; + // Refresh DNS check to reflect the new selector + this.fetchDNSRecords(); + } else { + const err = await response.json(); + this.selectorError = err.detail || 'Failed to add selector'; + } + } catch (error) { + this.selectorError = 'Network error — could not add selector'; + console.error('Error adding selector:', error); + } + }, + + async deleteSelector(selector) { + try { + const response = await fetch( + `/api/v1/domains/${this.domainId}/selectors/${encodeURIComponent(selector)}`, + { method: 'DELETE' } + ); + if (response.ok) { + const data = await response.json(); + this.selectors = data.selectors || []; + // Refresh DNS check after removing a selector + this.fetchDNSRecords(); + } else { + console.error('Error deleting selector:', response.status); + } + } catch (error) { + console.error('Error deleting selector:', error); + } + }, + async fetchReports() { try { const response = await fetch(`/api/v1/domains/${this.domainId}/reports?limit=10`); diff --git a/backend/app/templates/domains.html b/backend/app/templates/domains.html index a844326..1152f0f 100644 --- a/backend/app/templates/domains.html +++ b/backend/app/templates/domains.html @@ -99,25 +99,25 @@ function domainsApp() { return { domains: [], - + init() { // Fetch domains from server this.fetchDomains(); }, - + async fetchDomains() { try { const response = await fetch('/api/v1/domains/summary'); if (response.ok) { const data = await response.json(); - - // Format domains for display + + // Map API fields — DNS status comes directly from live lookups this.domains = data.domains.map(domain => ({ name: domain.domain_name, - dmarc_status: true, // In Milestone 1, assume DMARC is configured if we have reports - dmarc_policy: domain.policy || 'p=none', - spf_status: true, // In future milestones, this will come from DNS checks - dkim_status: true, // In future milestones, this will come from DNS checks + dmarc_status: domain.dmarc_status ?? false, + dmarc_policy: domain.dmarc_policy || 'Not configured', + spf_status: domain.spf_status ?? false, + dkim_status: domain.dkim_status ?? false, reports_count: domain.report_count, emails_count: domain.total_emails, compliance_rate: domain.pass_rate diff --git a/backend/app/tests/test_dns_endpoints.py b/backend/app/tests/test_dns_endpoints.py new file mode 100644 index 0000000..8326372 --- /dev/null +++ b/backend/app/tests/test_dns_endpoints.py @@ -0,0 +1,252 @@ +""" +Integration tests for the DKIM selector management API endpoints. + +These tests use the in-memory SQLite test database via the ``client`` fixture +(which overrides ``get_db``) and populate the ``ReportStore`` singleton so +that the endpoints can find the test domain. + +DNS lookups are mocked so no real network calls are made. +""" + +from unittest.mock import AsyncMock, patch + +import pytest +from fastapi.testclient import TestClient + +from app.services.dns_resolver import DomainDNSResult +from app.services.report_store import ReportStore + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + +DOMAIN = "example.com" + +# A minimal parsed DMARC report that populates the ReportStore +MINIMAL_REPORT = { + "domain": DOMAIN, + "report_id": "test-001", + "org_name": "Test Org", + "policy": {"p": "none", "sp": "", "pct": "100"}, + "records": [ + { + "source_ip": "1.2.3.4", + "count": 5, + "disposition": "none", + "dkim_result": "pass", + "spf_result": "pass", + "dkim": [{"domain": DOMAIN, "result": "pass", "selector": "google"}], + "spf": [{"domain": DOMAIN, "result": "pass"}], + } + ], + "summary": {"total_count": 5, "passed_count": 5, "failed_count": 0, "pass_rate": 100.0}, +} + +# DomainDNSResult returned by the mocked DNS provider +MOCK_DNS_RESULT = DomainDNSResult( + dmarc=True, + dmarc_record="v=DMARC1; p=none; rua=mailto:dmarc@example.com", + spf=True, + spf_record="v=spf1 include:_spf.google.com ~all", + dkim=True, + dkim_selector="google", + dkim_record="v=DKIM1; k=rsa; p=ABC", +) + + +@pytest.fixture(autouse=True) +def _seed_report_store(): + """Put a domain into the ReportStore for every test in this module.""" + store = ReportStore.get_instance() + store.add_report(MINIMAL_REPORT) + yield + + +def _mock_dns(result: DomainDNSResult = MOCK_DNS_RESULT): + """Return a context manager that patches the DNS provider's check_domain.""" + return patch( + "app.api.api_v1.endpoints.domains.get_default_provider", + return_value=AsyncMock(check_domain=AsyncMock(return_value=result)), + ) + + +# --------------------------------------------------------------------------- +# GET /api/v1/domains/{domain_id}/selectors +# --------------------------------------------------------------------------- + + +def test_get_selectors_empty(client: TestClient): + """Returns an empty list when no selectors have been configured.""" + response = client.get(f"/api/v1/domains/{DOMAIN}/selectors") + assert response.status_code == 200 + assert response.json() == {"selectors": []} + + +def test_get_selectors_unknown_domain(client: TestClient): + """Returns 404 for a domain not in the ReportStore.""" + response = client.get("/api/v1/domains/unknown.example.com/selectors") + assert response.status_code == 404 + + +# --------------------------------------------------------------------------- +# POST /api/v1/domains/{domain_id}/selectors +# --------------------------------------------------------------------------- + + +def test_add_selector(client: TestClient): + """Adding a selector persists it and returns the updated list.""" + response = client.post( + f"/api/v1/domains/{DOMAIN}/selectors", + json={"selector": "mysel"}, + ) + assert response.status_code == 201 + data = response.json() + assert "mysel" in data["selectors"] + + +def test_add_selector_deduplication(client: TestClient): + """Adding the same selector twice should not create duplicates.""" + client.post(f"/api/v1/domains/{DOMAIN}/selectors", json={"selector": "dup"}) + response = client.post(f"/api/v1/domains/{DOMAIN}/selectors", json={"selector": "dup"}) + assert response.status_code == 201 + assert response.json()["selectors"].count("dup") == 1 + + +def test_add_selector_invalid_empty(client: TestClient): + """An empty selector string should be rejected.""" + response = client.post( + f"/api/v1/domains/{DOMAIN}/selectors", + json={"selector": " "}, + ) + assert response.status_code == 422 + + +def test_add_selector_unknown_domain(client: TestClient): + """Adding a selector to an unknown domain returns 404.""" + response = client.post( + "/api/v1/domains/unknown.example.com/selectors", + json={"selector": "google"}, + ) + assert response.status_code == 404 + + +def test_add_multiple_selectors(client: TestClient): + """Multiple distinct selectors can be added and all are returned.""" + for sel in ("sel1", "sel2", "sel3"): + r = client.post(f"/api/v1/domains/{DOMAIN}/selectors", json={"selector": sel}) + assert r.status_code == 201 + + response = client.get(f"/api/v1/domains/{DOMAIN}/selectors") + assert response.status_code == 200 + selectors = response.json()["selectors"] + assert "sel1" in selectors + assert "sel2" in selectors + assert "sel3" in selectors + + +# --------------------------------------------------------------------------- +# DELETE /api/v1/domains/{domain_id}/selectors/{selector} +# --------------------------------------------------------------------------- + + +def test_delete_selector(client: TestClient): + """Deleting a selector removes it from the persisted list.""" + client.post(f"/api/v1/domains/{DOMAIN}/selectors", json={"selector": "todelete"}) + response = client.delete(f"/api/v1/domains/{DOMAIN}/selectors/todelete") + assert response.status_code == 200 + assert "todelete" not in response.json()["selectors"] + + +def test_delete_nonexistent_selector(client: TestClient): + """Deleting a selector that was never added returns 404.""" + # Ensure the domain exists in DB (via add then delete) + client.post(f"/api/v1/domains/{DOMAIN}/selectors", json={"selector": "dummy"}) + response = client.delete(f"/api/v1/domains/{DOMAIN}/selectors/ghost") + assert response.status_code == 404 + + +def test_delete_selector_unknown_domain(client: TestClient): + """Deleting from an unknown domain returns 404.""" + response = client.delete("/api/v1/domains/unknown.example.com/selectors/google") + assert response.status_code == 404 + + +# --------------------------------------------------------------------------- +# GET /api/v1/domains/{domain_id}/dns (real DNS replaced by mock) +# --------------------------------------------------------------------------- + + +def test_dns_endpoint_returns_real_data(client: TestClient): + """The /dns endpoint should return the mocked DNS check result.""" + with _mock_dns(): + response = client.get(f"/api/v1/domains/{DOMAIN}/dns") + + assert response.status_code == 200 + data = response.json() + assert data["dmarc"] is True + assert data["spf"] is True + assert data["dkim"] is True + assert "p=none" in data["dmarcRecord"] + + +def test_dns_endpoint_uses_manual_selectors(client: TestClient): + """Manually added selectors should be forwarded to check_domain.""" + # Add a custom selector + client.post(f"/api/v1/domains/{DOMAIN}/selectors", json={"selector": "customsel"}) + + captured_selectors = [] + + async def _fake_check_domain(domain, selectors=None): + captured_selectors.extend(selectors or []) + return MOCK_DNS_RESULT + + with patch( + "app.api.api_v1.endpoints.domains.get_default_provider", + return_value=AsyncMock(check_domain=_fake_check_domain), + ): + client.get(f"/api/v1/domains/{DOMAIN}/dns") + + assert "customsel" in captured_selectors + + +def test_dns_endpoint_404_for_unknown_domain(client: TestClient): + with _mock_dns(): + response = client.get("/api/v1/domains/unknown.example.com/dns") + assert response.status_code == 404 + + +# --------------------------------------------------------------------------- +# GET /api/v1/domains/summary (DNS fields included) +# --------------------------------------------------------------------------- + + +def test_summary_includes_dns_fields(client: TestClient): + """The summary endpoint should include dmarc_status, spf_status, dkim_status.""" + with _mock_dns(): + response = client.get("/api/v1/domains/summary") + + assert response.status_code == 200 + data = response.json() + assert data["total_domains"] == 1 + domain = data["domains"][0] + assert "dmarc_status" in domain + assert "spf_status" in domain + assert "dkim_status" in domain + assert domain["dmarc_status"] is True + assert domain["spf_status"] is True + assert domain["dkim_status"] is True + assert domain["dmarc_policy"] == "none" + + +def test_summary_dns_failure_defaults_false(client: TestClient): + """If DNS check fails, status fields default to False rather than crashing.""" + empty_result = DomainDNSResult() + + with _mock_dns(result=empty_result): + response = client.get("/api/v1/domains/summary") + + assert response.status_code == 200 + domain = response.json()["domains"][0] + assert domain["dmarc_status"] is False + assert domain["spf_status"] is False + assert domain["dkim_status"] is False diff --git a/backend/app/tests/test_dns_resolver.py b/backend/app/tests/test_dns_resolver.py new file mode 100644 index 0000000..127332a --- /dev/null +++ b/backend/app/tests/test_dns_resolver.py @@ -0,0 +1,269 @@ +""" +Unit tests for app.services.dns_resolver. + +DNS network I/O is mocked at the ``lookup_txt`` level so no real DNS queries +are made during testing. +""" + +from unittest.mock import AsyncMock, patch + +import pytest + +from app.services.dns_resolver import ( + BaseDNSProvider, + CloudflareDNSProvider, + DomainDNSResult, + SystemDNSProvider, + extract_dmarc_policy, + get_default_provider, +) + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +class FakeDNSProvider(BaseDNSProvider): + """Concrete provider backed by a simple dict for deterministic tests.""" + + def __init__(self, records: dict): + self._records = records + + async def lookup_txt(self, name: str): + if name in self._records: + return self._records[name] + raise LookupError(f"NXDOMAIN: {name}") + + +# --------------------------------------------------------------------------- +# extract_dmarc_policy +# --------------------------------------------------------------------------- + + +def test_extract_dmarc_policy_none(): + record = "v=DMARC1; p=none; rua=mailto:dmarc@example.com" + assert extract_dmarc_policy(record) == "none" + + +def test_extract_dmarc_policy_quarantine(): + record = "v=DMARC1; p=quarantine; pct=100" + assert extract_dmarc_policy(record) == "quarantine" + + +def test_extract_dmarc_policy_reject(): + assert extract_dmarc_policy("v=DMARC1; p=reject") == "reject" + + +def test_extract_dmarc_policy_missing_tag(): + assert extract_dmarc_policy("v=DMARC1; rua=mailto:dmarc@example.com") is None + + +def test_extract_dmarc_policy_none_input(): + assert extract_dmarc_policy(None) is None + + +# --------------------------------------------------------------------------- +# BaseDNSProvider helpers via FakeDNSProvider +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_check_dmarc_found(): + provider = FakeDNSProvider( + {"_dmarc.example.com": ["v=DMARC1; p=quarantine; rua=mailto:dmarc@example.com"]} + ) + found, record = await provider.check_dmarc("example.com") + assert found is True + assert record is not None + assert "p=quarantine" in record + + +@pytest.mark.asyncio +async def test_check_dmarc_not_found(): + provider = FakeDNSProvider({}) + found, record = await provider.check_dmarc("example.com") + assert found is False + assert record is None + + +@pytest.mark.asyncio +async def test_check_spf_found(): + provider = FakeDNSProvider( + {"example.com": ["v=spf1 include:_spf.google.com ~all", "some-other-record"]} + ) + found, record = await provider.check_spf("example.com") + assert found is True + assert record is not None + assert record.startswith("v=spf1") + + +@pytest.mark.asyncio +async def test_check_spf_not_found(): + provider = FakeDNSProvider({"example.com": ["some-other-record"]}) + found, record = await provider.check_spf("example.com") + assert found is False + assert record is None + + +@pytest.mark.asyncio +async def test_check_dkim_found_first_selector(): + provider = FakeDNSProvider( + {"google._domainkey.example.com": ["v=DKIM1; k=rsa; p=MIGfMA0GCSqGSIb3"]} + ) + found, selector, record = await provider.check_dkim("example.com", ["google", "mail"]) + assert found is True + assert selector == "google" + assert record is not None + + +@pytest.mark.asyncio +async def test_check_dkim_found_second_selector(): + provider = FakeDNSProvider( + {"mail._domainkey.example.com": ["v=DKIM1; k=rsa; p=MIGfMA0GCSqGSIb3"]} + ) + found, selector, record = await provider.check_dkim("example.com", ["google", "mail"]) + assert found is True + assert selector == "mail" + + +@pytest.mark.asyncio +async def test_check_dkim_not_found(): + provider = FakeDNSProvider({}) + found, selector, record = await provider.check_dkim("example.com", ["google", "mail"]) + assert found is False + assert selector is None + assert record is None + + +@pytest.mark.asyncio +async def test_check_domain_all_present(): + provider = FakeDNSProvider( + { + "_dmarc.example.com": ["v=DMARC1; p=none; rua=mailto:dmarc@example.com"], + "example.com": ["v=spf1 include:_spf.google.com ~all"], + "google._domainkey.example.com": ["v=DKIM1; k=rsa; p=MIGfMA0GCSqGSIb3"], + } + ) + result = await provider.check_domain("example.com", selectors=["google"]) + assert isinstance(result, DomainDNSResult) + assert result.dmarc is True + assert result.spf is True + assert result.dkim is True + assert result.dkim_selector == "google" + + +@pytest.mark.asyncio +async def test_check_domain_none_present(): + provider = FakeDNSProvider({}) + result = await provider.check_domain("missing.example.com") + assert result.dmarc is False + assert result.spf is False + assert result.dkim is False + + +@pytest.mark.asyncio +async def test_check_domain_uses_common_selectors_as_fallback(): + """When no selectors are passed, common selectors should be tried.""" + # Use 'default' which is in COMMON_DKIM_SELECTORS + provider = FakeDNSProvider({"default._domainkey.example.com": ["v=DKIM1; k=rsa; p=ABC"]}) + result = await provider.check_domain("example.com", selectors=[]) + assert result.dkim is True + assert result.dkim_selector == "default" + + +@pytest.mark.asyncio +async def test_check_domain_manual_selectors_take_priority(): + """Manually supplied selectors must be checked before common ones.""" + # Only the manual selector 'custom' has a record + provider = FakeDNSProvider({"custom._domainkey.example.com": ["v=DKIM1; k=rsa; p=XYZ"]}) + result = await provider.check_domain("example.com", selectors=["custom"]) + assert result.dkim is True + assert result.dkim_selector == "custom" + + +# --------------------------------------------------------------------------- +# SystemDNSProvider +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_system_provider_returns_txt_records(): + """SystemDNSProvider.lookup_txt should decode dnspython rdata correctly.""" + mock_string = b"v=DMARC1; p=none" + + class FakeRdata: + strings = [mock_string] + + class FakeAnswers: + def __iter__(self): + return iter([FakeRdata()]) + + with patch("dns.asyncresolver.resolve", new=AsyncMock(return_value=FakeAnswers())): + provider = SystemDNSProvider() + records = await provider.lookup_txt("_dmarc.example.com") + + assert records == ["v=DMARC1; p=none"] + + +@pytest.mark.asyncio +async def test_system_provider_raises_lookup_error_on_dns_exception(): + import dns.exception # type: ignore[import] + + with patch( + "dns.asyncresolver.resolve", + new=AsyncMock(side_effect=dns.exception.DNSException("NXDOMAIN")), + ): + provider = SystemDNSProvider() + with pytest.raises(LookupError): + await provider.lookup_txt("nonexistent.example.com") + + +# --------------------------------------------------------------------------- +# CloudflareDNSProvider +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_cloudflare_provider_parses_doh_response(): + """CloudflareDNSProvider should parse the Cloudflare DoH JSON response.""" + from unittest.mock import MagicMock + + fake_response_data = { + "Answer": [ + {"type": 16, "data": '"v=DMARC1; p=reject"'}, + {"type": 1, "data": "93.184.216.34"}, # A record — should be ignored + ] + } + + mock_response = AsyncMock() + mock_response.raise_for_status = MagicMock() # raise_for_status is synchronous in httpx + mock_response.json = lambda: fake_response_data + + with patch("httpx.AsyncClient.get", new=AsyncMock(return_value=mock_response)): + provider = CloudflareDNSProvider() + records = await provider.lookup_txt("_dmarc.example.com") + + assert records == ["v=DMARC1; p=reject"] + + +@pytest.mark.asyncio +async def test_cloudflare_provider_raises_on_http_error(): + import httpx + + with patch( + "httpx.AsyncClient.get", + new=AsyncMock(side_effect=httpx.RequestError("connection refused")), + ): + provider = CloudflareDNSProvider() + with pytest.raises(LookupError): + await provider.lookup_txt("_dmarc.example.com") + + +# --------------------------------------------------------------------------- +# get_default_provider +# --------------------------------------------------------------------------- + + +def test_get_default_provider_returns_system(): + provider = get_default_provider() + assert isinstance(provider, SystemDNSProvider)