"""DNS propagation checks for the ACME dns-01 challenge.""" from __future__ import annotations import logging import socket import time import dns.exception import dns.flags import dns.message import dns.query import dns.rdatatype import dns.resolver log = logging.getLogger(__name__) QUERY_TIMEOUT = 5.0 def _resolver(nameservers: list[str] | None = None) -> dns.resolver.Resolver: res = dns.resolver.Resolver(configure=not nameservers) if nameservers: res.nameservers = nameservers res.lifetime = QUERY_TIMEOUT * 2 res.timeout = QUERY_TIMEOUT return res def authoritative_servers(zone: str, fallback_resolvers: list[str]) -> list[str]: """IP addresses of the authoritative nameservers of *zone*.""" ips: list[str] = [] names: list[str] = [] for res in (_resolver(), _resolver(fallback_resolvers)): try: answer = res.resolve(zone, "NS", raise_on_no_answer=False) names = sorted({str(rdata.target).rstrip(".") for rdata in answer}) if names: break except (dns.exception.DNSException, OSError) as exc: log.debug("NS lookup for %s failed: %s", zone, exc) lookup = _resolver(fallback_resolvers) for name in names: for rtype in ("A", "AAAA"): try: for rdata in lookup.resolve(name, rtype, raise_on_no_answer=False): ips.append(str(rdata)) except (dns.exception.DNSException, OSError): continue unique = list(dict.fromkeys(ips)) usable = [ip for ip in unique if reachable(ip)] dropped = [ip for ip in unique if ip not in usable] if dropped: # Typical inside a container without IPv6: the NS have AAAA records we # cannot talk to. Their IPv4 addresses answer the same zone data. log.info("Skipping unreachable nameserver address(es): %s", ", ".join(dropped)) if usable: log.info("Authoritative nameservers for %s: %s (%s)", zone, ", ".join(names), ", ".join(usable)) else: log.warning("Could not reach any authoritative nameserver for %s, using public resolvers.", zone) return usable def reachable(server: str) -> bool: """Can this machine send a UDP packet to *server* at all? ``connect()`` on a UDP socket sends nothing, it only resolves the route - exactly what we need to weed out IPv6 addresses on an IPv4-only host. """ family = socket.AF_INET6 if ":" in server else socket.AF_INET try: with socket.socket(family, socket.SOCK_DGRAM) as sock: sock.connect((server, 53)) return True except OSError as exc: log.debug("Nameserver %s is not reachable from here: %s", server, exc) return False def txt_values(name: str, server: str) -> set[str] | None: """TXT values for *name* as seen by *server*. None if the server is unusable.""" query = dns.message.make_query(name, dns.rdatatype.TXT) values: set[str] = set() try: response = dns.query.udp(query, server, timeout=QUERY_TIMEOUT) if response.flags & dns.flags.TC: response = dns.query.tcp(query, server, timeout=QUERY_TIMEOUT) except dns.exception.DNSException as exc: log.debug("TXT query %s @%s failed: %s", name, server, exc) return values except OSError as exc: # No route to that address (e.g. IPv6 without IPv6 connectivity). log.debug("TXT query %s @%s not possible: %s", name, server, exc) return None for rrset in response.answer: if rrset.rdtype != dns.rdatatype.TXT: continue for rdata in rrset: values.add(b"".join(rdata.strings).decode("utf-8", "replace")) return values def wait_for_txt( expected: dict[str, set[str]], servers: list[str], timeout: int, interval: int, ) -> bool: """Block until every server serves every expected TXT value, or *timeout* expires.""" if not servers: log.warning("No DNS servers to verify against - waiting %ss blindly.", interval * 2) time.sleep(interval * 2) return False deadline = time.monotonic() + timeout alive = list(servers) attempt = 0 while True: attempt += 1 missing: list[str] = [] unusable: list[str] = [] for name, wanted in expected.items(): for server in alive: seen = txt_values(name, server) if seen is None: unusable.append(server) elif not wanted <= seen: missing.append(f"{name} @{server}") if unusable: for server in dict.fromkeys(unusable): alive.remove(server) log.warning("Nameserver %s cannot be queried from here - ignoring it.", server) if not alive: log.error("None of the nameservers could be queried at all.") return False continue # judge the round again, now without the dead servers if not missing: log.info("DNS propagation confirmed on all %d nameserver(s).", len(alive)) return True remaining = deadline - time.monotonic() if remaining <= 0: log.error("DNS propagation timed out. Still missing: %s", ", ".join(sorted(set(missing)))) return False log.info( "Attempt %d: waiting for DNS propagation (%d pending, %ds left)...", attempt, len(set(missing)), int(remaining), ) time.sleep(min(interval, max(1, int(remaining))))