diff --git a/dns_compare.py b/dns_compare.py index 3fea78f..030ae61 100755 --- a/dns_compare.py +++ b/dns_compare.py @@ -10,6 +10,7 @@ from dns import message from dns import query from dns import rdatatype from dns import name +from dns import flags # Get list of supported record types supported_types = sorted([rdatatype.to_text(rdtype) for rdtype in rdatatype.RdataType]) @@ -91,8 +92,14 @@ def perform_query(full_domain, ns, record_type): try: qname = name.from_text(full_domain) q = message.make_query(qname, record_type) + + # Try UDP first r = query.udp(q, ns) + # If truncated, retry with TCP to get all results + if r.flags & flags.TC: + r = query.tcp(q, ns) + results = [] answers = r.answer if answers: @@ -145,7 +152,7 @@ with ThreadPoolExecutor(max_workers=thread_count) as executor: 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 + results_map[(full_domain, ns)] = resolved_ips # Build the output structure output = {} @@ -170,12 +177,16 @@ for full_domain in full_domains: 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)] + results = sorted(results_map[(full_domain, ns)]) + # Store as single value if only one result, otherwise as array + output[full_domain]["records"][provider][ns] = results[0] if len(results) == 1 else results 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)] + results = sorted(results_map[(full_domain, ns)]) + # Store as single value if only one result, otherwise as array + output[full_domain]["records"][ns] = results[0] if len(results) == 1 else results # Calculate and add timing information end_time = time.time()