From 57e164d282e2ee2a45900f0da323e460c8ee532a Mon Sep 17 00:00:00 2001 From: Christian Krakau-Louis Date: Sat, 23 May 2026 16:16:55 +0200 Subject: [PATCH] feat: add mta-sts posture checks --- backend/app/api/api_v1/endpoints/domains.py | 112 +++++++- backend/app/services/dns_cache.py | 27 +- backend/app/services/mta_sts.py | 216 ++++++++++++++ backend/app/templates/domain_details.html | 49 +++- backend/app/tests/test_dns_endpoints.py | 161 ++++++++++- backend/app/tests/test_mta_sts.py | 303 ++++++++++++++++++++ docs/milestones.md | 4 +- docs/user_guide/domains.md | 11 +- 8 files changed, 866 insertions(+), 17 deletions(-) create mode 100644 backend/app/services/mta_sts.py create mode 100644 backend/app/tests/test_mta_sts.py diff --git a/backend/app/api/api_v1/endpoints/domains.py b/backend/app/api/api_v1/endpoints/domains.py index f4bf63b..2c9274d 100644 --- a/backend/app/api/api_v1/endpoints/domains.py +++ b/backend/app/api/api_v1/endpoints/domains.py @@ -29,6 +29,7 @@ from app.services.dns_resolver import ( extract_dmarc_policy, get_default_provider, ) +from app.services.mta_sts import MTAStsResult, check_mta_sts_cached from app.services.report_persistence import ( delete_persisted_domain, hydrate_report_store_from_db, @@ -129,6 +130,22 @@ class DNSHealthResponse(BaseModel): recommendations: List[DNSHealthRecommendation] +class MTAStsResponse(BaseModel): + """MTA-STS posture result for a domain.""" + + status: str + dns_record: Optional[str] = None + policy_url: Optional[str] = None + policy_text: Optional[str] = None + mode: Optional[str] = None + max_age: Optional[int] = None + mx: List[str] = Field(default_factory=list) + errors: List[str] = Field(default_factory=list) + warnings: List[str] = Field(default_factory=list) + cached: bool = False + checked_at: Optional[str] = None + + class CloudflareZoneResponse(BaseModel): """Cloudflare zone available for import.""" @@ -445,6 +462,54 @@ def _enforcement_recommendation( ) +def _mta_sts_check(result: MTAStsResult) -> DNSHealthCheck: + evidence = [ + _record_evidence("MTA-STS TXT", result.dns_record, "#dns-records"), + _record_evidence("Policy URL", result.policy_url, "#mta-sts-posture"), + ] + if result.mode: + evidence.append(_record_evidence("Mode", result.mode, "#mta-sts-posture")) + if result.mx: + evidence.append(_record_evidence("MX patterns", ", ".join(result.mx), "#mta-sts-posture")) + message = ( + "MTA-STS DNS and HTTPS policy are valid." + if result.status == "pass" + else (result.errors[0] if result.errors else "MTA-STS posture needs attention.") + ) + return DNSHealthCheck( + key="mta_sts", + label="MTA-STS", + status=result.status, + message=message, + evidence=evidence, + ) + + +def _mta_sts_recommendation(result: MTAStsResult) -> Optional[DNSHealthRecommendation]: + if result.status == "pass" and not result.warnings: + return None + severity = "warning" if result.status == "pass" else "error" + title = "MTA-STS policy needs review" if result.status == "pass" else "Publish MTA-STS" + detail = ( + "; ".join(result.warnings) + if result.status == "pass" + else "; ".join(result.errors or ["MTA-STS is not configured or is not valid."]) + ) + action = ( + "Move the policy to mode: enforce once MX coverage is confirmed." + if result.status == "pass" + else "Publish _mta-sts TXT and a valid HTTPS policy at the well-known URL." + ) + return DNSHealthRecommendation( + type="mta_sts_review" if result.status == "pass" else "missing_mta_sts", + severity=severity, + title=title, + detail=detail, + action=action, + evidence=_mta_sts_check(result).evidence, + ) + + @router.get("/summary", response_model=DomainSummaryResponse) async def get_domains_summary(db: Session = Depends(get_db)): """ @@ -769,6 +834,12 @@ async def get_domain_dns_health( selectors=combined_selectors, refresh=refresh, ) + mta_sts_result, _, _ = await check_mta_sts_cached( + db, + provider, + domain_id, + refresh=refresh, + ) summary = store.get_domain_summary(domain_id) policy = extract_dmarc_policy(result.dmarc_record) or "none" checks = [ @@ -802,10 +873,11 @@ async def get_domain_dns_health( _record_evidence("DKIM TXT", result.dkim_record), ], ), + _mta_sts_check(mta_sts_result), ] recommendations: List[DNSHealthRecommendation] = [] for check in checks: - if check.status == "fail": + if check.status == "fail" and check.key != "mta_sts": recommendations.append( DNSHealthRecommendation( type=f"missing_{check.key}", @@ -817,6 +889,9 @@ async def get_domain_dns_health( ) ) recommendations.append(_enforcement_recommendation(policy, summary)) + mta_sts_recommendation = _mta_sts_recommendation(mta_sts_result) + if mta_sts_recommendation: + recommendations.append(mta_sts_recommendation) failed_checks = sum(1 for check in checks if check.status == "fail") health_status = ( @@ -833,6 +908,41 @@ async def get_domain_dns_health( ) +@router.get("/{domain_id}/dns/mta-sts", response_model=MTAStsResponse) +async def get_domain_mta_sts( + domain_id: str = Path(..., title="The domain ID or name"), + refresh: bool = Query(False, title="Refresh cached MTA-STS result"), + db: Session = Depends(get_db), +): + """Return cached MTA-STS DNS and HTTPS policy posture for a domain.""" + store = ReportStore.get_instance() + hydrate_report_store_from_db(db, store) + if not _domain_exists(db, store, domain_id): + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail="Domain not found", + ) + result, cached, checked_at = await check_mta_sts_cached( + db, + get_default_provider(db), + domain_id, + refresh=refresh, + ) + return MTAStsResponse( + status=result.status, + dns_record=result.dns_record, + policy_url=result.policy_url, + policy_text=result.policy_text, + mode=result.mode, + max_age=result.max_age, + mx=result.mx, + errors=result.errors, + warnings=result.warnings, + cached=cached, + checked_at=checked_at.isoformat(), + ) + + @router.get("/cloudflare/discover", response_model=List[CloudflareZoneResponse]) async def discover_cloudflare_domains(db: Session = Depends(get_db)): """Discover active Cloudflare zones visible to the configured API token.""" diff --git a/backend/app/services/dns_cache.py b/backend/app/services/dns_cache.py index 042f045..0f7ead7 100644 --- a/backend/app/services/dns_cache.py +++ b/backend/app/services/dns_cache.py @@ -8,6 +8,7 @@ from dataclasses import asdict from datetime import datetime, timedelta, timezone from typing import List, Tuple +from sqlalchemy.exc import IntegrityError from sqlalchemy.orm import Session from app.models.dns_cache import DNSCache @@ -74,18 +75,38 @@ async def resolve_domain_dns_cached( return _result_from_json(row.result_json), True, row.checked_at result = await provider.check_domain(domain, selectors=selectors) + payload = _result_to_json(result) if row is None: row = DNSCache( domain=domain, provider=provider_name, selectors_key=selectors_key, - result_json=_result_to_json(result), + result_json=payload, checked_at=now, ) db.add(row) else: - row.result_json = _result_to_json(result) + row.result_json = payload row.checked_at = now - db.commit() + + try: + db.commit() + except IntegrityError: + db.rollback() + row = ( + db.query(DNSCache) + .filter( + DNSCache.domain == domain, + DNSCache.provider == provider_name, + DNSCache.selectors_key == selectors_key, + ) + .first() + ) + if row is None: + raise + row.result_json = payload + row.checked_at = now + db.commit() + db.refresh(row) return result, False, row.checked_at diff --git a/backend/app/services/mta_sts.py b/backend/app/services/mta_sts.py new file mode 100644 index 0000000..cdb2c91 --- /dev/null +++ b/backend/app/services/mta_sts.py @@ -0,0 +1,216 @@ +"""MTA-STS posture checks for monitored domains.""" + +from __future__ import annotations + +import json +from dataclasses import asdict, dataclass, field +from datetime import datetime, timedelta, timezone +from typing import Any, Dict, List, Optional, Tuple + +import httpx +from sqlalchemy.exc import IntegrityError +from sqlalchemy.orm import Session + +from app.models.dns_cache import DNSCache +from app.services.dns_cache import DEFAULT_DNS_CACHE_TTL_SECONDS +from app.services.dns_resolver import BaseDNSProvider + +_CACHE_KEY = "mta-sts-v1" +_POLICY_TIMEOUT_SECONDS = 5.0 +_VALID_MODES = {"enforce", "testing", "none"} + + +@dataclass +class MTAStsResult: + """Operator-facing MTA-STS posture evidence.""" + + status: str = "fail" + dns_record: Optional[str] = None + policy_url: Optional[str] = None + policy_text: Optional[str] = None + mode: Optional[str] = None + max_age: Optional[int] = None + mx: List[str] = field(default_factory=list) + errors: List[str] = field(default_factory=list) + warnings: List[str] = field(default_factory=list) + + +def _utcnow_naive() -> datetime: + return datetime.now(timezone.utc).replace(tzinfo=None) + + +def _is_fresh(row: DNSCache, ttl_seconds: int, now: datetime) -> bool: + return row.checked_at >= now - timedelta(seconds=ttl_seconds) + + +def _result_from_json(value: str) -> MTAStsResult: + data = json.loads(value) + return MTAStsResult( + status=str(data.get("status") or "fail"), + dns_record=data.get("dns_record"), + policy_url=data.get("policy_url"), + policy_text=data.get("policy_text"), + mode=data.get("mode"), + max_age=data.get("max_age"), + mx=list(data.get("mx") or []), + errors=list(data.get("errors") or []), + warnings=list(data.get("warnings") or []), + ) + + +def parse_mta_sts_record(records: List[str]) -> Tuple[Optional[str], List[str], List[str]]: + """Return the selected MTA-STS TXT record, warnings, and errors.""" + sts_records = [record for record in records if record.lower().startswith("v=stsv1")] + if not sts_records: + return None, [], ["No _mta-sts TXT record was found."] + warnings = [] + if len(sts_records) > 1: + warnings.append("Multiple _mta-sts TXT records were found; publish exactly one.") + record = sts_records[0] + tags = { + part.split("=", 1)[0].strip().lower(): part.split("=", 1)[1].strip() + for part in record.split(";") + if "=" in part + } + errors = [] + if tags.get("v", "").lower() != "stsv1": + errors.append("The _mta-sts TXT record must start with v=STSv1.") + if not tags.get("id"): + errors.append("The _mta-sts TXT record must include a non-empty id tag.") + return record, warnings, errors + + +def parse_mta_sts_policy( # noqa: C901 + policy_text: str, +) -> Tuple[Dict[str, Any], List[str], List[str]]: + """Parse and validate an MTA-STS policy file.""" + data: Dict[str, Any] = {"mx": []} + for raw_line in policy_text.splitlines(): + line = raw_line.strip() + if not line or line.startswith("#") or ":" not in line: + continue + key, value = line.split(":", 1) + key = key.strip().lower() + value = value.strip() + if key == "mx": + data.setdefault("mx", []).append(value) + else: + data[key] = value + + errors = [] + warnings = [] + if str(data.get("version", "")).upper() != "STSV1": + errors.append("The policy file must contain version: STSv1.") + mode = str(data.get("mode", "")).lower() + if mode not in _VALID_MODES: + errors.append("The policy file must contain mode: enforce, testing, or none.") + elif mode in {"testing", "none"}: + warnings.append(f"MTA-STS policy is valid but not enforcing mail delivery ({mode}).") + try: + max_age = int(str(data.get("max_age", ""))) + if max_age <= 0: + errors.append("The policy max_age must be greater than zero.") + data["max_age"] = max_age + except ValueError: + errors.append("The policy file must contain an integer max_age value.") + if not data.get("mx"): + errors.append("The policy file must contain at least one mx entry.") + return data, warnings, errors + + +async def check_mta_sts(domain: str, provider: BaseDNSProvider) -> MTAStsResult: + """Resolve the MTA-STS TXT record and validate the HTTPS policy file.""" + result = MTAStsResult(policy_url=f"https://mta-sts.{domain}/.well-known/mta-sts.txt") + try: + records = await provider.lookup_txt(f"_mta-sts.{domain}") + except LookupError as exc: + result.errors.append(f"MTA-STS DNS lookup failed: {exc}") + return result + + record, warnings, errors = parse_mta_sts_record(records) + result.dns_record = record + result.warnings.extend(warnings) + result.errors.extend(errors) + if record is None: + return result + + try: + async with httpx.AsyncClient( + timeout=_POLICY_TIMEOUT_SECONDS, follow_redirects=False + ) as client: + response = await client.get(result.policy_url) + response.raise_for_status() + result.policy_text = response.text + except (httpx.RequestError, httpx.HTTPStatusError, httpx.TimeoutException) as exc: + result.errors.append(f"MTA-STS policy fetch failed: {exc}") + return result + + policy, policy_warnings, policy_errors = parse_mta_sts_policy(result.policy_text or "") + result.warnings.extend(policy_warnings) + result.errors.extend(policy_errors) + result.mode = policy.get("mode") + result.max_age = policy.get("max_age") + result.mx = list(policy.get("mx") or []) + result.status = "pass" if not result.errors else "fail" + return result + + +async def check_mta_sts_cached( + db: Session, + provider: BaseDNSProvider, + domain: str, + *, + ttl_seconds: int = DEFAULT_DNS_CACHE_TTL_SECONDS, + refresh: bool = False, +) -> Tuple[MTAStsResult, bool, datetime]: + """Resolve MTA-STS posture, reusing the shared DNS cache semantics.""" + now = _utcnow_naive() + provider_name = f"{provider.__class__.__name__}:mta-sts" + row = ( + db.query(DNSCache) + .filter( + DNSCache.domain == domain, + DNSCache.provider == provider_name, + DNSCache.selectors_key == _CACHE_KEY, + ) + .first() + ) + if row and not refresh and _is_fresh(row, ttl_seconds, now): + return _result_from_json(row.result_json), True, row.checked_at + + result = await check_mta_sts(domain, provider) + payload = json.dumps(asdict(result), sort_keys=True, separators=(",", ":")) + if row is None: + row = DNSCache( + domain=domain, + provider=provider_name, + selectors_key=_CACHE_KEY, + result_json=payload, + checked_at=now, + ) + db.add(row) + else: + row.result_json = payload + row.checked_at = now + + try: + db.commit() + except IntegrityError: + db.rollback() + row = ( + db.query(DNSCache) + .filter( + DNSCache.domain == domain, + DNSCache.provider == provider_name, + DNSCache.selectors_key == _CACHE_KEY, + ) + .first() + ) + if row is None: + raise + row.result_json = payload + row.checked_at = now + db.commit() + + db.refresh(row) + return result, False, row.checked_at diff --git a/backend/app/templates/domain_details.html b/backend/app/templates/domain_details.html index 38e0921..f1468c8 100644 --- a/backend/app/templates/domain_details.html +++ b/backend/app/templates/domain_details.html @@ -143,7 +143,7 @@ {% endcall %} {% endcall %} {% call card_content() %} -
+