from __future__ import annotations

import os
from datetime import datetime, timedelta
from typing import Any, Dict, List, Optional

from sqlalchemy import select
from sqlalchemy.orm import Session

from db.channel_exports import CHANNEL_CODES, resolve_pipeline_channel_code
from db.models import ChannelPipelineReadiness


READINESS_STATUSES = frozenset({"missing", "queued", "running", "ready", "failed", "stale"})
DEFAULT_REFRESH_HOURS = int(os.getenv("CHANNEL_READINESS_REFRESH_HOURS", "4"))
STALE_RUNNING_MINUTES = int(os.getenv("CHANNEL_READINESS_STALE_RUNNING_MINUTES", "30"))


def _utcnow_naive() -> datetime:
    return datetime.utcnow()


def _as_utc_naive(value: Optional[datetime]) -> Optional[datetime]:
    if value is None:
        return None
    if value.tzinfo is not None:
        return value.replace(tzinfo=None)
    return value


def readiness_scope_key(
    *,
    only_assigned: bool = True,
    connection_id: Optional[int] = None,
    magento_connection_id: Optional[int] = None,
    shopify_connection_id: Optional[int] = None,
) -> str:
    return "|".join(
        [
            f"assigned={1 if only_assigned else 0}",
            f"conn={connection_id or ''}",
            f"magento={magento_connection_id or ''}",
            f"shopify={shopify_connection_id or ''}",
        ]
    )


def default_pipeline_scopes() -> List[Dict[str, Any]]:
    """Default dashboard scopes: assigned-only, default connections."""
    return [
        {
            "channel_code": channel,
            "only_assigned": True,
            "connection_id": None,
            "magento_connection_id": None,
            "shopify_connection_id": None,
        }
        for channel in sorted(CHANNEL_CODES)
    ]


def get_readiness_cache_row(
    session: Session,
    channel_code: str,
    *,
    only_assigned: bool = True,
    connection_id: Optional[int] = None,
    magento_connection_id: Optional[int] = None,
    shopify_connection_id: Optional[int] = None,
) -> Optional[ChannelPipelineReadiness]:
    channel = resolve_pipeline_channel_code(channel_code)
    scope_key = readiness_scope_key(
        only_assigned=only_assigned,
        connection_id=connection_id,
        magento_connection_id=magento_connection_id,
        shopify_connection_id=shopify_connection_id,
    )
    return session.scalar(
        select(ChannelPipelineReadiness)
        .where(ChannelPipelineReadiness.channel_code == channel)
        .where(ChannelPipelineReadiness.scope_key == scope_key)
    )


def queue_readiness_refresh(
    session: Session,
    channel_code: str,
    *,
    only_assigned: bool = True,
    connection_id: Optional[int] = None,
    magento_connection_id: Optional[int] = None,
    shopify_connection_id: Optional[int] = None,
    force: bool = False,
) -> ChannelPipelineReadiness:
    """Mark a readiness scope for background recomputation."""
    channel = resolve_pipeline_channel_code(channel_code)
    scope_key = readiness_scope_key(
        only_assigned=only_assigned,
        connection_id=connection_id,
        magento_connection_id=magento_connection_id,
        shopify_connection_id=shopify_connection_id,
    )
    row = get_readiness_cache_row(
        session,
        channel,
        only_assigned=only_assigned,
        connection_id=connection_id,
        magento_connection_id=magento_connection_id,
        shopify_connection_id=shopify_connection_id,
    )
    now = _utcnow_naive()
    if row is None:
        row = ChannelPipelineReadiness(
            channel_code=channel,
            scope_key=scope_key,
            only_assigned=only_assigned,
            connection_id=connection_id,
            magento_connection_id=magento_connection_id,
            shopify_connection_id=shopify_connection_id,
            status="queued",
            refresh_requested_at=now,
        )
        session.add(row)
        session.flush()
        return row

    if row.status == "running" and not force:
        return row
    if row.status == "queued" and not force:
        return row

    row.status = "queued"
    row.refresh_requested_at = now
    row.refresh_started_at = None
    row.last_error = None
    row.locked_at = None
    row.locked_by = None
    session.flush()
    return row


def queue_default_pipeline_readiness_refreshes(session: Session, *, force: bool = False) -> int:
    queued = 0
    for scope in default_pipeline_scopes():
        row = queue_readiness_refresh(session, force=force, **scope)
        if row.status == "queued":
            queued += 1
    return queued


def enqueue_stale_readiness_refreshes(
    session: Session,
    *,
    max_age_hours: int = DEFAULT_REFRESH_HOURS,
) -> int:
    """Queue scopes whose cache is older than max_age_hours (periodic refresh)."""
    cutoff = _utcnow_naive() - timedelta(hours=max_age_hours)
    enqueued = 0
    for scope in default_pipeline_scopes():
        row = get_readiness_cache_row(session, **scope)
        if row is None:
            queue_readiness_refresh(session, **scope)
            enqueued += 1
            continue
        if row.status in {"queued", "running"}:
            continue
        if row.computed_at is None or _as_utc_naive(row.computed_at) < cutoff:
            queue_readiness_refresh(session, **scope, force=True)
            enqueued += 1
    return enqueued


def _reset_stale_running_rows(session: Session, *, worker_id: str) -> None:
    cutoff = _utcnow_naive() - timedelta(minutes=STALE_RUNNING_MINUTES)
    for row in session.scalars(
        select(ChannelPipelineReadiness)
        .where(ChannelPipelineReadiness.status == "running")
        .where(ChannelPipelineReadiness.refresh_started_at < cutoff)
    ).all():
        row.status = "queued"
        row.refresh_requested_at = _utcnow_naive()
        row.locked_at = None
        row.locked_by = None
        row.last_error = "Reset after stale running state"


def claim_next_readiness_refresh(session: Session, *, worker_id: str) -> Optional[ChannelPipelineReadiness]:
    _reset_stale_running_rows(session, worker_id=worker_id)
    row = session.scalars(
        select(ChannelPipelineReadiness)
        .where(ChannelPipelineReadiness.status == "queued")
        .order_by(ChannelPipelineReadiness.refresh_requested_at.nullsfirst(), ChannelPipelineReadiness.id)
        .limit(1)
        .with_for_update(skip_locked=True)
    ).first()
    if row is None:
        return None
    now = _utcnow_naive()
    row.status = "running"
    row.refresh_started_at = now
    row.locked_at = now
    row.locked_by = worker_id
    session.flush()
    return row


def run_readiness_refresh(session: Session, row: ChannelPipelineReadiness) -> Dict[str, Any]:
    from db.channel_push_readiness import build_channel_push_readiness

    report = build_channel_push_readiness(
        session,
        row.channel_code,
        connection_id=row.connection_id,
        magento_connection_id=row.magento_connection_id or row.connection_id,
        shopify_connection_id=row.shopify_connection_id or row.connection_id,
        only_assigned=row.only_assigned,
    )
    row.report = report
    row.status = "ready"
    row.computed_at = _utcnow_naive()
    row.last_error = None
    row.locked_at = None
    row.locked_by = None
    session.flush()
    return report


def process_one_readiness_refresh(session: Session, *, worker_id: str) -> bool:
    row = claim_next_readiness_refresh(session, worker_id=worker_id)
    if row is None:
        return False
    try:
        run_readiness_refresh(session, row)
    except Exception as exc:
        row.status = "failed"
        row.last_error = str(exc)
        row.locked_at = None
        row.locked_by = None
        session.flush()
        raise
    return True


def readiness_row_snapshot(row: Optional[ChannelPipelineReadiness]) -> Optional[Dict[str, Any]]:
    """Copy cache fields while the ORM row is still session-bound."""
    if row is None:
        return None
    report = row.report
    return {
        "status": row.status,
        "computed_at": row.computed_at,
        "refresh_requested_at": row.refresh_requested_at,
        "refresh_started_at": row.refresh_started_at,
        "last_error": row.last_error,
        "report": dict(report) if report else None,
    }


def cache_meta(
    row: Optional[ChannelPipelineReadiness] = None,
    *,
    snapshot: Optional[Dict[str, Any]] = None,
) -> Dict[str, Any]:
    snap = snapshot if snapshot is not None else readiness_row_snapshot(row)
    if snap is None:
        return {
            "state": "missing",
            "computed_at": None,
            "refresh_requested_at": None,
            "refresh_started_at": None,
            "last_error": None,
        }
    state = snap.get("status") or "missing"
    computed_at = snap.get("computed_at")
    if state == "ready" and computed_at:
        cutoff = _utcnow_naive() - timedelta(hours=DEFAULT_REFRESH_HOURS)
        if _as_utc_naive(computed_at) < cutoff:
            state = "stale"
    return {
        "state": state,
        "computed_at": computed_at.isoformat() if computed_at else None,
        "refresh_requested_at": snap["refresh_requested_at"].isoformat()
        if snap.get("refresh_requested_at")
        else None,
        "refresh_started_at": snap["refresh_started_at"].isoformat()
        if snap.get("refresh_started_at")
        else None,
        "last_error": snap.get("last_error"),
    }


def _recompute_pending(snapshot: Dict[str, Any], meta: Dict[str, Any]) -> bool:
    if meta["state"] not in {"queued", "running"}:
        return False
    requested_at = snapshot.get("refresh_requested_at")
    computed_at = snapshot.get("computed_at")
    if requested_at is None:
        return True
    if computed_at is None:
        return True
    return _as_utc_naive(requested_at) > _as_utc_naive(computed_at)


def readiness_api_payload(
    row: Optional[ChannelPipelineReadiness] = None,
    *,
    snapshot: Optional[Dict[str, Any]] = None,
    auto_queue: bool = False,
) -> Dict[str, Any]:
    snap = snapshot if snapshot is not None else readiness_row_snapshot(row)
    meta = cache_meta(snapshot=snap)
    pending_recompute = snap is not None and _recompute_pending(snap, meta)
    if snap is None or snap.get("report") is None or pending_recompute:
        payload: Dict[str, Any] = {
            "status": "pending" if meta["state"] in {"missing", "queued", "running"} else "ok",
            "cache": meta,
            "ready": False,
            "blockers": ["readiness_cache_not_ready"],
            "recommendations": ["Readiness is computed in the background; retry shortly or POST .../readiness/refresh."],
            "checks": {},
            "samples": {},
        }
        if auto_queue and meta["state"] == "missing":
            payload["cache"]["state"] = "queued"
        return payload

    report = dict(snap.get("report") or {})
    report["status"] = "ok" if meta["state"] == "ready" else meta["state"]
    report["cache"] = meta
    return report
