diff --git a/dns_compare.py b/dns_compare.py index c5fbb34..a76b235 100755 --- a/dns_compare.py +++ b/dns_compare.py @@ -1,24 +1,70 @@ #!/usr/bin/env python import argparse -from shlex import join -import dns.query -import dns.rdatatype -import dns.name import json -import re +import os +import time +from concurrent.futures import ThreadPoolExecutor +from dns import message +from dns import query +from dns import rdatatype +from dns import name # Get list of supported record types -supported_types = sorted([dns.rdatatype.to_text(rdtype) for rdtype in dns.rdatatype.RdataType]) +supported_types = sorted([rdatatype.to_text(rdtype) for rdtype in rdatatype.RdataType]) + +# Dictionary of public DNS servers +builtin_nameservers = { + 'google': ['8.8.8.8', '8.8.4.4'], + 'cloudflare': ['1.1.1.1', '1.0.0.1', '1.1.1.2', '1.0.0.2'], + 'quad9': ['9.9.9.9', '149.112.112.112'], + 'gte': ['4.2.2.1', '4.2.2.2', '4.2.2.3', '4.2.2.4', '4.2.2.5', '4.2.2.6'], + 'opendns': ['208.67.222.222', '208.67.220.220', '208.67.222.123', '208.67.220.123'], + 'verisign': ['64.6.64.6', '64.6.65.6'], + 'comodo': ['8.26.56.26', '8.20.247.20'], + 'level3': ['209.244.0.3', '209.244.0.4'], +} 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('--ns-builtin', nargs='+', help=f'Built-in public DNS servers. Options: {", ".join(builtin_nameservers.keys())}') 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() +# Build nameserver map: IP -> provider name +nameserver_map = {} # ip -> provider_key +builtin_providers = [] # list of provider keys used + +# Resolve builtin nameservers to actual IPs +if args.ns_builtin: + for provider in args.ns_builtin: + if provider.lower() not in builtin_nameservers: + print(f"Error: Unknown DNS provider '{provider}'") + print(f"Available providers: {', '.join(builtin_nameservers.keys())}") + exit(1) + provider_key = provider.lower() + builtin_providers.append(provider_key) + for ip in builtin_nameservers[provider_key]: + nameserver_map[ip] = provider_key + +# Track manual nameservers separately +manual_nameservers = [] +if args.nameservers: + manual_nameservers = args.nameservers + for ip in args.nameservers: + nameserver_map[ip] = "manual_selection" + +# Combine all nameservers for querying +all_nameservers = list(nameserver_map.keys()) + +# Use resolved nameservers, or error if none provided +if not all_nameservers: + print("Error: Either --nameservers or --ns-builtin must be provided") + exit(1) + # Handle listing record types if args.list_record_types: print("Supported DNS record types:") @@ -27,54 +73,106 @@ if args.list_record_types: # Convert record type string to dns.rdatatype constant try: - record_type = dns.rdatatype.from_text(args.record_type) -except dns.rdatatype.UnknownRdatatype: + record_type = rdatatype.from_text(args.record_type) +except 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 = {} +# Function to perform a single DNS query +def perform_query(full_domain, ns, record_type): + """Perform a DNS query and return results""" + try: + qname = name.from_text(full_domain) + q = message.make_query(qname, record_type) + r = query.udp(q, ns) + + results = [] + answers = r.answer + if answers: + for rrset in answers: + for item in rrset: + results.append(str(item)) + + return (full_domain, ns, results) + except Exception as e: + return (full_domain, ns, None) + +# Start timer +start_time = time.time() # Determine what to query if args.domain: # If records are provided, use them as subdomains; otherwise use the domain itself records_to_query = args.records if args.records else [args.domain] + full_domains = [] 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) - # 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) + full_domains.append(full_domain) else: # Query each record as a full domain if not args.records: print("Error: Either --domain or --records must be provided") exit(1) + full_domains = args.records + +# Build list of query tasks +query_tasks = [] +for full_domain in full_domains: + for ns in all_nameservers: + query_tasks.append((full_domain, ns, record_type)) + +# Determine thread count: CPU count - 1, minimum 1 +thread_count = max(1, os.cpu_count() - 1) + +# Execute queries in parallel +results_map = {} # (full_domain, ns) -> resolved_ip +with ThreadPoolExecutor(max_workers=thread_count) as executor: + futures = [executor.submit(perform_query, full_domain, ns, record_type) + for full_domain, ns, _ in query_tasks] - for rec in args.records: - qname = dns.name.from_text(rec) - q = dns.message.make_query(qname, record_type) - output[rec] = {"records": {}} - for ns in args.nameservers: - r = dns.query.udp(q, 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) + for future in futures: + full_domain, ns, resolved_ips = future.result() + if resolved_ips: + results_map[(full_domain, ns)] = resolved_ips[0] # Use first result + +# Build the output structure +output = {} + +# Initialize output structure based on provider grouping +for full_domain in full_domains: + output[full_domain] = {"records": {}} + + # Group by provider if using builtin + if builtin_providers: + for provider in builtin_providers: + output[full_domain]["records"][provider] = {} + if manual_nameservers: + output[full_domain]["records"]["manual_selection"] = {} + else: + # Just use the flat nameserver structure if no builtin used + pass + +# Populate results +for full_domain in full_domains: + if builtin_providers: + for ns in all_nameservers: + provider = nameserver_map.get(ns) + if (full_domain, ns) in results_map: + output[full_domain]["records"][provider][ns] = results_map[(full_domain, ns)] + else: + # Flat structure for manual nameservers only + for ns in all_nameservers: + if (full_domain, ns) in results_map: + output[full_domain]["records"][ns] = results_map[(full_domain, ns)] + +# Calculate and add runtime +elapsed_time = time.time() - start_time +output["runtime"] = f"{elapsed_time:.2f}s" print(json.dumps(output, indent=4))