tests/unit/test_dns_cname.py (view raw)
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 |
"""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"])
|