from datetime import datetime, timedelta from fastapi import APIRouter, Depends, Query from sqlalchemy import func, select from sqlalchemy.orm import Session from app.db import get_db from app.deps import require_admin from app.models import ( ApiCallLog, AuditLog, Document, Message, MessageReference, User, ) from app.schemas.audit import StatsOverviewResponse router = APIRouter() _RANGE_SPEC: dict[str, tuple[timedelta, int, str]] = { "1h": (timedelta(hours=1), 12, "5min"), "12h": (timedelta(hours=12), 12, "hour"), "24h": (timedelta(hours=24), 24, "hour"), "7d": (timedelta(days=7), 7, "day"), } def _resolve_range(range_code: str, now: datetime) -> tuple[datetime, datetime, int, str]: span, bucket_count, unit = _RANGE_SPEC[range_code] return now - span, now, bucket_count, unit def _bucket_index(dt: datetime, start: datetime, unit: str) -> int | None: if dt < start: return None delta = (dt - start).total_seconds() if unit == "5min": return int(delta // 300) if unit == "hour": return int(delta // 3600) if unit == "day": return int(delta // 86400) return None def _bucket_label(start: datetime, idx: int, unit: str) -> str: if unit == "5min": return (start + timedelta(minutes=5 * idx)).strftime("%H:%M") if unit == "hour": return (start + timedelta(hours=idx)).strftime("%H:00") if unit == "day": return (start + timedelta(days=idx)).strftime("%m-%d") return str(idx) def _pct(part: int, total: int) -> float: return round(part / total * 100, 1) if total else 0.0 def _delta_pct(current: int, previous: int) -> float | None: if previous == 0: return None return round((current - previous) / previous * 100, 1) @router.get("/overview", response_model=StatsOverviewResponse) def overview(db: Session = Depends(get_db), _: User = Depends(require_admin)) -> StatsOverviewResponse: question_count = db.scalar(select(func.count(Message.id))) or 0 active_user_count = db.scalar(select(func.count(User.id)).where(User.status == "active")) or 0 document_count = db.scalar(select(func.count(Document.id)).where(Document.status != "deleted")) or 0 audit_count = db.scalar(select(func.count(AuditLog.id))) or 0 token_usage = db.scalar(select(func.coalesce(func.sum(Message.token_usage), 0))) or 0 pu_usage = db.scalar(select(func.coalesce(func.sum(Message.pu_usage), 0))) or 0 return StatsOverviewResponse( question_count=question_count, active_user_count=active_user_count, document_count=document_count, audit_count=audit_count, token_usage=token_usage, pu_usage=pu_usage, ) @router.get("/dashboard/business") def dashboard_business( range_code: str = Query("24h", alias="range", pattern="^(1h|12h|24h|7d)$"), db: Session = Depends(get_db), _: User = Depends(require_admin), ) -> dict: now = datetime.utcnow() start, end, bucket_count, unit = _resolve_range(range_code, now) prev_start = start - (end - start) rows = db.execute( select( Message.user_id, Message.status, Message.latency_ms, Message.created_at, ).where(Message.created_at >= start, Message.created_at < end) ).all() prev_user_count = ( db.scalar( select(func.count(func.distinct(Message.user_id))).where( Message.created_at >= prev_start, Message.created_at < start ) ) or 0 ) total = len(rows) success = sum(1 for r in rows if r.status == "done") failed = sum(1 for r in rows if r.status == "failed") latencies = [r.latency_ms for r in rows if r.latency_ms] active_users = {r.user_id for r in rows} active_user_count = len(active_users) trend = [ {"bucket": _bucket_label(start, i, unit), "asked": 0, "success": 0, "failed": 0} for i in range(bucket_count) ] for r in rows: idx = _bucket_index(r.created_at, start, unit) if idx is None or idx >= bucket_count: continue trend[idx]["asked"] += 1 if r.status == "done": trend[idx]["success"] += 1 elif r.status == "failed": trend[idx]["failed"] += 1 top_rows = db.execute( select(Message.question, func.count(Message.id).label("cnt")) .where(Message.created_at >= start, Message.created_at < end) .group_by(Message.question) .order_by(func.count(Message.id).desc()) .limit(10) ).all() top_questions = [{"question": (r[0] or "").strip()[:120], "count": int(r[1])} for r in top_rows] return { "range": { "code": range_code, "start": start.isoformat(), "end": end.isoformat(), "bucket_unit": unit, }, "kpis": { "question_count": total, "active_user_count": active_user_count, "prev_active_user_count": int(prev_user_count), "user_delta_pct": _delta_pct(active_user_count, int(prev_user_count)), "avg_questions_per_user": round(total / active_user_count, 1) if active_user_count else 0, "avg_latency_ms": int(sum(latencies) / len(latencies)) if latencies else 0, "success_count": success, "success_ratio": _pct(success, total), "fail_count": failed, "fail_ratio": _pct(failed, total), }, "trend": trend, "top_questions": top_questions, } @router.get("/dashboard/resources") def dashboard_resources( range_code: str = Query("24h", alias="range", pattern="^(1h|12h|24h|7d)$"), db: Session = Depends(get_db), _: User = Depends(require_admin), ) -> dict: now = datetime.utcnow() start, end, bucket_count, unit = _resolve_range(range_code, now) document_count = ( db.scalar(select(func.count(Document.id)).where(Document.status != "deleted")) or 0 ) total_tokens = ( db.scalar( select(func.coalesce(func.sum(Message.token_usage), 0)).where( Message.created_at >= start, Message.created_at < end ) ) or 0 ) total_pu = ( db.scalar( select(func.coalesce(func.sum(Message.pu_usage), 0)).where( Message.created_at >= start, Message.created_at < end ) ) or 0 ) api_total = ( db.scalar( select(func.count(ApiCallLog.id)).where( ApiCallLog.created_at >= start, ApiCallLog.created_at < end ) ) or 0 ) api_success = ( db.scalar( select(func.count(ApiCallLog.id)).where( ApiCallLog.created_at >= start, ApiCallLog.created_at < end, ApiCallLog.status == "success", ) ) or 0 ) api_avg_latency = db.scalar( select(func.avg(ApiCallLog.latency_ms)).where( ApiCallLog.created_at >= start, ApiCallLog.created_at < end, ApiCallLog.latency_ms.is_not(None), ) ) msg_rows = db.execute( select(Message.created_at, Message.token_usage, Message.pu_usage).where( Message.created_at >= start, Message.created_at < end ) ).all() token_trend = [ {"bucket": _bucket_label(start, i, unit), "tokens": 0, "pu": 0} for i in range(bucket_count) ] for r in msg_rows: idx = _bucket_index(r.created_at, start, unit) if idx is None or idx >= bucket_count: continue token_trend[idx]["tokens"] += int(r.token_usage or 0) token_trend[idx]["pu"] += int(r.pu_usage or 0) ref_rows = db.execute( select(MessageReference.source_title, func.count(MessageReference.id).label("cnt")) .join(Message, Message.id == MessageReference.message_id) .where(Message.created_at >= start, Message.created_at < end) .group_by(MessageReference.source_title) .order_by(func.count(MessageReference.id).desc()) .limit(10) ).all() top_documents = [ {"title": (r[0] or "").strip()[:120], "count": int(r[1])} for r in ref_rows ] return { "range": { "code": range_code, "start": start.isoformat(), "end": end.isoformat(), "bucket_unit": unit, }, "kpis": { "document_count": int(document_count), "total_tokens": int(total_tokens), "total_pu": int(total_pu), "api_call_count": int(api_total), "api_success_ratio": _pct(int(api_success), int(api_total)), "avg_api_latency_ms": int(api_avg_latency) if api_avg_latency else 0, }, "token_trend": token_trend, "top_documents": top_documents, }