2026-09-07 23:02:55 +08:00
|
|
|
"""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,
|
2026-09-09 20:24:47 +08:00
|
|
|
"description_model": row.description_model,
|
2026-09-07 23:02:55 +08:00
|
|
|
"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
|
2026-09-09 20:24:47 +08:00
|
|
|
if "description_model" in body.model_fields_set:
|
|
|
|
|
row.description_model = body.description_model
|
2026-09-07 23:02:55 +08:00
|
|
|
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
|