from datetime import datetime, timezone
import json
import logging
from typing import Any, Dict, List, Optional
from fastapi import APIRouter, Depends, HTTPException, Query
from pydantic import BaseModel
from sqlalchemy import desc, select
from sqlalchemy.ext.asyncio import AsyncSession

from database import get_db
from models import DnsCheckResult, DnsRecord, Domain, RegistrarAccount
from crypto import decrypt_string
from integrations import get_registrar
from dns_checker import check_dns_across_resolvers, check_all_record_types

logger = logging.getLogger(__name__)
router = APIRouter(prefix="/api/dns", tags=["dns"])


class DnsRecordItem(BaseModel):
    id: Optional[int] = None
    record_type: str
    name: str
    value: str
    ttl: Optional[int] = None
    priority: Optional[int] = None


class DnsSnapshotGroup(BaseModel):
    snapshotted_at: datetime
    records: List[DnsRecordItem]


class DnsCheckAnyRequest(BaseModel):
    domain: str
    record_type: str = "A"


@router.get("/domains/{domain_id}/history")
async def get_dns_history(domain_id: int, db: AsyncSession = Depends(get_db)):
    """
    Get DNS record history grouped by snapshot timestamp, with change diffs.
    """
    domain = await db.get(Domain, domain_id)
    if not domain:
        raise HTTPException(status_code=404, detail="Domain not found")

    stmt = (
        select(DnsRecord)
        .where(DnsRecord.domain_id == domain_id)
        .order_by(desc(DnsRecord.snapshotted_at), DnsRecord.record_type, DnsRecord.name)
    )
    records = (await db.execute(stmt)).scalars().all()

    # Group by snapshot timestamp
    groups_map: Dict[datetime, List[DnsRecordItem]] = {}
    for r in records:
        ts = r.snapshotted_at
        if ts not in groups_map:
            groups_map[ts] = []
        groups_map[ts].append(
            DnsRecordItem(
                id=r.id,
                record_type=r.record_type,
                name=r.name,
                value=r.value,
                ttl=r.ttl,
                priority=r.priority,
            )
        )

    # Build sorted snapshot list
    snapshots = []
    sorted_timestamps = sorted(groups_map.keys(), reverse=True)

    for i, ts in enumerate(sorted_timestamps):
        current_records = groups_map[ts]
        diff = {"added": [], "removed": [], "changed": []}

        # Compare with previous snapshot in time (i.e. i+1 in reverse sorted list)
        if i + 1 < len(sorted_timestamps):
            prev_ts = sorted_timestamps[i + 1]
            prev_records = groups_map[prev_ts]

            prev_map = {f"{r.record_type}:{r.name}": r.value for r in prev_records}
            curr_map = {f"{r.record_type}:{r.name}": r.value for r in current_records}

            for key, val in curr_map.items():
                if key not in prev_map:
                    diff["added"].append(key)
                elif prev_map[key] != val:
                    diff["changed"].append({"key": key, "old": prev_map[key], "new": val})

            for key in prev_map:
                if key not in curr_map:
                    diff["removed"].append(key)

        snapshots.append({
            "snapshotted_at": ts,
            "records": current_records,
            "diff": diff,
        })

    return {"domain_id": domain_id, "domain_name": domain.domain_name, "snapshots": snapshots}


@router.post("/domains/{domain_id}/snapshot")
async def take_dns_snapshot(domain_id: int, db: AsyncSession = Depends(get_db)):
    """
    Fetch current DNS records from registrar API and save a new snapshot.
    """
    query = (
        select(Domain, RegistrarAccount)
        .join(RegistrarAccount, Domain.account_id == RegistrarAccount.id)
        .where(Domain.id == domain_id)
    )
    result = (await db.execute(query)).first()
    if not result:
        raise HTTPException(status_code=404, detail="Domain not found")

    domain, account = result

    try:
        creds = json.loads(decrypt_string(account.credentials_encrypted))
        client = get_registrar(account.registrar, creds)
        records_data = await client.get_dns_records(domain.domain_name)
    except Exception as e:
        logger.error(f"Failed to fetch DNS records from {account.registrar}: {e}")
        raise HTTPException(
            status_code=500, detail=f"Failed to fetch DNS records from registrar: {str(e)}"
        )

    now = datetime.now(timezone.utc)
    saved_records = []
    for r in records_data:
        record = DnsRecord(
            domain_id=domain_id,
            record_type=r.get("record_type", "A").upper(),
            name=r.get("name", "@"),
            value=r.get("value", ""),
            ttl=r.get("ttl"),
            priority=r.get("priority"),
            snapshotted_at=now,
        )
        db.add(record)
        saved_records.append(record)

    await db.commit()
    return {
        "message": f"Successfully took snapshot with {len(saved_records)} records",
        "snapshotted_at": now,
        "record_count": len(saved_records),
    }


@router.post("/domains/{domain_id}/check")
async def run_live_dns_check(
    domain_id: int,
    record_type: str = Query("A"),
    db: AsyncSession = Depends(get_db),
):
    """
    Perform live DNS resolution check across multiple public resolvers and store results.
    """
    domain = await db.get(Domain, domain_id)
    if not domain:
        raise HTTPException(status_code=404, detail="Domain not found")

    results = await check_dns_across_resolvers(domain.domain_name, record_type)

    now = datetime.now(timezone.utc)
    for res in results:
        check_obj = DnsCheckResult(
            domain_id=domain_id,
            checked_at=now,
            resolver_used=res["resolver_used"],
            record_type=res["record_type"],
            resolved_values=json.dumps(res["resolved_values"]),
            is_reachable=res["is_reachable"],
            response_time_ms=res["response_time_ms"],
            error_message=res["error_message"],
        )
        db.add(check_obj)

    await db.commit()
    return {
        "domain_id": domain_id,
        "domain_name": domain.domain_name,
        "record_type": record_type,
        "checked_at": now,
        "results": results,
    }


@router.get("/domains/{domain_id}/checks")
async def get_check_history(
    domain_id: int,
    limit: int = Query(20, ge=1, le=100),
    db: AsyncSession = Depends(get_db),
):
    """Get history of live DNS checks for a domain."""
    domain = await db.get(Domain, domain_id)
    if not domain:
        raise HTTPException(status_code=404, detail="Domain not found")

    stmt = (
        select(DnsCheckResult)
        .where(DnsCheckResult.domain_id == domain_id)
        .order_by(desc(DnsCheckResult.checked_at))
        .limit(limit)
    )
    results = (await db.execute(stmt)).scalars().all()

    formatted = []
    for r in results:
        resolved = []
        if r.resolved_values:
            try:
                resolved = json.loads(r.resolved_values)
            except Exception:
                resolved = [r.resolved_values]

        formatted.append({
            "id": r.id,
            "checked_at": r.checked_at,
            "resolver_used": r.resolver_used,
            "record_type": r.record_type,
            "resolved_values": resolved,
            "is_reachable": r.is_reachable,
            "response_time_ms": r.response_time_ms,
            "error_message": r.error_message,
        })

    return {"domain_id": domain_id, "checks": formatted}


@router.post("/check/any")
async def check_any_domain(payload: DnsCheckAnyRequest):
    """Live DNS check for any domain across multiple resolvers without saving."""
    results = await check_dns_across_resolvers(payload.domain, payload.record_type)
    return {
        "domain": payload.domain,
        "record_type": payload.record_type,
        "results": results,
    }


@router.get("/historical")
async def get_historical_dns_records(
    domain: str = Query(..., description="Domain name e.g. example.com"),
    record_type: str = Query("A", description="A, AAAA, MX, NS, TXT, or ALL"),
    securitytrails_key: Optional[str] = Query(None, description="Optional SecurityTrails API Key"),
):
    """
    SecurityTrails-style Internet-wide Historical DNS Records:
    Retrieves past IP addresses, nameservers, mail servers, first_seen, last_seen,
    and ASN/organization metadata across historical data sources.
    """
    from historical_dns import get_aggregated_historical_dns
    data = await get_aggregated_historical_dns(
        domain=domain,
        record_type=record_type,
        securitytrails_api_key=securitytrails_key,
    )
    return data


@router.get("/subdomains")
async def get_domain_subdomains(
    domain: str = Query(..., description="Domain name e.g. example.com"),
):
    """
    Historical Subdomain Discovery:
    Discovers all subdomains and historical certificates for a domain.
    """
    from historical_dns import fetch_crtsh_subdomains
    subdomains = await fetch_crtsh_subdomains(domain)
    return {
        "domain": domain,
        "total": len(subdomains),
        "subdomains": subdomains,
    }

