import asyncio
import time
import logging
from typing import Any, Dict, List, Optional
import dns.asyncresolver
import dns.exception
import dns.rdatatype

logger = logging.getLogger(__name__)

# Standard public DNS resolvers for cross-network verification
RESOLVERS = [
    {"name": "Google DNS", "ip": "8.8.8.8"},
    {"name": "Cloudflare DNS", "ip": "1.1.1.1"},
    {"name": "Quad9 DNS", "ip": "9.9.9.9"},
    {"name": "OpenDNS", "ip": "208.67.222.222"},
]

COMMON_RECORD_TYPES = ["A", "AAAA", "CNAME", "MX", "NS", "TXT"]


async def check_single_resolver(
    domain_name: str,
    record_type: str,
    resolver_ip: str,
    timeout: float = 3.0,
) -> Dict[str, Any]:
    """
    Query a specific DNS resolver for a domain and record type.
    Measures latency and checks reachability.
    """
    clean_domain = domain_name.strip().rstrip(".")
    res = dns.asyncresolver.Resolver(configure=False)
    res.nameservers = [resolver_ip]
    res.timeout = timeout
    res.lifetime = timeout

    start_time = time.perf_counter()
    try:
        rdtype = dns.rdatatype.from_text(record_type.upper())
        answer = await res.resolve(clean_domain, rdtype)
        duration_ms = int((time.perf_counter() - start_time) * 1000)

        resolved_values = []
        for rdata in answer:
            resolved_values.append(rdata.to_text())

        return {
            "resolver_used": resolver_ip,
            "record_type": record_type.upper(),
            "resolved_values": resolved_values,
            "is_reachable": True,
            "response_time_ms": duration_ms,
            "error_message": None,
        }
    except dns.resolver.NXDOMAIN:
        duration_ms = int((time.perf_counter() - start_time) * 1000)
        return {
            "resolver_used": resolver_ip,
            "record_type": record_type.upper(),
            "resolved_values": [],
            "is_reachable": False,
            "response_time_ms": duration_ms,
            "error_message": "NXDOMAIN: Domain does not exist",
        }
    except dns.resolver.NoAnswer:
        duration_ms = int((time.perf_counter() - start_time) * 1000)
        return {
            "resolver_used": resolver_ip,
            "record_type": record_type.upper(),
            "resolved_values": [],
            "is_reachable": True,
            "response_time_ms": duration_ms,
            "error_message": f"No {record_type.upper()} records found for domain",
        }
    except (dns.resolver.Timeout, dns.exception.Timeout):
        duration_ms = int((time.perf_counter() - start_time) * 1000)
        return {
            "resolver_used": resolver_ip,
            "record_type": record_type.upper(),
            "resolved_values": [],
            "is_reachable": False,
            "response_time_ms": duration_ms,
            "error_message": "Query timed out",
        }
    except Exception as e:
        duration_ms = int((time.perf_counter() - start_time) * 1000)
        return {
            "resolver_used": resolver_ip,
            "record_type": record_type.upper(),
            "resolved_values": [],
            "is_reachable": False,
            "response_time_ms": duration_ms,
            "error_message": str(e),
        }


async def check_dns_across_resolvers(
    domain_name: str,
    record_type: str = "A",
    resolver_ips: Optional[List[str]] = None,
) -> List[Dict[str, Any]]:
    """
    Run concurrent queries for a domain across all standard public resolvers.
    """
    if not resolver_ips:
        resolver_ips = [r["ip"] for r in RESOLVERS]

    tasks = [
        check_single_resolver(domain_name, record_type, ip)
        for ip in resolver_ips
    ]
    results = await asyncio.gather(*tasks, return_exceptions=False)
    return results


async def check_all_record_types(
    domain_name: str,
    primary_resolver: str = "8.8.8.8",
) -> List[Dict[str, Any]]:
    """
    Query all common record types for a domain using a primary resolver.
    """
    tasks = [
        check_single_resolver(domain_name, rtype, primary_resolver)
        for rtype in COMMON_RECORD_TYPES
    ]
    results = await asyncio.gather(*tasks, return_exceptions=False)
    return results
