#!/usr/bin/env python3
# setptr -- change rDNS PTR records using DNS UPDATE
import argparse
import dns.name
import dns.resolver
import ipaddress
import logging
import subprocess

RPZ_ZONE = dns.name.from_text("rpz.nullroute.lt.")

def confirm(text):
    return input(f"{text} ").startswith("y")

def chase_one_cname(name):
    try:
        ans = dns.resolver.resolve(name, "CNAME")
    except dns.resolver.NoAnswer:
        return name
    except dns.resolver.NXDOMAIN:
        return name
    else:
        return ans.rrset[0].target

def master_for_zone(zone):
    # Because Bind, the SOA response might have an upper-case tail, e.g.
    # "star.nullroute.LT" for our RPZ zone, and nsupdate blindly gives that to
    # Kerberos which fails. So we manually look up the master server here and
    # casefold the result.
    ans = dns.resolver.resolve(zone, "SOA")
    return ans.rrset[0].mname.canonicalize()

def send_nsupdate(zone, cmds, *, server=None, use_gss=True, debug=False):
    if not server:
        server = master_for_zone(zone)
    logging.debug(f"Sending update to server \"{server}\"")

    cmds = [f"server {server}\n",
            f"zone {zone}\n",
            *cmds,
            f"send\n"]

    nsupdate_args = ["nsupdate"]
    if debug:
        nsupdate_args += ["-d"]
    if use_gss:
        nsupdate_args += ["-g"]

    subprocess.run(nsupdate_args,
                   input="".join(cmds).encode(),
                   check=True)

parser = argparse.ArgumentParser()
parser.add_argument("-e", "--edit", action="store_true",
                        help="interactively edit the zone")
parser.add_argument("-z", "--zone",
                        help="override automatically determined rDNS zone")
parser.add_argument("-R", "--rpz", action="store_true",
                        help="update the Response Policy Zone")
parser.add_argument("-l", "--ttl", type=int, default=3600,
                        help="set the TTL for created records")
parser.add_argument("-n", "--dry-run", action="store_true",
                        help="only show updates that would be done")
parser.add_argument("-x", "--no-gss", action="store_true",
                        help="disable Kerberos (GSS-TSIG) authentication")
parser.add_argument("-d", "--debug", action="count", default=0,
                        help="enable nsupdate debugging")
parser.add_argument("-v", "--verbose", action="store_true",
                        help="show detailed information")
parser.add_argument("address",
                        help="IP address to update the PTR for")
parser.add_argument("target", nargs="?",
                        help="PTR target domain name, or \".\" to remove")
args = parser.parse_args()

logging.basicConfig(level=[logging.INFO, logging.DEBUG][args.verbose],
                    format="%(message)s")

if args.zone and args.rpz:
    exit(f"setptr: --rpz and --zone are mutually exclusive")

if args.edit:
    if args.target:
        exit(f"setptr: target cannot be specified when using interactive mode")
    cmd = ["vireverse"]
    if args.debug:
        cmd += ["--debug"]
    if args.verbose:
        cmd += ["--verbose"]
    if args.no_gss:
        cmd += ["--no-gss"]
    if args.rpz:
        cmd += ["--rpz"]
    cmd += [args.address]
    exit(subprocess.call(cmd))
else:
    if not args.target:
        exit(f"setptr: target must be specified when not using interactive mode")

if "/" in args.address:
    addrs = ipaddress.ip_network(args.address)
    target = dns.name.from_text(args.target)

    if addrs.prefixlen < 24:
        # Safety check to prevent excessive deletions
        exit(f"setptr: address prefix too short")
    if str(target) != ".":
        exit(f"setptr: prefixes only allowed when removing records")

    rnames = [dns.name.from_text(addr.reverse_pointer)
              for addr in addrs]

    if args.rpz:
        zone = RPZ_ZONE
        rnames = [rname.relativize(dns.name.root) + zone
                  for rname in rnames]
    elif args.zone:
        zone = dns.name.from_text(args.zone)
    else:
        # Detect IPv4 classless delegations
        rname = chase_one_cname(rnames[0])
        zone = dns.resolver.zone_for_name(rname)

    logging.debug(f"Updating zone \"{zone}\"")

    cmds = [f"del {rname} {args.ttl} PTR\n"
            for rname in rnames]

    print(f"Removing {len(rnames)} PTRs for {addrs}")
    if not confirm("Continue?"):
        exit("Changes discarded.")

else:
    addr = ipaddress.ip_address(args.address)
    rname = dns.name.from_text(addr.reverse_pointer)
    target = dns.name.from_text(args.target)

    if args.rpz:
        zone = RPZ_ZONE
        rname = rname.relativize(dns.name.root) + zone
    elif args.zone:
        zone = dns.name.from_text(args.zone)
    else:
        # Detect IPv4 classless delegations
        rname = chase_one_cname(rname)
        zone = dns.resolver.zone_for_name(rname)

    logging.debug(f"Updating zone \"{zone}\"")

    if str(target) == ".":
        print(f"Removing PTR for [{addr}]")
        cmds = [f"del {rname} {args.ttl} PTR\n"]
    else:
        print(f"Changing PTR for [{addr}] to \"{target}\"")
        cmds = [f"del {rname} {args.ttl} PTR\n",
                f"add {rname} {args.ttl} PTR {target}\n"]

if args.debug >= 1:
    print("Commands:")
    print("\t" + "\t".join(cmds), end="")

if not args.dry_run:
    send_nsupdate(zone, cmds,
                  use_gss=(not args.no_gss),
                  debug=(args.debug >= 2))
