"""AI router - Kimi API integration endpoints (P3-M1).""" from typing import List from fastapi import APIRouter, HTTPException, Query from sqlalchemy import func from ..services.ai_client import get_kimi_client from ..schemas.ai import ( ChatRequest, ChatResponse, AIHealthResponse, AICallLogResponse, AIUsageSummary ) from ..database import SessionLocal from ..models.ai_call_log import AICallLog router = APIRouter(prefix="/api/ai", tags=["AI"]) @router.get("/health", response_model=AIHealthResponse) def ai_health_check(): """Check AI API key validity and list available models.""" client = get_kimi_client() return client.health_check() @router.post("/chat", response_model=ChatResponse) def ai_chat(request: ChatRequest): """Send a chat completion request to Kimi API.""" client = get_kimi_client() if not client.is_configured: raise HTTPException(status_code=503, detail="KIMI_API_KEY is not configured") try: result = client.chat( messages=[m.model_dump() for m in request.messages], model=request.model, temperature=request.temperature, max_tokens=request.max_tokens, system_prompt=request.system_prompt, ) return ChatResponse(**result) except Exception as e: raise HTTPException(status_code=500, detail=str(e)) @router.get("/logs", response_model=List[AICallLogResponse]) def list_ai_logs( limit: int = Query(50, ge=1, le=200), offset: int = Query(0, ge=0), status: str = Query(None, description="success/failed"), ): """List AI call logs for cost tracking and debugging.""" db = SessionLocal() query = db.query(AICallLog) if status: query = query.filter(AICallLog.status == status) logs = query.order_by(AICallLog.id.desc()).offset(offset).limit(limit).all() db.close() return [ AICallLogResponse( id=log.id, endpoint=log.endpoint, model=log.model, status=log.status, duration_ms=log.duration_ms, prompt_tokens=log.prompt_tokens, completion_tokens=log.completion_tokens, total_tokens=log.total_tokens, error=log.error, created_at=log.created_at.isoformat() if log.created_at else "", ) for log in logs ] @router.get("/usage", response_model=AIUsageSummary) def ai_usage_summary(): """Get AI usage summary (total calls, tokens, duration).""" db = SessionLocal() total = db.query(func.count(AICallLog.id)).scalar() or 0 success = db.query(func.count(AICallLog.id)).filter(AICallLog.status == "success").scalar() or 0 failed = total - success total_tokens = db.query(func.coalesce(func.sum(AICallLog.total_tokens), 0)).scalar() or 0 prompt_tokens = db.query(func.coalesce(func.sum(AICallLog.prompt_tokens), 0)).scalar() or 0 completion_tokens = db.query(func.coalesce(func.sum(AICallLog.completion_tokens), 0)).scalar() or 0 total_duration = db.query(func.coalesce(func.sum(AICallLog.duration_ms), 0)).scalar() or 0 db.close() return AIUsageSummary( total_calls=total, success_calls=success, failed_calls=failed, total_tokens=total_tokens, total_prompt_tokens=prompt_tokens, total_completion_tokens=completion_tokens, total_duration_ms=total_duration, )