all repos — rastro @ 6738696441c3bf8d4382193143952f0e7cdac6c5

BitTorrent tracker!

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"])