import asyncio
import logging
import uuid
from datetime import UTC, datetime

from sqlalchemy import delete, select
from sqlalchemy.ext.asyncio import AsyncSession

from .auth import _audit
from .computer_client import ComputerClient, HttpComputerClient
from .config import get_settings
from .database import session_factory
from .errors import AppError
from .models import (
    ApprovalRequest,
    ComputerAction,
    ComputerArtifact,
    ComputerControlLease,
    ComputerSession,
    IdempotencyRecord,
)

logger = logging.getLogger("hayva.maintenance")


async def run_maintenance_once(
    session: AsyncSession, *, now: datetime | None = None,
    client: ComputerClient | None = None,
) -> dict[str, int]:
    """Reconcile expired authority and retention records in bounded batches."""
    current = now or datetime.now(UTC)
    limit = get_settings().maintenance_batch_size
    counts = {
        "leases_released": 0,
        "approvals_expired": 0,
        "artifacts_deleted": 0,
        "artifacts_pending_deletion": 0,
        "idempotency_deleted": 0,
        "idempotency_unknown": 0,
    }

    leases = list((await session.scalars(
        select(ComputerControlLease).where(
            ComputerControlLease.released_at.is_(None),
            ComputerControlLease.expires_at <= current,
        ).order_by(ComputerControlLease.expires_at).limit(limit).with_for_update()
    )).all())
    for lease in leases:
        lease.released_at = current
        lease.release_reason = "lease_expired"
        record = await session.scalar(select(ComputerSession).where(
            ComputerSession.workspace_id == lease.workspace_id,
            ComputerSession.id == lease.session_id,
        ).with_for_update())
        if record and record.status not in {"stopped", "failed"}:
            record.status = "paused"
            record.control_owner = "none"
            record.paused_at = current
            record.failure_code = "COMPUTER_CONTROL_LEASE_EXPIRED"
            record.failure_message = "Control authority expired and requires an explicit resume."
            record.version += 1
            actions = list((await session.scalars(select(ComputerAction).where(
                ComputerAction.workspace_id == record.workspace_id,
                ComputerAction.session_id == record.id,
                ComputerAction.status == "running",
            ).with_for_update())).all())
            for action in actions:
                action.status = "unknown" if action.access_type == "write" else "failed"
                action.error_code = "COMPUTER_CONTROL_LEASE_EXPIRED"
                action.completed_at = current
            await _audit(
                session, workspace_id=record.workspace_id,
                actor_id=record.requested_by_user_id, actor_type="system",
                event_type="computer.lease_expired",
                request_id=f"maintenance-{uuid.uuid4().hex}",
                data={
                    "computer_session_id": str(record.id),
                    "lease_id": str(lease.id),
                    "fencing_token": lease.fencing_token,
                    "unknown_write_actions": sum(
                        item.status == "unknown" for item in actions
                    ),
                },
            )
        counts["leases_released"] += 1

    approvals = list((await session.scalars(
        select(ApprovalRequest).where(
            ApprovalRequest.status == "pending",
            ApprovalRequest.expires_at <= current,
        ).order_by(ApprovalRequest.expires_at).limit(limit).with_for_update()
    )).all())
    for approval in approvals:
        approval.status = "expired"
        approval.decided_at = current
        action = await session.scalar(select(ComputerAction).where(
            ComputerAction.workspace_id == approval.workspace_id,
            ComputerAction.id == approval.action_id,
        ).with_for_update())
        if action and action.status == "awaiting_approval":
            action.status = "blocked"
            action.error_code = "COMPUTER_APPROVAL_EXPIRED"
            action.completed_at = current
        await _audit(
            session, workspace_id=approval.workspace_id,
            actor_id=approval.requested_by_id, actor_type="system",
            event_type="computer.approval_expired",
            request_id=f"maintenance-{uuid.uuid4().hex}",
            data={
                "approval_id": str(approval.id),
                "action_id": str(approval.action_id),
            },
        )
        counts["approvals_expired"] += 1

    artifacts = list((await session.scalars(
        select(ComputerArtifact).where(
            ComputerArtifact.status.in_(("available", "expired")),
            ComputerArtifact.expires_at.is_not(None),
            ComputerArtifact.expires_at <= current,
        ).order_by(ComputerArtifact.expires_at).limit(limit).with_for_update()
    )).all())
    artifact_client = client or HttpComputerClient()
    for artifact in artifacts:
        artifact.status = "expired"
        try:
            await artifact_client.delete_artifact(
                workspace_id=artifact.workspace_id,
                session_id=artifact.session_id,
                artifact_id=artifact.id,
                request_id=f"maintenance-{uuid.uuid4().hex}",
            )
        except AppError:
            counts["artifacts_pending_deletion"] += 1
        else:
            artifact.status = "deleted"
            counts["artifacts_deleted"] += 1

    stale_in_progress = list((await session.scalars(
        select(IdempotencyRecord).where(
            IdempotencyRecord.status == "in_progress",
            IdempotencyRecord.expires_at <= current,
        ).limit(limit).with_for_update()
    )).all())
    for record in stale_in_progress:
        record.status = "unknown"
        record.response_status = 409
        record.response_body = {
            "success": False,
            "error": {"code": "IDEMPOTENCY_OUTCOME_UNKNOWN"},
        }
    counts["idempotency_unknown"] = len(stale_in_progress)
    result = await session.execute(delete(IdempotencyRecord).where(
        IdempotencyRecord.status.in_(("completed", "failed")),
        IdempotencyRecord.expires_at <= current,
    ))
    counts["idempotency_deleted"] = result.rowcount or 0
    await session.commit()
    return counts


async def maintenance_loop(stop_event: asyncio.Event) -> None:
    interval = get_settings().maintenance_interval_seconds
    while not stop_event.is_set():
        try:
            async with session_factory() as session:
                counts = await run_maintenance_once(session)
            if any(counts.values()):
                logger.info("maintenance_completed counts=%s", counts)
        except Exception:
            logger.exception("maintenance_failed")
        try:
            await asyncio.wait_for(stop_event.wait(), timeout=interval)
        except TimeoutError:
            pass
