Files

159 lines
6.6 KiB
Python
Raw Permalink Normal View History

"""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,
"description_model": row.description_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 "description_model" in body.model_fields_set:
row.description_model = body.description_model
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