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>
This commit is contained in:
copilot-swe-agent[bot]
2026-03-29 18:53:35 +00:00
parent 4e4db14d36
commit 0d1a4fdac3
6 changed files with 1096 additions and 39 deletions
+252
View File
@@ -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
+269
View File
@@ -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)