app/services/dns.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 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 |
from __future__ import annotations
import hashlib
import logging
import socket
from dataclasses import dataclass, field
import dns.asyncresolver
from app.services.ssrf import is_public_ip
logger = logging.getLogger(__name__)
MAX_CNAME_DEPTH = 10
@dataclass
class DnsSnapshot:
cname_chain: list[str] = field(default_factory=list)
terminal_cname: str | None = None
ipv4: list[str] = field(default_factory=list)
ipv6: list[str] = field(default_factory=list)
cycle_detected: bool = False
fingerprint: str | None = None
async def resolve_cname_chain(hostname: str, *, max_depth: int = MAX_CNAME_DEPTH) -> DnsSnapshot:
snapshot = DnsSnapshot()
resolver = dns.asyncresolver.Resolver()
current = hostname.rstrip(".").lower()
seen: set[str] = set()
for _ in range(max_depth):
if current in seen:
snapshot.cycle_detected = True
break
seen.add(current)
try:
answer = await resolver.resolve(current, "CNAME")
except Exception: # noqa: BLE001 — end of CNAME chain
break
if not answer:
break
target = str(answer[0].target).rstrip(".").lower()
snapshot.cname_chain.append(target)
current = target
if snapshot.cname_chain:
snapshot.terminal_cname = snapshot.cname_chain[-1]
lookup_host = snapshot.terminal_cname or hostname.rstrip(".").lower()
snapshot.ipv4 = await _resolve_addresses(resolver, lookup_host, "A")
snapshot.ipv6 = await _resolve_addresses(resolver, lookup_host, "AAAA")
snapshot.fingerprint = compute_fingerprint(
snapshot.terminal_cname, snapshot.ipv4 + snapshot.ipv6
)
return snapshot
async def _resolve_addresses(
resolver: dns.asyncresolver.Resolver, hostname: str, rdtype: str
) -> list[str]:
try:
answer = await resolver.resolve(hostname, rdtype)
except Exception: # noqa: BLE001
return []
out: list[str] = []
for rdata in answer:
ip = rdata.to_text()
if is_public_ip(ip):
out.append(ip)
return sorted(set(out))
def compute_fingerprint(terminal_cname: str | None, public_ips: list[str]) -> str | None:
if terminal_cname:
material = f"cname:{terminal_cname.lower()}"
elif public_ips:
material = "ips:" + ",".join(sorted(set(public_ips)))
else:
return None
return hashlib.sha256(material.encode("utf-8")).hexdigest()[:32]
def host_has_ipv6() -> bool:
"""Best-effort detection of local IPv6 connectivity."""
try:
sock = socket.socket(socket.AF_INET6, socket.SOCK_DGRAM)
try:
sock.connect(("2001:4860:4860::8888", 53))
return True
finally:
sock.close()
except OSError:
return False
def lookup_asn(ip: str, db_path: str) -> tuple[int | None, str | None, str | None]:
"""Optional MaxMind ASN lookup. Returns (asn, network_name, country_code)."""
if not db_path:
return None, None, None
try:
import geoip2.database # type: ignore
except ImportError:
return None, None, None
try:
with geoip2.database.Reader(db_path) as reader:
# Prefer ASN db; country may be unavailable in ASN-only files
try:
asn_resp = reader.asn(ip)
asn = asn_resp.autonomous_system_number
org = asn_resp.autonomous_system_organization
except Exception: # noqa: BLE001
asn, org = None, None
country = None
try:
country = reader.country(ip).country.iso_code # type: ignore[attr-defined]
except Exception: # noqa: BLE001
country = None
return asn, org, country
except Exception: # noqa: BLE001
logger.debug("GeoIP lookup unavailable for %s", ip, exc_info=True)
return None, None, None
|