Files
dns_compare/dns_compare.py
2026-09-10 15:06:42 -04:00

460 lines
18 KiB
Python
Executable File

#!/usr/bin/env python
import argparse
import json
import os
import sys
import time
from datetime import datetime
from concurrent.futures import ThreadPoolExecutor
import dns.message as message
import dns.query as query
import dns.rdatatype as rdatatype
import dns.rcode as rcode
import dns.name as name
import dns.flags as flags
import dns.exception as exception
# Get list of supported record types
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())}, all",
)
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(
"--threads",
type=int,
default=None,
help="Number of threads to use for queries. Default: CPU count - 1",
)
parser.add_argument(
"--timeout",
type=float,
default=15.0,
help="DNS query timeout in seconds. Default: 5.0",
)
parser.add_argument(
"--verbose", action="store_true", help="Enable verbose debugging output."
)
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:
# Handle 'all' special case
builtin_providers = (
list(builtin_nameservers.keys())
if "all" in [p.lower() for p in args.ns_builtin]
else [p.lower() for p in args.ns_builtin]
)
for provider in builtin_providers:
if provider not in builtin_nameservers:
print(f"Error: Unknown DNS provider '{provider}'")
print(
f"Available providers: {', '.join(builtin_nameservers.keys())}, all"
)
exit(1)
for ip in builtin_nameservers[provider]:
nameserver_map[ip] = provider
# 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:")
print(", ".join(supported_types))
exit(0)
# Convert record type string to dns.rdatatype constant
try:
record_type = rdatatype.from_text(args.record_type)
except rdatatype.UnknownRdatatype:
print(f"Error: Unknown record type '{args.record_type}'")
print("\nSupported record types:")
print(", ".join(supported_types))
exit(1)
# Function to perform a single DNS query
def perform_query(full_domain, ns, record_type, timeout):
"""Perform a DNS query and return results with protocol used"""
try:
qname = name.from_text(full_domain)
q = message.make_query(qname, record_type)
# Add EDNS0 extension for better compatibility (like dig does)
q.use_edns(edns=0, ednsflags=0, payload=4096)
protocol_used = None
# Try UDP first
try:
r = query.udp(q, ns, timeout=timeout)
protocol_used = "UDP"
if args.verbose:
print(
f"[DEBUG] UDP Response: Answer={len(r.answer)} RRsets, TC flag: {bool(r.flags & flags.TC)}",
file=sys.stderr,
)
# Check if UDP returned ETPA error - if so, try TCP
udp_rcode = r.rcode()
udp_rcode_name = rcode.to_text(udp_rcode)
if udp_rcode_name == "ETPA":
if args.verbose:
print(
"[DEBUG] UDP returned ETPA, retrying with TCP...",
file=sys.stderr,
)
try:
r = query.tcp(q, ns, timeout=timeout)
protocol_used = "TCP"
if args.verbose:
print(
f"[DEBUG] TCP Response: Answer={len(r.answer)} RRsets, TC flag: {bool(r.flags & flags.TC)}",
file=sys.stderr,
)
# If TCP also returned ETPA, try TLS
tcp_rcode = r.rcode()
tcp_rcode_name = rcode.to_text(tcp_rcode)
if tcp_rcode_name == "ETPA":
if args.verbose:
print(
"[DEBUG] TCP also returned ETPA, retrying with TLS...",
file=sys.stderr,
)
r = query.tls(q, ns, timeout=timeout)
protocol_used = "TLS"
if args.verbose:
print(
f"[DEBUG] TLS Response: Answer={len(r.answer)} RRsets, TC flag: {bool(r.flags & flags.TC)}",
file=sys.stderr,
)
except Exception as tcp_error:
# Fall back to TLS if TCP fails
if args.verbose:
print(
f"[DEBUG] TCP failed ({tcp_error}), trying TLS...",
file=sys.stderr,
)
r = query.tls(q, ns, timeout=timeout)
protocol_used = "TLS"
if args.verbose:
print(
f"[DEBUG] TLS Response: Answer={len(r.answer)} RRsets, TC flag: {bool(r.flags & flags.TC)}",
file=sys.stderr,
)
# If UDP returns truncated, retry with TCP
elif r.flags & flags.TC:
if args.verbose:
print(
"[DEBUG] UDP Response truncated (TC flag set), retrying with TCP...",
file=sys.stderr,
)
try:
r = query.tcp(q, ns, timeout=timeout)
protocol_used = "TCP"
if args.verbose:
print(
f"[DEBUG] TCP Response: Answer={len(r.answer)} RRsets, TC flag: {bool(r.flags & flags.TC)}",
file=sys.stderr,
)
except Exception as tcp_error:
# Fall back to TLS if TCP fails
if args.verbose:
print(
f"[DEBUG] TCP failed ({tcp_error}), trying TLS...",
file=sys.stderr,
)
r = query.tls(q, ns, timeout=timeout)
protocol_used = "TLS"
if args.verbose:
print(
f"[DEBUG] TLS Response: Answer={len(r.answer)} RRsets, TC flag: {bool(r.flags & flags.TC)}",
file=sys.stderr,
)
except Exception as udp_error:
# Fall back to TCP if UDP fails
if args.verbose:
print(
f"[DEBUG] UDP failed ({udp_error}), trying TCP...",
file=sys.stderr,
)
try:
r = query.tcp(q, ns, timeout=timeout)
protocol_used = "TCP"
if args.verbose:
print(
f"[DEBUG] TCP Response: Answer={len(r.answer)} RRsets, TC flag: {bool(r.flags & flags.TC)}",
file=sys.stderr,
)
# If TCP returned ETPA, try TLS
tcp_rcode = r.rcode()
tcp_rcode_name = rcode.to_text(tcp_rcode)
if tcp_rcode_name == "ETPA":
if args.verbose:
print(
"[DEBUG] TCP returned ETPA, retrying with TLS...",
file=sys.stderr,
)
r = query.tls(q, ns, timeout=timeout)
protocol_used = "TLS"
if args.verbose:
print(
f"[DEBUG] TLS Response: Answer={len(r.answer)} RRsets, TC flag: {bool(r.flags & flags.TC)}",
file=sys.stderr,
)
except Exception as tcp_error:
# Fall back to TLS if TCP fails
if args.verbose:
print(
f"[DEBUG] TCP failed ({tcp_error}), trying TLS...",
file=sys.stderr,
)
r = query.tls(q, ns, timeout=timeout)
protocol_used = "TLS"
if args.verbose:
print(
f"[DEBUG] TLS Response: Answer={len(r.answer)} RRsets, TC flag: {bool(r.flags & flags.TC)}",
file=sys.stderr,
)
# Verbose debugging
if args.verbose:
print(
f"\n[DEBUG] Query: {full_domain} (@{ns}) Type: {rdatatype.to_text(record_type)}",
file=sys.stderr,
)
print(f"[DEBUG] Response code: {rcode.to_text(r.rcode())}", file=sys.stderr)
print(f"[DEBUG] Answer section: {len(r.answer)} RRsets", file=sys.stderr)
print(
f"[DEBUG] Authority section: {len(r.authority)} RRsets",
file=sys.stderr,
)
print(
f"[DEBUG] Additional section: {len(r.additional)} RRsets",
file=sys.stderr,
)
print(f"[DEBUG] Full message: {r}", file=sys.stderr)
results = []
response_rcode = r.rcode()
# Check response code
if response_rcode != 0: # 0 = NOERROR
rcode_name = rcode.to_text(response_rcode)
return (full_domain, ns, f"RCODE: {rcode_name}", protocol_used)
# Check answer section first, then additional, then authority
for section in (r.answer, r.additional, r.authority):
if section:
for rrset in section:
for item in rrset:
results.append(str(item))
if results:
return (full_domain, ns, results, protocol_used)
# Truly no records found
return (full_domain, ns, "NO RECORDS", protocol_used)
except (exception.Timeout, TimeoutError):
return (full_domain, ns, "TIMED OUT", None)
except Exception as e:
error_msg = str(e) if str(e) else type(e).__name__
return (full_domain, ns, f"ERROR: {error_msg}", None)
# Start timer
start_time = time.time()
start_time_iso = datetime.fromtimestamp(start_time).isoformat()
# 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}"
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: use provided value or default to CPU count - 1, minimum 1
thread_count = args.threads if args.threads else max(1, os.cpu_count() - 1)
# Execute queries in parallel
results_map = {} # (full_domain, ns) -> (resolved_ip, protocol_used)
with ThreadPoolExecutor(max_workers=thread_count) as executor:
futures = [
executor.submit(perform_query, full_domain, ns, record_type, args.timeout)
for full_domain, ns, _ in query_tasks
]
for future in futures:
full_domain, ns, resolved_ips, protocol_used = future.result()
if resolved_ips is not None:
results_map[(full_domain, ns)] = (resolved_ips, protocol_used)
# Build the output structure
output = {"records": {}, "metadata": {"timing": {}, "statistics": {}}}
# Initialize records structure based on provider grouping
for full_domain in full_domains:
output["records"][full_domain] = {"results": {}}
# Group by provider if using builtin
if builtin_providers:
for provider in builtin_providers:
output["records"][full_domain]["results"][provider] = {}
if manual_nameservers:
output["records"][full_domain]["results"]["manual_selection"] = {}
# Populate results and collect statistics
provider_stats = {} # provider -> {result: count, protocol: count}
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:
results, protocol_used = results_map[(full_domain, ns)]
# Create entry with result and protocol
entry = {}
# Handle error/timeout/rcode cases (strings) vs list of results
match results:
case str():
entry["result"] = results
case list():
results = sorted(results)
# Store as single value if only one result, otherwise as array
entry["result"] = results[0] if len(results) == 1 else results
if protocol_used:
entry["protocol"] = protocol_used
output["records"][full_domain]["results"][provider][ns] = entry
# Collect statistics
if provider not in provider_stats:
provider_stats[provider] = {}
result_val = entry.get("result")
if result_val:
# Handle both single results and lists of results
if isinstance(result_val, list):
for item in result_val:
provider_stats[provider][item] = (
provider_stats[provider].get(item, 0) + 1
)
else:
provider_stats[provider][result_val] = (
provider_stats[provider].get(result_val, 0) + 1
)
if protocol_used:
provider_stats[provider][protocol_used] = (
provider_stats[provider].get(protocol_used, 0) + 1
)
else:
# Flat structure for manual nameservers only
for ns in all_nameservers:
if (full_domain, ns) in results_map:
results, protocol_used = results_map[(full_domain, ns)]
# Create entry with result and protocol
entry = {}
# Handle error/timeout/rcode cases (strings) vs list of results
match results:
case str():
entry["result"] = results
case list():
results = sorted(results)
# Store as single value if only one result, otherwise as array
entry["result"] = results[0] if len(results) == 1 else results
if protocol_used:
entry["protocol"] = protocol_used
output["records"][full_domain]["results"][ns] = entry
# Calculate timing information
end_time = time.time()
end_time_iso = datetime.fromtimestamp(end_time).isoformat()
elapsed_time = end_time - start_time
elapsed_ms = int(elapsed_time * 1000)
output["metadata"]["timing"]["start_time"] = start_time_iso
output["metadata"]["timing"]["end_time"] = end_time_iso
output["metadata"]["timing"]["runtime_ms"] = elapsed_ms
# Build statistics section
for provider, stats in provider_stats.items():
output["metadata"]["statistics"][provider] = stats
# Calculate totals
total_stats = {}
for provider_stats_dict in provider_stats.values():
for key, count in provider_stats_dict.items():
total_stats[key] = total_stats.get(key, 0) + count
output["metadata"]["statistics"]["total"] = total_stats
print(json.dumps(output, indent=4))