import json
from datetime import timedelta
from typing import List, Optional

from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy.exc import IntegrityError, SQLAlchemyError
from sqlalchemy.orm import Session

from .. import models, schemas
from ..crypto import encrypt_token
from ..db import get_db
from ..dns_providers import (
    DNSProviderError,
    delete_previous_dkim_record_for_domain,
    publish_dkim_record_for_domain,
    publish_mail_dns_for_domain,
    publish_previous_dkim_record_for_domain,
)
from ..cache import CACHE_TTL_SECONDS, get_cache, safe_cache_delete, safe_cache_set
from ..security import require_admin
from ..settings import get_settings
from ..utils import (
    generate_dkim_key,
    remove_domain_mailboxes,
    remove_managed_dkim_key,
    resolve_domain_overlap_days,
    resolve_domain_rotation_interval_days,
    sync_dkim_signing_maps,
    timestamp_dkim_selector,
    validate_selector,
)

router = APIRouter(prefix="/domains", tags=["domains"])
settings = get_settings()


def cache_domain(cache, domain: models.Domain) -> None:
    safe_cache_set(
        cache,
        f"domain:{domain.name}",
        json.dumps(
            {
                "id": domain.id,
                "name": domain.name,
                "is_active": domain.is_active,
                "max_users": domain.max_users,
                "dkim_selector": domain.dkim_selector,
                "dkim_public_key": domain.dkim_public_key,
                "dkim_previous_selector": domain.dkim_previous_selector,
                "dkim_rotated_at": domain.dkim_rotated_at.isoformat() if domain.dkim_rotated_at else None,
                "dkim_previous_expires_at": domain.dkim_previous_expires_at.isoformat() if domain.dkim_previous_expires_at else None,
                "dkim_auto_rotate": domain.dkim_auto_rotate,
                "dkim_rotation_interval_days": domain.dkim_rotation_interval_days,
                "dkim_overlap_days": domain.dkim_overlap_days,
                "dmarc_policy": domain.dmarc_policy,
                "dns_provider": domain.dns_provider,
                "dns_account_id": domain.dns_account_id,
                "dns_zone_id": domain.dns_zone_id,
                "dns_sync_enabled": domain.dns_sync_enabled,
                "dns_last_sync_at": domain.dns_last_sync_at.isoformat() if domain.dns_last_sync_at else None,
                "dns_last_sync_status": domain.dns_last_sync_status,
            }
        ),
        ex=CACHE_TTL_SECONDS,
    )


def get_domain_or_404(db: Session, domain_id: int) -> models.Domain:
    domain = db.query(models.Domain).filter(models.Domain.id == domain_id).first()
    if not domain:
        raise HTTPException(status_code=404, detail="Domain not found")
    return domain


def _set_dns_sync_status(domain: models.Domain, status_value: str) -> None:
    domain.dns_last_sync_at = models.utc_now()
    domain.dns_last_sync_status = status_value[:512]


def _publish_domain_dkim(domain: models.Domain, db: Session, cache) -> dict[str, object]:
    try:
        change = publish_dkim_record_for_domain(domain)
    except DNSProviderError as exc:
        _set_dns_sync_status(domain, f"error: {exc}")
        db.commit()
        cache_domain(cache, domain)
        raise HTTPException(status_code=502, detail=str(exc)) from exc
    previous_payload = None
    try:
        previous = publish_previous_dkim_record_for_domain(domain)
        if previous:
            previous_payload = previous.as_dict()
    except DNSProviderError as exc:
        # Current key is published; previous-selector retention is best-effort.
        previous_payload = {"status": "error", "detail": str(exc)}
    status_parts = [f"{change.action} {change.name}"]
    if isinstance(previous_payload, dict) and previous_payload.get("action"):
        status_parts.append(f"{previous_payload.get('action')} {previous_payload.get('name')}")
    elif isinstance(previous_payload, dict) and previous_payload.get("status") == "error":
        status_parts.append(f"previous-error: {previous_payload.get('detail')}")
    _set_dns_sync_status(domain, "ok: " + "; ".join(status_parts))
    db.commit()
    db.refresh(domain)
    cache_domain(cache, domain)
    # Flat current-record fields for backward-compatible API clients.
    payload = change.as_dict()
    if previous_payload:
        payload["previous"] = previous_payload
    return payload


def _try_auto_publish_domain_dkim(domain: models.Domain, db: Session, cache) -> dict[str, object] | None:
    if not domain.dns_sync_enabled:
        return None
    try:
        return _publish_domain_dkim(domain, db, cache)
    except HTTPException as exc:
        return {"status": "error", "detail": exc.detail}


def _publish_domain_mail_dns(domain: models.Domain, db: Session, cache) -> list[dict[str, object]]:
    try:
        changes = publish_mail_dns_for_domain(domain, settings)
    except DNSProviderError as exc:
        _set_dns_sync_status(domain, f"error: {exc}")
        db.commit()
        cache_domain(cache, domain)
        raise HTTPException(status_code=502, detail=str(exc)) from exc
    summary = ", ".join(f"{item.get('action')} {item.get('type')} {item.get('name')}" for item in changes[:8])
    _set_dns_sync_status(domain, f"ok: published {len(changes)} records ({summary})")
    db.commit()
    db.refresh(domain)
    cache_domain(cache, domain)
    return changes


def _retire_previous_dkim_if_expired(domain: models.Domain, *, force: bool = False) -> dict[str, object]:
    """Drop previous DKIM material after the overlap window (or immediately if force)."""
    if not domain.dkim_previous_selector and not domain.dkim_previous_private_path:
        return {"retired": False, "reason": "no previous key"}
    now = models.utc_now()
    expires = domain.dkim_previous_expires_at
    if not force and expires and expires > now:
        return {"retired": False, "reason": "overlap active", "expires_at": expires.isoformat()}

    previous_path = domain.dkim_previous_private_path
    previous_selector = domain.dkim_previous_selector
    dns_deleted: list[dict[str, object]] = []
    if domain.dns_sync_enabled and previous_selector:
        try:
            dns_deleted = [item.as_dict() for item in delete_previous_dkim_record_for_domain(domain)]
        except DNSProviderError as exc:
            # Keep local retirement best-effort even if remote DNS cleanup fails.
            dns_deleted = [{"status": "error", "detail": str(exc)}]

    if previous_path and previous_path != domain.dkim_private_path:
        remove_managed_dkim_key(previous_path)

    domain.dkim_previous_selector = None
    domain.dkim_previous_public_key = None
    domain.dkim_previous_private_path = None
    domain.dkim_previous_expires_at = None
    return {
        "retired": True,
        "previous_selector": previous_selector,
        "dns": dns_deleted,
    }


def perform_dkim_rotation(
    domain: models.Domain,
    *,
    selector: str | None = None,
    auto_selector: bool = False,
    db: Session,
    cache,
) -> dict[str, object]:
    """Rotate DKIM with overlap: keep previous key/DNS until expiry."""
    # Retire any already-expired previous key first so we don't accumulate more than one.
    _retire_previous_dkim_if_expired(domain, force=False)

    old_selector = domain.dkim_selector or "default"
    old_public = domain.dkim_public_key
    old_private_path = domain.dkim_private_path

    if auto_selector or not selector or selector.strip().lower() in {"", "auto"}:
        requested_selector = timestamp_dkim_selector("s")
    else:
        try:
            requested_selector = validate_selector(selector)
        except ValueError as exc:
            raise HTTPException(status_code=400, detail=str(exc)) from exc

    # Avoid overwriting the active key file when reusing the same selector.
    if requested_selector == old_selector and old_private_path:
        requested_selector = timestamp_dkim_selector(f"{requested_selector[:8]}r")

    dkim = generate_dkim_key(domain.name, requested_selector)
    now = models.utc_now()
    overlap_days = resolve_domain_overlap_days(domain)

    if old_private_path and old_public:
        domain.dkim_previous_selector = old_selector
        domain.dkim_previous_public_key = old_public
        domain.dkim_previous_private_path = old_private_path
        domain.dkim_previous_expires_at = now + timedelta(days=overlap_days) if overlap_days > 0 else now
    else:
        domain.dkim_previous_selector = None
        domain.dkim_previous_public_key = None
        domain.dkim_previous_private_path = None
        domain.dkim_previous_expires_at = None

    domain.dkim_private_path = dkim["path"]
    domain.dkim_public_key = dkim.get("dns")
    domain.dkim_selector = dkim["selector"]
    domain.dkim_rotated_at = now

    try:
        db.commit()
    except SQLAlchemyError as exc:
        db.rollback()
        if dkim.get("path") != old_private_path:
            remove_managed_dkim_key(dkim.get("path"))
        raise HTTPException(status_code=500, detail="Unable to save the rotated DKIM key") from exc

    db.refresh(domain)
    sync_dkim_signing_maps(db.query(models.Domain).all())
    cache_domain(cache, domain)
    dns_sync = _try_auto_publish_domain_dkim(domain, db, cache)
    return {
        "dns_record": dkim.get("dns"),
        "selector": dkim["selector"],
        "path": dkim.get("path"),
        "previous_selector": domain.dkim_previous_selector,
        "previous_expires_at": domain.dkim_previous_expires_at.isoformat() if domain.dkim_previous_expires_at else None,
        "overlap_days": overlap_days,
        "dns_sync": dns_sync,
    }


def domain_needs_auto_rotation(domain: models.Domain) -> bool:
    if not domain.dkim_auto_rotate or not domain.is_active:
        return False
    if not domain.dkim_private_path:
        return True
    interval = resolve_domain_rotation_interval_days(domain)
    if interval <= 0:
        return False
    rotated_at = domain.dkim_rotated_at or domain.created_at
    if not rotated_at:
        return True
    return models.utc_now() >= rotated_at + timedelta(days=interval)


@router.get("/", response_model=List[schemas.DomainOut])
def list_domains(db: Session = Depends(get_db), _: str = Depends(require_admin)):
    return db.query(models.Domain).all()


@router.post("/", response_model=schemas.DomainOut)
def create_domain(payload: schemas.DomainCreate, db: Session = Depends(get_db), cache=Depends(get_cache), _: str = Depends(require_admin)):
    existing = db.query(models.Domain).filter(models.Domain.name == payload.name).first()
    if existing:
        raise HTTPException(status_code=400, detail="Domain already exists")
    domain = models.Domain(
        name=payload.name,
        is_active=payload.is_active,
        max_users=payload.max_users,
        dmarc_policy=payload.dmarc_policy,
        dkim_selector=payload.dkim_selector,
        dkim_auto_rotate=payload.dkim_auto_rotate,
        dkim_rotation_interval_days=payload.dkim_rotation_interval_days,
        dkim_overlap_days=payload.dkim_overlap_days,
    )
    if payload.generate_dkim:
        dkim = generate_dkim_key(payload.name, payload.dkim_selector)
        domain.dkim_private_path = dkim["path"]
        domain.dkim_public_key = dkim.get("dns")
        domain.dkim_rotated_at = models.utc_now()
    db.add(domain)
    try:
        db.commit()
    except IntegrityError as exc:
        db.rollback()
        if domain.dkim_private_path:
            remove_managed_dkim_key(domain.dkim_private_path)
        raise HTTPException(status_code=409, detail="Domain already exists") from exc
    db.refresh(domain)
    sync_dkim_signing_maps(db.query(models.Domain).all())
    cache_domain(cache, domain)
    return domain


@router.get("/{domain_id}", response_model=schemas.DomainOut)
def get_domain(domain_id: int, db: Session = Depends(get_db), _: str = Depends(require_admin)):
    return get_domain_or_404(db, domain_id)


@router.patch("/{domain_id}", response_model=schemas.DomainOut)
def update_domain(domain_id: int, payload: schemas.DomainUpdate, db: Session = Depends(get_db), cache=Depends(get_cache), _: str = Depends(require_admin)):
    domain = get_domain_or_404(db, domain_id)
    if payload.is_active is not None:
        domain.is_active = payload.is_active
    if payload.max_users is not None:
        domain.max_users = payload.max_users
    if payload.dmarc_policy is not None:
        domain.dmarc_policy = payload.dmarc_policy
    if payload.dkim_auto_rotate is not None:
        domain.dkim_auto_rotate = payload.dkim_auto_rotate
    if payload.dkim_rotation_interval_days is not None:
        domain.dkim_rotation_interval_days = payload.dkim_rotation_interval_days
    if payload.dkim_overlap_days is not None:
        domain.dkim_overlap_days = payload.dkim_overlap_days
    db.commit()
    db.refresh(domain)
    cache_domain(cache, domain)
    return domain


@router.post("/{domain_id}/dns/provider", response_model=schemas.DomainOut)
def configure_dns_provider(
    domain_id: int,
    payload: schemas.DomainDnsProviderUpdate,
    db: Session = Depends(get_db),
    cache=Depends(get_cache),
    _: str = Depends(require_admin),
):
    domain = get_domain_or_404(db, domain_id)
    provider = payload.dns_provider
    if provider == "none":
        domain.dns_provider = None
        domain.dns_account_id = None
        domain.dns_zone_id = None
        domain.dns_api_token = None
        domain.dns_sync_enabled = False
        domain.dns_last_sync_status = "disabled"
    else:
        domain.dns_provider = provider or "cloudflare"
        if payload.dns_account_id is not None:
            domain.dns_account_id = payload.dns_account_id
        if payload.dns_zone_id is not None:
            domain.dns_zone_id = payload.dns_zone_id
        if payload.clear_api_token:
            domain.dns_api_token = None
        elif payload.dns_api_token:
            domain.dns_api_token = encrypt_token(payload.dns_api_token)
        if payload.dns_sync_enabled is not None:
            domain.dns_sync_enabled = payload.dns_sync_enabled
        if domain.dns_sync_enabled and (not domain.dns_account_id or not domain.dns_zone_id or not domain.dns_api_token):
            raise HTTPException(status_code=400, detail="Cloudflare account ID, zone ID and API token are required before enabling DNS sync")
    db.commit()
    db.refresh(domain)
    cache_domain(cache, domain)
    return domain


@router.post("/{domain_id}/dkim", response_model=dict)
def rotate_dkim(
    domain_id: int,
    selector: str = "auto",
    db: Session = Depends(get_db),
    cache=Depends(get_cache),
    _: str = Depends(require_admin),
    auto_selector: bool = False,
):
    # Positional compatibility: rotate_dkim(id, selector, db, cache, _=...).
    domain = get_domain_or_404(db, domain_id)
    return perform_dkim_rotation(
        domain,
        selector=selector,
        auto_selector=auto_selector or str(selector or "").strip().lower() in {"auto", ""},
        db=db,
        cache=cache,
    )


@router.post("/{domain_id}/dkim/retire-previous", response_model=dict)
def retire_previous_dkim(
    domain_id: int,
    force: bool = Query(default=False),
    db: Session = Depends(get_db),
    cache=Depends(get_cache),
    _: str = Depends(require_admin),
):
    domain = get_domain_or_404(db, domain_id)
    result = _retire_previous_dkim_if_expired(domain, force=force)
    db.commit()
    db.refresh(domain)
    sync_dkim_signing_maps(db.query(models.Domain).all())
    cache_domain(cache, domain)
    return result


@router.post("/{domain_id}/dns/sync-dkim", response_model=dict)
def sync_dkim_dns(domain_id: int, db: Session = Depends(get_db), cache=Depends(get_cache), _: str = Depends(require_admin)):
    domain = get_domain_or_404(db, domain_id)
    result = _publish_domain_dkim(domain, db, cache)
    return {"status": "ok", "record": result}


@router.post("/{domain_id}/dns/sync-all", response_model=dict)
def sync_all_dns(domain_id: int, db: Session = Depends(get_db), cache=Depends(get_cache), _: str = Depends(require_admin)):
    domain = get_domain_or_404(db, domain_id)
    if not domain.dns_provider or domain.dns_provider == "none":
        raise HTTPException(status_code=400, detail="Configure a DNS provider before publishing records")
    if not domain.dns_account_id or not domain.dns_zone_id or not domain.dns_api_token:
        raise HTTPException(status_code=400, detail="Cloudflare account ID, zone ID and API token are required")
    changes = _publish_domain_mail_dns(domain, db, cache)
    return {"status": "ok", "records": changes, "count": len(changes)}


@router.post("/dkim/maintenance", response_model=dict)
def dkim_maintenance(db: Session = Depends(get_db), cache=Depends(get_cache), _: str = Depends(require_admin)):
    """Run scheduled DKIM maintenance: retire expired previous keys and auto-rotate due domains."""
    rotated: list[dict[str, object]] = []
    retired: list[dict[str, object]] = []
    errors: list[dict[str, object]] = []
    domains = db.query(models.Domain).all()
    for domain in domains:
        try:
            retire_result = _retire_previous_dkim_if_expired(domain, force=False)
            if retire_result.get("retired"):
                retired.append({"domain": domain.name, **retire_result})
            if domain_needs_auto_rotation(domain):
                result = perform_dkim_rotation(domain, auto_selector=True, db=db, cache=cache)
                rotated.append({"domain": domain.name, "selector": result.get("selector")})
        except HTTPException as exc:
            errors.append({"domain": domain.name, "detail": exc.detail})
        except Exception as exc:  # pragma: no cover - defensive
            errors.append({"domain": domain.name, "detail": str(exc)})
    db.commit()
    sync_dkim_signing_maps(db.query(models.Domain).all())
    return {
        "status": "ok",
        "rotated": rotated,
        "retired": retired,
        "errors": errors,
        "rotated_count": len(rotated),
        "retired_count": len(retired),
    }


@router.delete("/{domain_id}", response_model=dict)
def delete_domain(domain_id: int, db: Session = Depends(get_db), cache=Depends(get_cache), _: str = Depends(require_admin)):
    domain = get_domain_or_404(db, domain_id)
    domain_name = domain.name
    for account in domain.accounts:
        safe_cache_delete(cache, f"account:{account.local_part}@{domain_name}")
    for alias in domain.aliases:
        safe_cache_delete(cache, f"alias:{alias.source_local}@{domain_name}")
    for redirect in domain.redirects:
        safe_cache_delete(cache, f"redirect:{redirect.source_local}@{domain_name}")
    safe_cache_delete(cache, f"domain:{domain_name}")
    private_key_path = domain.dkim_private_path
    previous_key_path = domain.dkim_previous_private_path
    db.delete(domain)
    db.commit()
    remove_managed_dkim_key(private_key_path)
    remove_managed_dkim_key(previous_key_path)
    remove_domain_mailboxes(domain_name)
    sync_dkim_signing_maps(db.query(models.Domain).all())
    return {"deleted": True, "domain": domain_name}


@router.get("/{domain_id}/dns", response_model=dict)
def suggested_dns(domain_id: int, db: Session = Depends(get_db), _: str = Depends(require_admin)):
    domain = get_domain_or_404(db, domain_id)
    mx_host = settings.hostname
    mail_host_record = f"mail.{domain.name}"
    spf = f"v=spf1 mx a:{mx_host} -all"
    dmarc = f"v=DMARC1; p={domain.dmarc_policy}; adkim=s; aspf=s; pct=100; rua=mailto:postmaster@{domain.name}"
    dkim = domain.dkim_public_key or "Generate a DKIM key to populate this record"
    arc = "ARC enabled via Rspamd (no DNS record needed). Keep DKIM signing active for forwarded mail."
    payload = {
        "mx": f"{domain.name} IN MX 10 {mx_host}.",
        **({"mail": f"{mail_host_record} IN CNAME {mx_host}."} if mail_host_record != mx_host else {}),
        "spf": f"{domain.name} IN TXT \"{spf}\"",
        **(
            {"srs_spf": f"{settings.srs_domain} IN TXT \"{spf}\""}
            if settings.enable_srs
            and settings.srs_domain
            and settings.srs_domain != domain.name
            and domain.name == settings.primary_domain
            else {}
        ),
        "dmarc": f"_dmarc.{domain.name} IN TXT \"{dmarc}\"",
        "dkim": f"{domain.dkim_selector}._domainkey.{domain.name} IN TXT \"{dkim}\"",
        "jmap": f"_jmap._tcp.{domain.name} IN SRV 0 1 443 {mx_host}.",
        "imap": f"_imap._tcp.{domain.name} IN SRV 0 1 143 {mx_host}.",
        "imaps": f"_imaps._tcp.{domain.name} IN SRV 0 1 993 {mx_host}.",
        "pop3": f"_pop3._tcp.{domain.name} IN SRV 0 1 110 {mx_host}.",
        "pop3s": f"_pop3s._tcp.{domain.name} IN SRV 0 1 995 {mx_host}.",
        "submission": f"_submission._tcp.{domain.name} IN SRV 0 1 587 {mx_host}.",
        "submissions": f"_submissions._tcp.{domain.name} IN SRV 0 1 465 {mx_host}.",
        "autoconfig": f"autoconfig.{domain.name} IN CNAME {mx_host}.",
        "autodiscover": f"autodiscover.{domain.name} IN CNAME {mx_host}.",
        "arc": arc,
        "ptr": f"Set the reverse DNS of the public mail IP to {mx_host}.",
    }
    if domain.dkim_previous_selector and domain.dkim_previous_public_key:
        payload["dkim_previous"] = (
            f"{domain.dkim_previous_selector}._domainkey.{domain.name} IN TXT \"{domain.dkim_previous_public_key}\""
        )
        if domain.dkim_previous_expires_at:
            payload["dkim_previous_expires_at"] = domain.dkim_previous_expires_at.isoformat()
    if settings.enable_mta_sts and domain.name == settings.primary_domain:
        if settings.mta_sts_host != mx_host:
            payload["mta_sts_host"] = f"{settings.mta_sts_host} IN CNAME {mx_host}."
        payload["mta_sts"] = f"_mta-sts.{domain.name} IN TXT \"v=STSv1; id={settings.mta_sts_id}\""
        payload["mta_sts_policy"] = f"https://{settings.mta_sts_host}/.well-known/mta-sts.txt"
    if settings.enable_tls_rpt and domain.name == settings.primary_domain:
        payload["tls_rpt"] = f"_smtp._tls.{domain.name} IN TXT \"v=TLSRPTv1; rua=mailto:{settings.tls_rpt_mailbox}\""
    return payload
