from __future__ import annotations

import os
from datetime import datetime, timedelta, timezone
from typing import Any, Optional

from sqlalchemy import or_, select
from sqlalchemy.orm import Session

from db.models_v2 import (
    V2ChannelPublishState,
    V2QueueMedia,
    V2QueueProductData,
    V2QueueRelations,
    V2SyncJob,
)


QUEUE_MODELS = {
    "product_data": V2QueueProductData,
    "media": V2QueueMedia,
    "relations": V2QueueRelations,
}

QUEUE_ORDER = ("product_data", "media", "relations")


def utcnow() -> datetime:
    return datetime.now(timezone.utc)


def stale_processing_minutes() -> int:
    return max(1, int(os.getenv("V2_SYNC_STALE_MINUTES", "120")))


def _queue_domain(model: type) -> str:
    for domain, candidate in QUEUE_MODELS.items():
        if candidate is model:
            return domain
    raise KeyError(f"Unknown V2 queue model: {model}")


def reset_stale_processing_rows(
    session: Session,
    *,
    domain: str | None = None,
) -> int:
    domains = [domain] if domain else list(QUEUE_ORDER)
    cutoff = utcnow() - timedelta(minutes=stale_processing_minutes())
    reset_count = 0
    for domain_name in domains:
        model = QUEUE_MODELS[domain_name]
        stale_rows = session.scalars(
            select(model)
            .where(model.status == "processing")
            .where(model.locked_at < cutoff)
            .with_for_update(skip_locked=True)
        ).all()
        for row in stale_rows:
            if (row.attempt_count or 0) >= (row.max_attempts or 5):
                row.status = "dead_letter"
                row.completed_at = utcnow()
                row.last_error = "Exceeded max attempts after stale processing reset"
            else:
                row.status = "pending"
                row.last_error = "Reset after stale processing lock"
            row.locked_at = None
            row.locked_by = None
            _set_publish_state_for_queue_row(
                session,
                row,
                status="failed" if row.status == "dead_letter" else "pending",
                last_error=row.last_error,
                verification_result={
                    "status": "stale_reset",
                    "queue_status": row.status,
                    "domain": domain_name,
                },
            )
            reset_count += 1
    if reset_count:
        session.flush()
    return reset_count


def claim_next_queue_item(
    session: Session,
    *,
    worker_id: str,
    domain: str | None = None,
) -> tuple[str, Any] | None:
    reset_stale_processing_rows(session, domain=domain)
    domains = [domain] if domain else list(QUEUE_ORDER)
    now = utcnow()
    for domain_name in domains:
        model = QUEUE_MODELS[domain_name]
        row = session.scalars(
            select(model)
            .where(model.status == "pending")
            .where(or_(model.scheduled_at.is_(None), model.scheduled_at <= now))
            .where(model.attempt_count < model.max_attempts)
            .order_by(model.priority.asc(), model.scheduled_at.asc(), model.id.asc())
            .limit(1)
            .with_for_update(skip_locked=True)
        ).first()
        if row is None:
            continue
        row.status = "processing"
        row.locked_at = now
        row.locked_by = worker_id
        row.attempt_count = int(row.attempt_count or 0) + 1
        row.last_error = None
        _set_publish_state_for_queue_row(
            session,
            row,
            status="processing",
            last_error=None,
        )
        _mark_sync_job_running(session, row.sync_job_id, started_at=now)
        session.flush()
        return domain_name, row
    return None


def mark_queue_item_done(
    session: Session,
    *,
    domain: str,
    row_id: int,
    response_payload: dict[str, Any] | None = None,
    channel_ref_id: str | None = None,
    verification_mode: str = "sampled",
) -> Any:
    model = QUEUE_MODELS[domain]
    row = session.get(model, row_id)
    if row is None:
        raise ValueError(f"Queue row {domain}:{row_id} not found")
    now = utcnow()
    row.status = "done"
    row.completed_at = now
    row.locked_at = None
    row.locked_by = None
    row.last_error = None
    row.last_response = response_payload or {"status": "done"}
    _set_publish_state_for_queue_row(
        session,
        row,
        status="verified",
        channel_ref_id=channel_ref_id,
        last_pushed_at=now,
        last_verified_at=now,
        verification_mode=verification_mode,
        verification_result=response_payload or {"status": "done"},
        last_error=None,
    )
    _increment_sync_job_counter(session, row.sync_job_id, success_delta=1)
    session.flush()
    return row


def mark_queue_item_failed(
    session: Session,
    *,
    domain: str,
    row_id: int,
    error: str,
    response_payload: dict[str, Any] | None = None,
) -> Any:
    model = QUEUE_MODELS[domain]
    row = session.get(model, row_id)
    if row is None:
        raise ValueError(f"Queue row {domain}:{row_id} not found")
    attempts = int(row.attempt_count or 0)
    max_attempts = int(row.max_attempts or 5)
    terminal_status = "dead_letter" if attempts >= max_attempts else "failed"
    now = utcnow()
    row.status = terminal_status
    row.completed_at = now if terminal_status == "dead_letter" else None
    row.locked_at = None
    row.locked_by = None
    row.last_error = str(error)[:4000]
    row.last_response = response_payload
    _set_publish_state_for_queue_row(
        session,
        row,
        status="failed",
        verification_result=response_payload or {
            "status": terminal_status,
            "error": str(error),
        },
        last_error=str(error)[:4000],
    )
    _increment_sync_job_counter(session, row.sync_job_id, failure_delta=1)
    session.flush()
    return row


def retry_failed_queue_item(
    session: Session,
    *,
    domain: str,
    row_id: int,
) -> Any:
    model = QUEUE_MODELS[domain]
    row = session.get(model, row_id)
    if row is None:
        raise ValueError(f"Queue row {domain}:{row_id} not found")
    row.status = "pending"
    row.locked_at = None
    row.locked_by = None
    row.completed_at = None
    _set_publish_state_for_queue_row(
        session,
        row,
        status="pending",
        last_error=None,
        verification_result=None,
    )
    session.flush()
    return row


def _set_publish_state_for_queue_row(
    session: Session,
    row: Any,
    *,
    status: str,
    channel_ref_id: str | None = None,
    last_pushed_at: datetime | None = None,
    last_verified_at: datetime | None = None,
    verification_mode: str | None = None,
    verification_result: dict[str, Any] | None = None,
    last_error: str | None = None,
) -> None:
    domain_name = _queue_domain(type(row))
    state = session.get(
        V2ChannelPublishState,
        {
            "product_id": int(row.product_id),
            "channel": str(row.channel),
            "domain": domain_name,
        },
    )
    if state is None:
        state = V2ChannelPublishState(
            product_id=int(row.product_id),
            channel=str(row.channel),
            domain=domain_name,
            status=status,
        )
        session.add(state)
    state.status = status
    state.block_reason = None
    state.last_queue_item_id = int(row.id)
    state.updated_at = utcnow()
    if channel_ref_id is not None:
        state.channel_ref_id = channel_ref_id
    if last_pushed_at is not None:
        state.last_pushed_at = last_pushed_at
    if last_verified_at is not None:
        state.last_verified_at = last_verified_at
    if verification_mode is not None:
        state.verification_mode = verification_mode
    if verification_result is not None:
        state.verification_result = verification_result
    state.last_error = last_error


def _mark_sync_job_running(session: Session, sync_job_id: str | None, *, started_at: datetime) -> None:
    if not sync_job_id:
        return
    row = session.get(V2SyncJob, sync_job_id)
    if row is None:
        return
    if row.status == "pending":
        row.status = "running"
    if row.started_at is None:
        row.started_at = started_at


def _increment_sync_job_counter(
    session: Session,
    sync_job_id: str | None,
    *,
    success_delta: int = 0,
    failure_delta: int = 0,
) -> None:
    if not sync_job_id:
        return
    row = session.get(V2SyncJob, sync_job_id)
    if row is None:
        return
    row.success_count = int(row.success_count or 0) + int(success_delta)
    row.failure_count = int(row.failure_count or 0) + int(failure_delta)
    _finalize_sync_job_if_complete(session, row)


def _finalize_sync_job_if_complete(session: Session, row: V2SyncJob) -> None:
    if not row.id or int(row.item_count or 0) <= 0:
        return
    pending_total = 0
    for model in QUEUE_MODELS.values():
        pending_total += int(
            session.scalar(
                select(model.id)
                .where(model.sync_job_id == row.id)
                .where(model.status.in_(("pending", "processing")))
                .limit(1)
            )
            is not None
        )
    if pending_total:
        return
    row.completed_at = utcnow()
    row.status = "failed" if int(row.failure_count or 0) > 0 else "completed"
