diff --git a/dns_compare.py b/dns_compare.py index e882964..71c4548 100755 --- a/dns_compare.py +++ b/dns_compare.py @@ -7,39 +7,73 @@ import sys import time from datetime import datetime from concurrent.futures import ThreadPoolExecutor -from dns import message -from dns import query -from dns import rdatatype -from dns import rcode -from dns import name -from dns import flags -from dns import exception +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]) +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'], + "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.') +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 @@ -49,15 +83,17 @@ builtin_providers = [] # list of provider keys used # Resolve builtin nameservers to actual IPs if args.ns_builtin: # Handle 'all' special case - if 'all' in [p.lower() for p in args.ns_builtin]: + if "all" in [p.lower() for p in args.ns_builtin]: builtin_providers = list(builtin_nameservers.keys()) else: builtin_providers = [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") + print( + f"Available providers: {', '.join(builtin_nameservers.keys())}, all" + ) exit(1) for ip in builtin_nameservers[provider]: nameserver_map[ip] = provider @@ -80,7 +116,7 @@ if not all_nameservers: # Handle listing record types if args.list_record_types: print("Supported DNS record types:") - print((", ").join(supported_types)) + print(", ".join(supported_types)) exit(0) # Convert record type string to dns.rdatatype constant @@ -88,8 +124,8 @@ try: 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)) + print("\nSupported record types:") + print(", ".join(supported_types)) exit(1) # Function to perform a single DNS query @@ -98,113 +134,170 @@ def perform_query(full_domain, ns, record_type, timeout): 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) - + 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(f"[DEBUG] UDP returned ETPA, retrying with TCP...", file=sys.stderr) + 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) + 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(f"[DEBUG] TCP also returned ETPA, retrying with TLS...", file=sys.stderr) + 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) + 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) + 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) + 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(f"[DEBUG] UDP Response truncated (TC flag set), retrying with TCP...", file=sys.stderr) + 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) + 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) + 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) + 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) + 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) + 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(f"[DEBUG] TCP returned ETPA, retrying with TLS...", file=sys.stderr) + 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) + 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) + 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) - + 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"\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] 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 - response_rcode = r.rcode() if response_rcode != 0: # 0 = NOERROR rcode_name = rcode.to_text(response_rcode) return (full_domain, ns, f"RCODE: {rcode_name}", protocol_used) - - results = [] - + # Check answer section first answers = r.answer if answers: @@ -212,7 +305,7 @@ def perform_query(full_domain, ns, record_type, timeout): for item in rrset: results.append(str(item)) return (full_domain, ns, results, protocol_used) - + # If no answer, check additional section (some servers put records there) additional = r.additional if additional: @@ -221,7 +314,7 @@ def perform_query(full_domain, ns, record_type, timeout): results.append(str(item)) if results: return (full_domain, ns, results, protocol_used) - + # If still no results, check authority section authority = r.authority if authority: @@ -230,7 +323,7 @@ def perform_query(full_domain, ns, record_type, timeout): 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): @@ -247,13 +340,13 @@ start_time_iso = datetime.fromtimestamp(start_time).isoformat() 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_domain = f"{rec}.{args.domain}" full_domains.append(full_domain) else: # Query each record as a full domain @@ -274,32 +367,32 @@ 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] - + 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 = {} +output = {"records": {}, "metadata": {"timing": {}, "statistics": {}}} -# Initialize output structure based on provider grouping +# Initialize records structure based on provider grouping for full_domain in full_domains: - output[full_domain] = {"records": {}} - + output["records"][full_domain] = {"results": {}} + # Group by provider if using builtin if builtin_providers: for provider in builtin_providers: - output[full_domain]["records"][provider] = {} + output["records"][full_domain]["results"][provider] = {} if manual_nameservers: - output[full_domain]["records"]["manual_selection"] = {} - else: - # Just use the flat nameserver structure if no builtin used - pass + output["records"][full_domain]["results"]["manual_selection"] = {} -# Populate results +# 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: @@ -317,7 +410,27 @@ for full_domain in full_domains: entry["result"] = results[0] if len(results) == 1 else results if protocol_used: entry["protocol"] = protocol_used - output[full_domain]["records"][provider][ns] = entry + 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: @@ -334,17 +447,35 @@ for full_domain in full_domains: entry["result"] = results[0] if len(results) == 1 else results if protocol_used: entry["protocol"] = protocol_used - output[full_domain]["records"][ns] = entry + output["records"][full_domain]["results"][ns] = entry -# Calculate and add timing information +# 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) +<<<<<<< HEAD output["metadata"] = { "start_time": start_time_iso, "end_time": end_time_iso, "runtime": f"{elapsed_time:.2f}s" } +======= +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 +>>>>>>> b2565dd (Updated to have better json format and statistics.) print(json.dumps(output, indent=4))