#!/usr/bin/env python3
# Copyright (C) 2026 - License: GNU General Public License v2
"""Special agent 'dnssec_health'.

Reports the DNSSEC status of arbitrary domains (not necessarily mail domains):
  - signed:    whether the domain publishes DNSKEY records
  - validated: whether the queried recursive resolver validated the
               signatures (the AD / Authenticated Data bit in the response)

Each configured domain is queried against every configured DNS server, so a
separate record (domain x nameserver) is emitted. The 'validated' result
therefore reflects whether that specific resolver performs DNSSEC validation.

Implemented with the Python standard library only - no dnspython required.
"""

from __future__ import annotations

import argparse
import json
import random
import socket
import struct
import sys
from dataclasses import dataclass, field

TYPE_SOA = 6
TYPE_DNSKEY = 48

RCODE_NOERROR = 0
RCODE_NXDOMAIN = 3

_RCODE_NAMES = {
    0: "NOERROR",
    1: "FORMERR",
    2: "SERVFAIL",
    3: "NXDOMAIN",
    4: "NOTIMP",
    5: "REFUSED",
}


class DnsError(Exception):
    """Transport-level DNS error (timeout, network trouble, malformed reply)."""


class _Truncated(Exception):
    pass


@dataclass(frozen=True)
class DnsResponse:
    rcode: int
    authenticated: bool  # AD bit
    answer_count: int

    @property
    def rcode_name(self) -> str:
        return _RCODE_NAMES.get(self.rcode, f"RCODE{self.rcode}")


def encode_name(name: str) -> bytes:
    out = b""
    for label in name.rstrip(".").split("."):
        raw = label.encode("idna") if any(ord(c) > 127 for c in label) else label.encode("ascii")
        if not 0 < len(raw) < 64:
            raise DnsError(f"invalid label in {name!r}")
        out += bytes([len(raw)]) + raw
    return out + b"\x00"


def build_query(name: str, rtype: int, txid: int) -> bytes:
    # RD=1 (0x0100) + AD=1 (0x0020): request DNSSEC-validated status.
    flags = 0x0120
    header = struct.pack(">HHHHHH", txid, flags, 1, 0, 0, 0)
    return header + encode_name(name) + struct.pack(">HH", rtype, 1)


def parse_response(msg: bytes, expected_txid: int) -> DnsResponse:
    if len(msg) < 12:
        raise DnsError("short response")
    txid, flags, _qd, ancount, _ns, _ar = struct.unpack(">HHHHHH", msg[:12])
    if txid != expected_txid:
        raise DnsError("transaction id mismatch")
    if flags & 0x0200:  # TC bit
        raise _Truncated()
    return DnsResponse(
        rcode=flags & 0x000F,
        authenticated=bool(flags & 0x0020),
        answer_count=ancount,
    )


def _recv_exact(sock: socket.socket, count: int) -> bytes:
    data = b""
    while len(data) < count:
        chunk = sock.recv(count - len(data))
        if not chunk:
            raise DnsError("connection closed")
        data += chunk
    return data


class Resolver:
    def __init__(self, nameservers: list[str], timeout: float, retries: int = 2) -> None:
        if not nameservers:
            raise DnsError("no nameservers configured")
        self.nameservers = nameservers
        self.timeout = timeout
        self.retries = retries

    def query(self, name: str, rtype: int) -> DnsResponse:
        last_error: Exception = DnsError("no nameserver reachable")
        for server in self.nameservers:
            for _attempt in range(self.retries):
                txid = random.randint(0, 0xFFFF)
                request = build_query(name, rtype, txid)
                try:
                    return self._query_udp(server, request, txid)
                except _Truncated:
                    try:
                        return self._query_tcp(server, request, txid)
                    except (OSError, DnsError) as exc:
                        last_error = exc
                except (OSError, DnsError) as exc:
                    last_error = exc
        raise DnsError(str(last_error) or type(last_error).__name__)

    def _family(self, server: str) -> int:
        return socket.AF_INET6 if ":" in server else socket.AF_INET

    def _query_udp(self, server: str, request: bytes, txid: int) -> DnsResponse:
        with socket.socket(self._family(server), socket.SOCK_DGRAM) as sock:
            sock.settimeout(self.timeout)
            sock.sendto(request, (server, 53))
            msg, _addr = sock.recvfrom(4096)
        return parse_response(msg, txid)

    def _query_tcp(self, server: str, request: bytes, txid: int) -> DnsResponse:
        with socket.create_connection((server, 53), timeout=self.timeout) as sock:
            sock.settimeout(self.timeout)
            sock.sendall(struct.pack(">H", len(request)) + request)
            (length,) = struct.unpack(">H", _recv_exact(sock, 2))
            msg = _recv_exact(sock, length)
        return parse_response(msg, txid)


def system_nameservers() -> list[str]:
    servers: list[str] = []
    try:
        with open("/etc/resolv.conf", encoding="utf-8") as handle:
            for line in handle:
                parts = line.split()
                if len(parts) >= 2 and parts[0] == "nameserver":
                    servers.append(parts[1])
    except OSError:
        pass
    return servers or ["127.0.0.1"]


def collect_dnssec(resolver: Resolver, domain: str, nameserver: str) -> dict:
    signed = None
    dnskey_error = None
    try:
        dnskey = resolver.query(domain, TYPE_DNSKEY)
        if dnskey.rcode == RCODE_NOERROR:
            signed = dnskey.answer_count > 0
        elif dnskey.rcode == RCODE_NXDOMAIN:
            signed = False
        else:
            dnskey_error = dnskey.rcode_name
    except DnsError as exc:
        dnskey_error = str(exc)

    validated = None
    soa_error = None
    try:
        soa = resolver.query(domain, TYPE_SOA)
        if soa.rcode == RCODE_NOERROR:
            validated = soa.authenticated
        else:
            soa_error = soa.rcode_name
    except DnsError as exc:
        soa_error = str(exc)

    return {
        "domain": domain,
        "nameserver": nameserver,
        "signed": signed,
        "validated": validated,
        "error": dnskey_error or soa_error,
    }


@dataclass
class Args:
    domains: list[str] = field(default_factory=list)
    nameservers: list[str] = field(default_factory=list)
    timeout: float = 5.0
    debug: bool = False


def parse_arguments(argv: list[str]) -> Args:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--domain", action="append", default=[], dest="domains", metavar="DOMAIN")
    parser.add_argument(
        "--nameserver", action="append", default=[], dest="nameservers", metavar="IP"
    )
    parser.add_argument("--timeout", type=float, default=5.0)
    parser.add_argument("--debug", action="store_true")
    return Args(**vars(parser.parse_args(argv)))


def main(argv: list[str] | None = None) -> int:
    args = parse_arguments(sys.argv[1:] if argv is None else argv)
    nameservers = args.nameservers or system_nameservers()
    try:
        if args.domains:
            sys.stdout.write("<<<dnssec_health:sep(0)>>>\n")
            for domain in args.domains:
                for nameserver in nameservers:
                    resolver = Resolver([nameserver], timeout=args.timeout)
                    record = collect_dnssec(resolver, domain, nameserver)
                    sys.stdout.write(json.dumps(record) + "\n")
    except Exception as exc:  # pylint: disable=broad-except
        if args.debug:
            raise
        sys.stderr.write(f"agent_dnssec_health: {exc}\n")
        return 1
    return 0


if __name__ == "__main__":
    sys.exit(main())
