"""DNS CNAME chain / cycle unit tests with mocked resolver.""" from __future__ import annotations from types import SimpleNamespace from unittest.mock import AsyncMock, patch import pytest from app.services.dns import compute_fingerprint, resolve_cname_chain class _Rdata: def __init__(self, target: str): self.target = target def to_text(self) -> str: return self.target @pytest.mark.asyncio async def test_cname_chain_and_cycle(): calls = {"n": 0} async def fake_resolve(name, rdtype): calls["n"] += 1 host = str(name).rstrip(".").lower() if not isinstance(name, str) else name.lower() # dnspython passes string hostname in our code host = name if isinstance(name, str) else str(name) host = host.rstrip(".").lower() if rdtype == "CNAME": mapping = { "a.example": "b.example", "b.example": "a.example", # cycle } if host in mapping: return [_Rdata(mapping[host])] raise Exception("no cname") if rdtype == "A": return [] if rdtype == "AAAA": return [] if rdtype == "TXT": raise Exception("no txt") raise Exception("unexpected") with patch("app.services.dns.dns.asyncresolver.Resolver") as resolver_cls: inst = resolver_cls.return_value inst.resolve = AsyncMock(side_effect=fake_resolve) snap = await resolve_cname_chain("a.example") assert snap.cycle_detected is True assert "b.example" in snap.cname_chain @pytest.mark.asyncio async def test_cname_terminal_fingerprint(): async def fake_resolve(name, rdtype): host = name if isinstance(name, str) else str(name) host = host.rstrip(".").lower() if rdtype == "CNAME": if host == "start.example": return [_Rdata("end.example.")] raise Exception("nx") if rdtype == "A": if host == "end.example": obj = SimpleNamespace(to_text=lambda: "203.0.113.10") return [obj] return [] if rdtype == "AAAA": return [] if rdtype == "TXT": raise Exception("no") raise Exception("x") with patch("app.services.dns.dns.asyncresolver.Resolver") as resolver_cls: inst = resolver_cls.return_value inst.resolve = AsyncMock(side_effect=fake_resolve) snap = await resolve_cname_chain("start.example") assert snap.terminal_cname == "end.example" assert snap.fingerprint == compute_fingerprint("end.example", ["203.0.113.10"])