Cybersecurity-Projects/PROJECTS/beginner/dns-lookup/dnslookup/resolver.py

425 lines
12 KiB
Python

"""
ⒸAngelaMos | 2026
resolver.py
Async DNS resolution with record type support
Core DNS engine for the tool. Provides async functions for forward lookup,
reverse lookup, and trace using dnspython. Forward lookups fire all
record type queries concurrently via asyncio.gather. The trace function
walks the DNS hierarchy starting from root servers, following NS referrals
until it reaches an authoritative server with the final answer.
Key exports:
RecordType - StrEnum of supported record types (A, AAAA, MX, NS, TXT, CNAME, SOA, PTR)
ALL_RECORD_TYPES - Default list used by forward lookups (excludes PTR)
DNSRecord - Single record with type, value, TTL, and optional priority
DNSResult - Full lookup result with records, errors, timing, and nameserver
TraceHop - One server queried during a trace with zone, IP, and response summary
TraceResult - Full trace path including all hops and the final resolved answer
lookup() - Async forward lookup for one domain across multiple record types
reverse_lookup() - Async PTR record lookup for an IP address
trace_dns() - Synchronous DNS trace from root servers to authoritative servers
batch_lookup() - Async concurrent lookup for a list of domains
Connects to:
cli.py - all public functions and constants imported here
output.py - DNSResult, TraceResult, RecordType imported for display formatting
"""
from __future__ import annotations
import asyncio
import time
from dataclasses import dataclass, field
from enum import StrEnum
from typing import Any
import dns.asyncresolver
import dns.exception
import dns.message
import dns.name
import dns.query
import dns.rcode
import dns.rdatatype
import dns.resolver
import dns.reversename
class RecordType(StrEnum):
"""
Supported DNS record types
"""
A = "A"
AAAA = "AAAA"
MX = "MX"
NS = "NS"
TXT = "TXT"
CNAME = "CNAME"
SOA = "SOA"
PTR = "PTR"
ALL_RECORD_TYPES = [
RecordType.A,
RecordType.AAAA,
RecordType.MX,
RecordType.NS,
RecordType.TXT,
RecordType.CNAME,
RecordType.SOA,
]
@dataclass
class DNSRecord:
"""
Represents a single DNS record
"""
record_type: RecordType
value: str
ttl: int
priority: int | None = None
@dataclass
class DNSResult:
"""
Result of a DNS lookup
"""
domain: str
records: list[DNSRecord] = field(default_factory = list)
errors: list[str] = field(default_factory = list)
query_time_ms: float = 0.0
nameserver: str | None = None
@dataclass
class TraceHop:
"""
Represents a single hop in DNS resolution trace
"""
zone: str
server: str
server_ip: str
response: str
is_authoritative: bool = False
@dataclass
class TraceResult:
"""
Result of a DNS trace
"""
domain: str
hops: list[TraceHop] = field(default_factory = list)
final_answer: str | None = None
error: str | None = None
def create_resolver(
nameserver: str | None = None,
timeout: float = 5.0,
) -> dns.asyncresolver.Resolver:
"""
Create a configured async DNS resolver
"""
resolver = dns.asyncresolver.Resolver()
resolver.timeout = timeout
resolver.lifetime = timeout * 2
if nameserver:
resolver.nameservers = [nameserver]
return resolver
def extract_record_value(rdata: Any,
record_type: RecordType
) -> tuple[str,
int | None]:
"""
Extract value and priority from rdata based on record type
"""
priority = None
if record_type == RecordType.A or record_type == RecordType.AAAA:
value = rdata.address
elif record_type == RecordType.MX:
value = str(rdata.exchange).rstrip(".")
priority = rdata.preference
elif record_type in (RecordType.NS,
RecordType.CNAME,
RecordType.PTR):
value = str(rdata.target).rstrip(".")
elif record_type == RecordType.TXT:
value = rdata.to_text()
elif record_type == RecordType.SOA:
value = f"NS: {str(rdata.mname).rstrip('.')}, Serial: {rdata.serial}"
else:
value = rdata.to_text()
return value, priority
async def query_record_type(
domain: str,
record_type: RecordType,
resolver: dns.asyncresolver.Resolver,
) -> list[DNSRecord]:
"""
Query a single record type for a domain
"""
records = []
try:
answers = await resolver.resolve(domain, record_type.value)
for rdata in answers:
value, priority = extract_record_value(rdata, record_type)
records.append(
DNSRecord(
record_type = record_type,
value = value,
ttl = answers.rrset.ttl,
priority = priority,
)
)
except (dns.resolver.NXDOMAIN,
dns.resolver.NoAnswer,
dns.resolver.NoNameservers):
pass
except dns.exception.Timeout:
pass
return records
async def lookup(
domain: str,
record_types: list[RecordType] | None = None,
nameserver: str | None = None,
timeout: float = 5.0,
) -> DNSResult:
"""
Perform DNS lookup for specified record types
"""
if record_types is None:
record_types = ALL_RECORD_TYPES
resolver = create_resolver(nameserver, timeout)
result = DNSResult(domain = domain, nameserver = nameserver)
start_time = time.perf_counter()
tasks = [
query_record_type(domain,
rt,
resolver) for rt in record_types
]
query_results = await asyncio.gather(
*tasks,
return_exceptions = True
)
for i, query_result in enumerate(query_results):
if isinstance(query_result, Exception):
result.errors.append(f"{record_types[i]}: {query_result}")
else:
result.records.extend(query_result)
result.query_time_ms = (time.perf_counter() - start_time) * 1000
return result
async def reverse_lookup(
ip_address: str,
nameserver: str | None = None,
timeout: float = 5.0,
) -> DNSResult:
"""
Perform reverse DNS lookup for an IP address
"""
resolver = create_resolver(nameserver, timeout)
result = DNSResult(domain = ip_address, nameserver = nameserver)
start_time = time.perf_counter()
try:
answers = await resolver.resolve_address(ip_address)
for rdata in answers:
result.records.append(
DNSRecord(
record_type = RecordType.PTR,
value = str(rdata.target).rstrip("."),
ttl = answers.rrset.ttl,
)
)
except dns.resolver.NXDOMAIN:
result.errors.append("No PTR record found")
except dns.resolver.NoAnswer:
result.errors.append("No answer from nameserver")
except dns.resolver.NoNameservers:
result.errors.append("No nameservers available")
except dns.exception.Timeout:
result.errors.append("Query timed out")
except dns.exception.DNSException as e:
result.errors.append(str(e))
result.query_time_ms = (time.perf_counter() - start_time) * 1000
return result
def trace_dns(domain: str, record_type: str = "A") -> TraceResult:
"""
Trace DNS resolution path from root to authoritative servers
"""
result = TraceResult(domain = domain)
try:
name = dns.name.from_text(domain)
rdtype = dns.rdatatype.from_text(record_type)
root_servers = [
("a.root-servers.net",
"198.41.0.4"),
("b.root-servers.net",
"170.247.170.2"),
("c.root-servers.net",
"192.33.4.12"),
]
current_servers = root_servers
current_zone = "."
while True:
server_name, server_ip = current_servers[0]
try:
query = dns.message.make_query(name, rdtype)
response = dns.query.udp(
query,
server_ip,
timeout = 3.0
)
rcode = response.rcode()
if rcode != dns.rcode.NOERROR:
result.error = f"DNS error: {dns.rcode.to_text(rcode)}"
break
if response.answer:
for rrset in response.answer:
for rdata in rrset:
result.final_answer = str(rdata)
break
result.hops.append(
TraceHop(
zone = current_zone,
server = server_name,
server_ip = server_ip,
response =
f"{record_type}: {result.final_answer}",
is_authoritative = True,
)
)
break
if response.authority:
ns_records = []
for rrset in response.authority:
if rrset.rdtype == dns.rdatatype.NS:
for rdata in rrset:
ns_name = str(rdata.target
).rstrip(".")
ns_records.append(ns_name)
new_zone = str(rrset.name).rstrip(".")
if not new_zone:
new_zone = "."
if ns_records:
referral_msg = f"Referred to {new_zone or 'next'} servers"
result.hops.append(
TraceHop(
zone = current_zone,
server = server_name,
server_ip = server_ip,
response = referral_msg,
)
)
glue_ips = {}
if response.additional:
for rrset in response.additional:
if rrset.rdtype == dns.rdatatype.A:
for rdata in rrset:
glue_ips[str(
rrset.name
).rstrip(".")] = rdata.address
new_servers = []
for ns in ns_records:
if ns in glue_ips:
new_servers.append((ns, glue_ips[ns]))
else:
try:
answers = dns.resolver.resolve(
ns,
"A"
)
for rdata in answers:
new_servers.append(
(ns,
rdata.address)
)
break
except dns.exception.DNSException:
continue
if new_servers:
current_servers = new_servers
current_zone = new_zone
else:
result.error = "Could not resolve nameserver IPs"
break
else:
result.error = "No NS records in authority section"
break
else:
result.error = "No answer or authority in response"
break
except dns.exception.Timeout:
result.error = f"Timeout querying {server_name}"
break
except dns.exception.DNSException as e:
result.error = str(e)
break
except Exception as e:
result.error = str(e)
return result
async def batch_lookup(
domains: list[str],
record_types: list[RecordType] | None = None,
nameserver: str | None = None,
timeout: float = 5.0,
) -> list[DNSResult]:
"""
Perform DNS lookups for multiple domains concurrently
"""
tasks = [
lookup(domain,
record_types,
nameserver,
timeout) for domain in domains
]
return await asyncio.gather(*tasks)