"""LLM usage + cost tracking, monthly budget guard.""" from __future__ import annotations from datetime import UTC, datetime from typing import Any from sqlalchemy import func, select from app.core.config import get_settings from app.db import SessionLocal from app.llm.router import router from app.models import LLMUsage async def record_usage(task: str, course: str, usage: dict[str, Any], latency_ms: int) -> None: async with SessionLocal() as session: session.add(LLMUsage( generation_id=usage.get("generation_id") or None, task=task, course_code=course, model=usage.get("model", ""), provider=usage.get("provider", ""), tokens_in=int(usage.get("input_tokens", 0)), tokens_out=int(usage.get("output_tokens", 0)), cost_usd=float(usage.get("cost_usd", 0.0)), latency_ms=latency_ms)) await session.commit() async def month_cost() -> float: start = datetime.now(UTC).replace(day=1, hour=0, minute=0, second=0, microsecond=0, tzinfo=None) async with SessionLocal() as session: total = await session.scalar(select(func.coalesce(func.sum(LLMUsage.cost_usd), 0.0)) .where(LLMUsage.created_at >= start)) return float(total or 0.0) async def budget_status(budget_override: float | None = None) -> dict[str, Any]: budget = budget_override if budget_override is not None else get_settings().LLM_MONTHLY_BUDGET_USD spent = await month_cost() ratio = spent / budget if budget > 0 else 0.0 router.budget_exceeded = ratio >= 1.0 return {"budget_usd": budget, "spent_usd": round(spent, 4), "ratio": round(ratio, 3), "warning": ratio >= 0.8, "exceeded": ratio >= 1.0} async def daily_costs(days: int = 30) -> list[dict[str, Any]]: async with SessionLocal() as session: day = func.date(LLMUsage.created_at) rows = await session.execute( select(day.label("day"), LLMUsage.model, func.sum(LLMUsage.cost_usd), func.sum(LLMUsage.tokens_in), func.sum(LLMUsage.tokens_out), func.count()) .group_by(day, LLMUsage.model).order_by(day.desc()).limit(days * 8)) return [{"day": str(d), "model": m, "cost_usd": round(float(c or 0), 5), "tokens_in": int(ti or 0), "tokens_out": int(to or 0), "calls": int(n)} for d, m, c, ti, to, n in rows] async def costs_by_course() -> list[dict[str, Any]]: async with SessionLocal() as session: rows = await session.execute( select(LLMUsage.course_code, func.sum(LLMUsage.cost_usd), func.count()) .group_by(LLMUsage.course_code)) return [{"course": c or "—", "cost_usd": round(float(s or 0), 4), "calls": int(n)} for c, s, n in rows]