diff --git a/dns_compare.py b/dns_compare.py index 46558ff..f529e6f 100755 --- a/dns_compare.py +++ b/dns_compare.py @@ -1,28 +1,81 @@ #!/usr/bin/env python import argparse +from shlex import join import dns.query +import dns.rdatatype +import dns.name +import json import re +# Get list of supported record types +supported_types = sorted([dns.rdatatype.to_text(rdtype) for rdtype in dns.rdatatype.RdataType]) + parser = argparse.ArgumentParser() parser.add_argument('--domain', help='Helps if all records to query are on the same domain.') parser.add_argument('--nameservers', nargs='+', help='IPv4 or IPv6 address of nameserver(s) to be queried.') -parser.add_argument('--records', nargs='+', help='Hostname to query.') +parser.add_argument('--records', nargs='*', help='Hostname(s) to query. If not provided with --domain, queries the domain itself.') +parser.add_argument('--record-type', default='A', help='DNS record type to query. Default: A') +parser.add_argument('--list-record-types', action='store_true', help='Print list of supported record types and exit.') args = parser.parse_args() +# Handle listing record types +if args.list_record_types: + print("Supported DNS record types:") + for rtype in supported_types: + print(f" {rtype}") + exit(0) + +# Convert record type string to dns.rdatatype constant +try: + record_type = dns.rdatatype.from_text(args.record_type) +except dns.rdatatype.UnknownRdatatype: + print(f"Error: Unknown record type '{args.record_type}'") + print(f"\nSupported record types:") + print((', ').join(supported_types)) + exit(1) + +# Build the output structure +output = {} + +# Determine what to query if args.domain: - for rec in args.records: - qname = dns.name.from_text(f'{rec}.{args.domain}') - q = dns.message.make_query(qname, dns.rdatatype.A) + # If records are provided, use them as subdomains; otherwise use the domain itself + records_to_query = args.records if args.records else [args.domain] + + for rec in records_to_query: + if rec == args.domain: + full_domain = args.domain + else: + full_domain = f'{rec}.{args.domain}' + qname = dns.name.from_text(full_domain) + q = dns.message.make_query(qname, record_type) + output[full_domain] = {"records": {}} for ns in args.nameservers: r = dns.query.udp(q, ns) - x = re.findall(fr'{str(rec)}.*', str(r)) - print(x, ns) + # Extract the resolved IP address from the response + answers = r.answer + if answers: + for rrset in answers: + for item in rrset: + output[full_domain]["records"][ns] = str(item) else: + # Query each record as a full domain + if not args.records: + print("Error: Either --domain or --records must be provided") + exit(1) + for rec in args.records: qname = dns.name.from_text(rec) - q = dns.message.make_query(qname, dns.rdatatype.A) + q = dns.message.make_query(qname, record_type) + output[rec] = {"records": {}} for ns in args.nameservers: r = dns.query.udp(q, ns) - x = re.findall(fr'{str(rec)}.*', str(r)) - print(x, ns) + # Extract the resolved IP address from the response + answers = r.answer + if answers: + for rrset in answers: + for item in rrset: + output[rec]["records"][ns] = str(item) + +print(json.dumps(output, indent=4))