#!/usr/bin/env python3

"""A very simple DNS resolver in Python, using the dnspython
library. This is more or less the equivalent of Julia Evans' nice Go
version
<https://jvns.ca/blog/2022/02/01/a-dns-resolver-in-80-lines-of-go/>. The
goal is to teach DNS, not to be a real production-ready resolver. (It
has many limits and weaknesses.)"""

# https://www.dnspython.org/
import dns.message
import dns.query
import dns.rdata
import dns.rdataclass

import sys

root = "2001:7fd::1"  # K-root server. A real resolver would do priming (RFC 8109)
verbose = True
MAXDEPTH = 6
ADDRESSFAMILYNAMESERVERS = dns.rdatatype.A

def usage(msg=None):
    if msg is not None:
        print(msg, file=sys.stderr)
    print("Usage: %s domain-name [dns-record-type]" % sys.argv[0])

def query(qname, qtype=dns.rdatatype.AAAA, server=root, depth=0):
    """Qname and qtype MUST be int, not strings. Returns a tuple {result,
       value}."""
    if depth > MAXDEPTH:
        return ("TOO DEEP", MAXDEPTH)
    if verbose:
        print("Querying %s/%s at %s, depth %i" % (qname, dns.rdatatype.to_text(qtype), server, depth))
    msg = dns.message.make_query(qname, qtype,
                                 use_edns=0, want_dnssec=False, payload=4096) 
    msg.flags = 0 # No RD flag
    reply = dns.query.udp(msg, server, timeout=2)
    if reply.rcode() == dns.rcode.NXDOMAIN:
        return ("NXDOMAIN", None)
    elif reply.rcode() == dns.rcode.SERVFAIL:
        return ("SERVFAIL", None) # A real resolver would try another
                                  # authoritative name server (same
                                  # thing for the next error code)
    elif reply.rcode() == dns.rcode.REFUSED:
        return ("REFUSED", server)
    elif reply.rcode() != dns.rcode.NOERROR:
        return ("UNKNOWN ERROR", reply.rcode())
    answer = reply.get_rrset(dns.message.ANSWER, qname, dns.rdataclass.IN, qtype)
    if answer is not None:
        return ("OK", answer[0])
    else:
        answer = reply.get_rrset(dns.message.ANSWER, qname, dns.rdataclass.IN, dns.rdatatype.CNAME)
        if answer is not None:
            return query(answer[0].target, qtype, root, 0)
        if len(reply.section_from_number(dns.message.ADDITIONAL)) > 0:
            server = reply.section_from_number(dns.message.ADDITIONAL)[0]
            # We are lazy, we use only the first server of the list. A
            # real resolver would choose at random, then memorize the
            # fastest one, and be able to switch to another one if
            # there is a timeout.
            server_addr = server[0].address # Note the address can be
                                            # v4 or v6, whatever
                                            # ADDRESSFAMILYNAMESERVERS
                                            # is.
        else:
            server_name = reply.section_from_number(dns.message.AUTHORITY)
            if server_name[0].rdtype == dns.rdatatype.SOA:
                return ("NODATA", None)
            # We are lazy, we use only the first server of the list. A
            # real resolver would choose at random, then memorize the
            # fastest one, and be able to switch to another one if
            # there is a timeout.
            result = query(server_name[0][0].target, ADDRESSFAMILYNAMESERVERS, root, 0)
            if result[0] == "OK":
                server_addr = result[1].address
            else:
                return ("NOGLUE", server_name[0][0].target)
        return query(qname, qtype, server_addr, depth+1)
    
if len(sys.argv) > 3 or len(sys.argv) < 2:
    usage()
    sys.exit(1)
qname = dns.name.from_text(sys.argv[1])
if len(sys.argv) == 2:
    qtype = dns.rdatatype.AAAA
else:
    qtype = dns.rdatatype.from_text(sys.argv[2])
r = query(qname, qtype, root)
print("Value of %s/%s is %s" % (qname, dns.rdatatype.to_text(qtype), r))
