from datetime import datetime, timedelta, timezone
from uuid import uuid4

from sqlalchemy import func, select
from sqlalchemy.orm import Session

from app.core.models import AIUsageLog, AIUsageSource, AIUsageStatus, TenantAIQuota


class AIQuotaExceeded(Exception):
    def __init__(self, window: str):
        self.window = window
        super().__init__(f"AI {window} credit limit reached")


def _now() -> datetime:
    return datetime.now(timezone.utc).replace(tzinfo=None)


def _month_start(now: datetime) -> datetime:
    return now.replace(day=1, hour=0, minute=0, second=0, microsecond=0)


def _week_start(now: datetime) -> datetime:
    midnight = now.replace(hour=0, minute=0, second=0, microsecond=0)
    return midnight - timedelta(days=midnight.weekday())


def _usage_count(db: Session, tenant_id: int, started_at: datetime) -> int:
    return db.scalar(
        select(func.count(AIUsageLog.id)).where(
            AIUsageLog.tenant_id == tenant_id,
            AIUsageLog.created_at >= started_at,
            AIUsageLog.status.in_((AIUsageStatus.reserved, AIUsageStatus.completed)),
        )
    ) or 0


def _monthly_completed_count(db: Session, tenant_id: int, now: datetime) -> int:
    return db.scalar(
        select(func.count(AIUsageLog.id)).where(
            AIUsageLog.tenant_id == tenant_id,
            AIUsageLog.created_at >= _month_start(now),
            AIUsageLog.status == AIUsageStatus.completed,
        )
    ) or 0


def reserve_ai_credit(
    db: Session,
    *,
    tenant_id: int,
    source: AIUsageSource,
    user_id: int | None = None,
    now: datetime | None = None,
) -> AIUsageLog:
    now = now or _now()
    quota = db.execute(
        select(TenantAIQuota).where(TenantAIQuota.tenant_id == tenant_id).with_for_update()
    ).scalar_one_or_none()
    if quota is None or not quota.enabled:
        raise AIQuotaExceeded("access")
    if not quota.unlimited:
        for window, started_at, limit in (
            ("5-hour", now - timedelta(hours=5), quota.five_hour_credit_limit),
            ("weekly", _week_start(now), quota.weekly_credit_limit),
            ("monthly", _month_start(now), quota.monthly_credit_limit),
        ):
            if limit and _usage_count(db, tenant_id, started_at) >= limit:
                raise AIQuotaExceeded(window)

    usage = AIUsageLog(
        tenant_id=tenant_id,
        user_id=user_id,
        request_id=uuid4().hex,
        source=source,
        credits=1,
        status=AIUsageStatus.reserved,
        estimated=False,
        created_at=now,
    )
    db.add(usage)
    db.commit()
    db.refresh(usage)
    return usage


def complete_ai_credit(db: Session, usage: AIUsageLog) -> None:
    usage.status = AIUsageStatus.completed
    quota = db.execute(
        select(TenantAIQuota).where(TenantAIQuota.tenant_id == usage.tenant_id).with_for_update()
    ).scalar_one_or_none()
    if quota is not None:
        quota.used_credits = _monthly_completed_count(db, usage.tenant_id, _now())
    db.commit()


def fail_ai_credit(db: Session, usage: AIUsageLog, error: str) -> None:
    usage.status = AIUsageStatus.failed
    usage.error_summary = error[:1000]
    db.commit()
