ai.py 3.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384
  1. """AI router - Kimi API integration endpoints (P3-M1)."""
  2. from typing import List
  3. from fastapi import APIRouter, HTTPException, Query
  4. from sqlalchemy import func
  5. from ..services.ai_client import get_kimi_client
  6. from ..schemas.ai import (
  7. ChatRequest, ChatResponse, AIHealthResponse,
  8. AICallLogResponse, AIUsageSummary
  9. )
  10. from ..database import SessionLocal
  11. from ..models.ai_call_log import AICallLog
  12. router = APIRouter(prefix="/api/ai", tags=["AI"])
  13. @router.get("/health", response_model=AIHealthResponse)
  14. def ai_health_check():
  15. """Check AI API key validity and list available models."""
  16. client = get_kimi_client()
  17. return client.health_check()
  18. @router.post("/chat", response_model=ChatResponse)
  19. def ai_chat(request: ChatRequest):
  20. """Send a chat completion request to Kimi API."""
  21. client = get_kimi_client()
  22. if not client.is_configured:
  23. raise HTTPException(status_code=503, detail="KIMI_API_KEY is not configured")
  24. try:
  25. result = client.chat(
  26. messages=[m.model_dump() for m in request.messages],
  27. model=request.model,
  28. temperature=request.temperature,
  29. max_tokens=request.max_tokens,
  30. system_prompt=request.system_prompt,
  31. )
  32. return ChatResponse(**result)
  33. except Exception as e:
  34. raise HTTPException(status_code=500, detail=str(e))
  35. @router.get("/logs", response_model=List[AICallLogResponse])
  36. def list_ai_logs(
  37. limit: int = Query(50, ge=1, le=200),
  38. offset: int = Query(0, ge=0),
  39. status: str = Query(None, description="success/failed"),
  40. ):
  41. """List AI call logs for cost tracking and debugging."""
  42. db = SessionLocal()
  43. query = db.query(AICallLog)
  44. if status:
  45. query = query.filter(AICallLog.status == status)
  46. logs = query.order_by(AICallLog.id.desc()).offset(offset).limit(limit).all()
  47. db.close()
  48. return [
  49. AICallLogResponse(
  50. id=log.id, endpoint=log.endpoint, model=log.model,
  51. status=log.status, duration_ms=log.duration_ms,
  52. prompt_tokens=log.prompt_tokens, completion_tokens=log.completion_tokens,
  53. total_tokens=log.total_tokens, error=log.error,
  54. created_at=log.created_at.isoformat() if log.created_at else "",
  55. )
  56. for log in logs
  57. ]
  58. @router.get("/usage", response_model=AIUsageSummary)
  59. def ai_usage_summary():
  60. """Get AI usage summary (total calls, tokens, duration)."""
  61. db = SessionLocal()
  62. total = db.query(func.count(AICallLog.id)).scalar() or 0
  63. success = db.query(func.count(AICallLog.id)).filter(AICallLog.status == "success").scalar() or 0
  64. failed = total - success
  65. total_tokens = db.query(func.coalesce(func.sum(AICallLog.total_tokens), 0)).scalar() or 0
  66. prompt_tokens = db.query(func.coalesce(func.sum(AICallLog.prompt_tokens), 0)).scalar() or 0
  67. completion_tokens = db.query(func.coalesce(func.sum(AICallLog.completion_tokens), 0)).scalar() or 0
  68. total_duration = db.query(func.coalesce(func.sum(AICallLog.duration_ms), 0)).scalar() or 0
  69. db.close()
  70. return AIUsageSummary(
  71. total_calls=total, success_calls=success, failed_calls=failed,
  72. total_tokens=total_tokens, total_prompt_tokens=prompt_tokens,
  73. total_completion_tokens=completion_tokens, total_duration_ms=total_duration,
  74. )