"""Authenticated AI HTTP surface; streams project only server-owned state.""" import asyncio from fastapi import APIRouter, Depends, HTTPException, Request from fastapi.encoders import jsonable_encoder from fastapi.responses import StreamingResponse from sqlalchemy import select from ..models import AIConversation, AIMessage, AIRun, AISettings from ..security import cipher, require_auth, token_hash from .contracts import Decision, ModelSettingsInput, RunInput from .provider import public_error, test_capabilities from .runtime import uid def settings_output(row): return { "base_url": row.base_url, "model": row.model, "protocol": row.protocol, "configured": bool(row.api_key_encrypted), "enabled": row.enabled, "ready": row.tested_revision == row.revision, "test_results": row.test_results, } def router(runtime): api = APIRouter(prefix="/api/v1/ai", tags=["ai"], dependencies=[Depends(require_auth)]) @api.get("/settings") async def get_settings(): async with runtime.sessions() as db: return settings_output(await db.get(AISettings, 1)) @api.put("/settings") async def save_settings(body: ModelSettingsInput): async with runtime.lock: async with runtime.sessions.begin() as db: row = await db.get(AISettings, 1) key = body.api_key.get_secret_value() if body.api_key else None if (row.base_url != body.base_url or not row.api_key_encrypted) and not key: raise HTTPException(422, "首次配置或更换 Base URL 时必须重新输入 API Key") changed = any( getattr(row, k) != getattr(body, k) for k in ("base_url", "model", "protocol") ) or bool(key) if changed: row.revision += 1 row.tested_revision, row.test_results = None, {} row.base_url, row.model, row.protocol = body.base_url, body.model, body.protocol if key: row.api_key_encrypted = cipher(runtime.settings).encrypt(key.encode()).decode() row.enabled = body.enabled and row.tested_revision == row.revision return settings_output(row) @api.post("/settings/test") async def test_settings(request: Request): token = token_hash(request.cookies["wq_session"]) async with runtime.sessions() as db: row = await db.get(AISettings, 1) if not row.api_key_encrypted: raise HTTPException(409, "请先保存模型配置") revision = row.revision try: async with asyncio.timeout(runtime.settings.ai_timeout): async with runtime.model_factory(row, runtime.settings) as model: results = await test_capabilities(model) except Exception as exc: results = {k: {"ok": False, "message": public_error(exc)} for k in ("answer", "stream", "tools")} await runtime.authorize(token) async with runtime.lock: async with runtime.sessions.begin() as db: row = await db.get(AISettings, 1) if row.revision != revision: raise HTTPException(409, "测试期间配置已变化,请重新测试") row.test_results = results row.tested_revision = revision if all(v["ok"] for v in results.values()) else None if row.tested_revision is None: row.enabled = False return settings_output(row) @api.get("/conversations") async def conversations(): async with runtime.sessions() as db: rows = ( await db.scalars( select(AIConversation) .where(AIConversation.admin_id == 1) .order_by(AIConversation.updated_at.desc()) ) ).all() return [{"id": r.id, "title": r.title, "updated_at": r.updated_at} for r in rows] @api.post("/conversations", status_code=201) async def create_conversation(): async with runtime.sessions.begin() as db: row = AIConversation(id=uid()) db.add(row) await db.flush() return {"id": row.id, "title": row.title} @api.get("/conversations/{conversation_id}") async def conversation(conversation_id: str): async with runtime.sessions() as db: row = await db.get(AIConversation, conversation_id) if not row or row.admin_id != 1: raise HTTPException(404, "会话不存在") messages = ( await db.scalars( select(AIMessage) .where(AIMessage.conversation_id == conversation_id) .order_by(AIMessage.created_at, AIMessage.id) ) ).all() runs = ( await db.scalars( select(AIRun).where(AIRun.conversation_id == conversation_id).order_by(AIRun.created_at) ) ).all() data = { "id": row.id, "title": row.title, "messages": [{"id": m.id, "role": m.role, "parts": m.parts} for m in messages], } data["runs"] = [await runtime.snapshot(r.id) for r in runs] return jsonable_encoder(data) def stream(run_id): return StreamingResponse( runtime.events(run_id), media_type="text/event-stream", headers={"x-vercel-ai-ui-message-stream": "v1", "X-Accel-Buffering": "no", "X-AI-Run-ID": run_id}, ) @api.post("/conversations/{conversation_id}/runs") async def create_run(conversation_id: str, body: RunInput, request: Request): run_id = await runtime.create_run(conversation_id, body, token_hash(request.cookies["wq_session"])) return stream(run_id) @api.get("/runs/{run_id}") async def get_run(run_id: str): return await runtime.snapshot(run_id) @api.post("/runs/{run_id}/cancel") async def cancel_run(run_id: str): await runtime.cancel(run_id) return await runtime.snapshot(run_id) @api.post("/approvals/{approval_id}/decision") async def decide(approval_id: str, body: Decision, request: Request): run_id = await runtime.decision(approval_id, body.approved, token_hash(request.cookies["wq_session"])) return stream(run_id) return api