Files
worldquant-alpha-system/backend/app/ai/runtime.py
T

659 lines
30 KiB
Python
Raw Normal View History

"""Single-process AI executor with authoritative history and transactional confirmations.
The stream is a view of an independently owned task. Losing the browser connection
does not cancel it. Model calls are never retried by replaying business mutations.
"""
import asyncio
import json
import logging
import time
from dataclasses import asdict, dataclass, field
from uuid import uuid4
from fastapi import HTTPException
from fastapi.encoders import jsonable_encoder
from pydantic import ValidationError
from pydantic_ai import Agent, AgentRunResultEvent, CallDeferred, ModelRetry
from pydantic_ai.messages import (
ModelMessagesTypeAdapter,
ModelRequest,
ModelResponse,
PartDeltaEvent,
PartStartEvent,
RetryPromptPart,
TextPart,
TextPartDelta,
ToolCallPart,
ToolReturnPart,
UserPromptPart,
)
from pydantic_ai.tools import DeferredToolRequests, DeferredToolResults, Tool
from pydantic_ai.usage import RunUsage, UsageLimits
from pydantic_core import to_jsonable_python
from sqlalchemy import select, update
from ..business import Business
from ..models import AIConversation, AIMessage, AIRun, AISettings, AIToolCall, LoginSession, now
from .capabilities import ToolContext, model_result
from .provider import ensure_complete, model_connection, public_error
from .tools import CAPABILITIES, INSTRUCTIONS, presentation
logger = logging.getLogger(__name__)
def uid():
return str(uuid4())
@dataclass
class LiveRun:
message_id: str
events: list = field(default_factory=list)
condition: asyncio.Condition = field(default_factory=asyncio.Condition)
task: asyncio.Task | None = None
done: bool = False
parts: list = field(default_factory=list)
text_id: str | None = None
last_save: float = 0
async def emit(self, event):
async with self.condition:
self.events.append(event)
self.condition.notify_all()
class AIRuntime:
def __init__(self, sessions, settings, runner, model_factory=None):
self.sessions, self.settings, self.runner = sessions, settings, runner
self.model_factory = model_factory or model_connection
self.lock = asyncio.Lock()
self.live: dict[str, LiveRun] = {}
self.stopping = False
async def authorize(self, token):
async with self.sessions() as db:
session = await db.get(LoginSession, token)
if session is None or session.expires_at.replace(tzinfo=now().tzinfo) <= now():
raise HTTPException(401, "系统登录已过期,请重新登录")
async def start(self):
async with self.sessions.begin() as db:
if not await db.get(AISettings, 1):
db.add(AISettings(id=1))
await db.execute(
update(AIRun)
.where(AIRun.status == "running")
.values(
status="interrupted", error="服务已重启,本轮执行中断;已完成操作保留", updated_at=now()
)
)
# A crash before the deferred SDK checkpoint cannot leave executable orphan approvals.
await db.execute(
update(AIToolCall)
.where(
AIToolCall.status == "pending",
AIToolCall.run_id.in_(select(AIRun.id).where(AIRun.status == "interrupted")),
)
.values(status="cancelled")
)
async def stop(self):
self.stopping = True
tasks = [live.task for live in self.live.values() if live.task]
for task in tasks:
task.cancel()
await asyncio.gather(*tasks, return_exceptions=True)
async def config(self, db):
config = await db.get(AISettings, 1)
if not config or not config.enabled or config.tested_revision != config.revision:
raise HTTPException(409, "请在个人信息中配置、测试并启用大模型服务")
return config
async def create_run(self, conversation_id, body, token):
await self.authorize(token)
async with self.lock:
async with self.sessions.begin() as db:
conversation = await db.get(AIConversation, conversation_id)
if not conversation or conversation.admin_id != 1:
raise HTTPException(404, "会话不存在")
previous = await db.scalar(
select(AIRun).where(
AIRun.conversation_id == conversation_id, AIRun.request_id == body.request_id
)
)
if previous:
if previous.context != body.context.model_dump(mode="json"):
raise HTTPException(409, "请求标识已用于其他上下文")
user = await db.scalar(
select(AIMessage).where(AIMessage.run_id == previous.id, AIMessage.role == "user")
)
if not user or user.parts[0]["text"] != body.message:
raise HTTPException(409, "请求标识已用于其他消息")
return previous.id
active = await db.scalar(
select(AIRun.id).where(
AIRun.conversation_id == conversation_id,
AIRun.status.in_(("running", "waiting_approval")),
)
)
if active:
raise HTTPException(409, "请先完成、拒绝或停止当前执行")
config = await self.config(db)
history = []
past = (
await db.scalars(
select(AIRun)
.where(
AIRun.conversation_id == conversation_id,
AIRun.status.in_(("completed", "failed", "cancelled", "interrupted")),
)
.order_by(AIRun.created_at.desc())
.limit(10)
)
).all()
for item in reversed(past):
if item.status == "completed":
history.extend(item.model_messages[item.history_start :])
else:
# A failed continuation may follow a committed write. Preserve those facts
# without replaying incomplete provider tool calls or claiming generated text.
user = await db.scalar(
select(AIMessage).where(AIMessage.run_id == item.id, AIMessage.role == "user")
)
calls = (
await db.scalars(select(AIToolCall).where(AIToolCall.run_id == item.id))
).all()
if user:
facts = {
"run_status": item.status,
"error": item.error,
"tool_records": [
{"name": c.name, "status": c.status, "result": model_result(c.result)}
for c in calls
],
}
history.extend(
to_jsonable_python(
[
ModelRequest(
parts=[
UserPromptPart(
user.parts[0]["text"]
+ "\n页面上下文:"
+ json.dumps(item.context, ensure_ascii=False)
)
]
),
ModelResponse(
parts=[
TextPart(
"服务端执行记录(非模型回答):"
+ json.dumps(facts, ensure_ascii=False)
)
]
),
]
)
)
run = AIRun(
id=uid(),
conversation_id=conversation_id,
request_id=body.request_id,
context=body.context.model_dump(mode="json"),
model=config.model,
settings_revision=config.revision,
model_messages=history,
history_start=len(history),
)
db.add(run)
await db.flush()
db.add(
AIMessage(
id=uid(),
conversation_id=conversation_id,
run_id=run.id,
role="user",
parts=[{"type": "text", "text": body.message}],
)
)
conversation.updated_at = now()
if conversation.title == "新会话":
conversation.title = body.message[:60]
prompt = (
body.message
+ "\n\n页面上下文(仅数据引用):"
+ json.dumps(run.context, ensure_ascii=False)
)
await self.launch(run.id, token, prompt)
return run.id
async def launch(self, run_id, token, prompt=None):
live = LiveRun(message_id=uid())
async with self.sessions.begin() as db:
run = await db.get(AIRun, run_id)
db.add(
AIMessage(
id=live.message_id,
run_id=run_id,
conversation_id=run.conversation_id,
role="assistant",
parts=[],
)
)
self.live[run_id] = live
await live.emit({"type": "start", "messageId": live.message_id})
await live.emit({"type": "data-run", "data": await self.snapshot(run_id), "transient": True})
live.task = asyncio.create_task(self.execute(run_id, token, prompt, live))
# Ensure the task enters its try/finally before an immediate cancel request can arrive.
await asyncio.sleep(0)
async def save_parts(self, live, force=False):
if not force and time.monotonic() - live.last_save < 0.25:
return
async with self.sessions.begin() as db:
message = await db.get(AIMessage, live.message_id)
message.parts = jsonable_encoder(live.parts)
live.last_save = time.monotonic()
async def card(self, live, call):
data = {
"id": call.id,
"name": call.name,
"status": call.status,
"presentation": presentation(call.name),
"preview": call.preview if call.status == "pending" else {},
"result": call.result,
}
part = {"type": "data-tool", "id": call.id, "data": jsonable_encoder(data)}
live.parts.append(part)
await self.save_parts(live, True)
await live.emit(part)
async def tool(self, run_id, token, live, name, call_id, kwargs):
await self.authorize(token)
capability = CAPABILITIES.get(name)
if capability is None:
raise ModelRetry("此能力不可用,请使用当前提供的工具")
try:
args = capability.schema.model_validate(kwargs)
except ValidationError:
raise ModelRetry("参数不符合工具契约,请检查字段、范围和类型") from None
async with self.sessions.begin() as db:
ai_run = await db.get(AIRun, run_id)
business = Business(db, {"conversation_id": ai_run.conversation_id, "ai_run_id": run_id})
call = AIToolCall(
id=uid(),
run_id=run_id,
call_id=call_id,
name=name,
arguments=args.model_dump(mode="json", exclude_unset=True),
)
try:
# A prepare handler may flush a new artifact before a later validation fails.
# Roll back business changes while retaining a durable failed audit record.
async with db.begin_nested():
context = ToolContext(business, self.runner.client)
if capability.requires_confirmation:
call.preview = jsonable_encoder(await capability.preview(context, args))
call.status = "pending"
else:
call.result = await capability.invoke(context, kwargs)
call.status = "completed"
except HTTPException as exc:
call.result, call.status = {"error": exc.detail}, "failed"
except (ValueError, ValidationError):
call.result, call.status = {"error": "操作参数或目标数据不符合业务规则"}, "failed"
db.add(call)
await self.card(live, call)
if call.status == "pending":
raise CallDeferred(metadata={"approval_id": call.id})
return model_result(call.result)
async def execute(self, run_id, token, prompt, live):
started = time.monotonic()
status, error, usage = "completed", None, RunUsage()
token_usage_known = False
try:
await self.authorize(token)
async with self.sessions() as db:
run = await db.get(AIRun, run_id)
config = await self.config(db)
if config.revision != run.settings_revision:
raise HTTPException(409, "模型配置已变化,请开始新一轮对话")
history = ModelMessagesTypeAdapter.validate_python(run.model_messages)
usage = RunUsage(**{k: v for k, v in run.usage.items() if k in RunUsage.__dataclass_fields__})
token_usage_known = run.usage.get("token_usage_known", True)
old_elapsed = run.elapsed_ms
calls = (await db.scalars(select(AIToolCall).where(AIToolCall.run_id == run_id))).all()
unresolved = {
part.tool_call_id
for message in history
for part in message.parts
if isinstance(part, ToolCallPart)
} - {
part.tool_call_id
for message in history
for part in message.parts
if isinstance(part, (ToolReturnPart, RetryPromptPart))
}
deferred = (
DeferredToolResults(
calls={
c.call_id: model_result(c.result)
for c in calls
if c.call_id in unresolved and c.status != "pending"
}
)
if prompt is None
else None
)
tools = []
for name, capability in CAPABILITIES.items():
def bind(tool_name):
async def handler(ctx, **kwargs):
return await self.tool(run_id, token, live, tool_name, ctx.tool_call_id, kwargs)
return handler
tools.append(
Tool.from_schema(
bind(name),
name,
capability.description,
capability.schema.model_json_schema(),
takes_ctx=True,
sequential=True,
)
)
async with asyncio.timeout(max(0.01, self.settings.ai_timeout - old_elapsed / 1000)):
async with self.model_factory(config, self.settings) as model:
agent = Agent(
model,
tools=tools,
instructions=INSTRUCTIONS,
output_type=[str, DeferredToolRequests],
tool_retries=1,
output_retries=1,
model_settings={
"max_tokens": self.settings.ai_output_tokens,
"parallel_tool_calls": False,
},
)
async with agent.run_stream_events(
prompt,
message_history=history,
deferred_tool_results=deferred,
usage=usage,
usage_limits=UsageLimits(
request_limit=self.settings.ai_request_limit,
tool_calls_limit=self.settings.ai_tool_limit,
),
) as stream:
async for event in stream:
if isinstance(event, PartStartEvent) and isinstance(event.part, TextPart):
if live.text_id:
await live.emit({"type": "text-end", "id": live.text_id})
live.text_id = uid()
live.parts.append({"type": "text", "text": event.part.content})
await live.emit({"type": "text-start", "id": live.text_id})
if event.part.content:
await live.emit(
{
"type": "text-delta",
"id": live.text_id,
"delta": event.part.content,
}
)
elif isinstance(event, PartDeltaEvent) and isinstance(event.delta, TextPartDelta):
live.parts[-1]["text"] += event.delta.content_delta
await live.emit(
{
"type": "text-delta",
"id": live.text_id,
"delta": event.delta.content_delta,
}
)
elif isinstance(event, AgentRunResultEvent):
ensure_complete(model, event.result)
token_usage_known = token_usage_known and all(
bool(m.usage.input_tokens or m.usage.output_tokens)
for m in event.result.new_messages()
if isinstance(m, ModelResponse)
)
status = (
"waiting_approval"
if isinstance(event.result.output, DeferredToolRequests)
else "completed"
)
async with self.sessions.begin() as db:
row = await db.get(AIRun, run_id)
row.model_messages = to_jsonable_python(event.result.all_messages())
await self.save_parts(live)
except asyncio.CancelledError:
status = "interrupted" if self.stopping else "cancelled"
error = "执行已停止;已完成的操作和已创建的业务任务保留"
except HTTPException as exc:
status, error = "failed", str(exc.detail)
except Exception as exc:
from pydantic_ai.exceptions import UsageLimitExceeded
status = "failed"
error = (
"本轮已达到模型或工具调用上限,请缩小请求范围"
if isinstance(exc, UsageLimitExceeded)
else public_error(exc)
)
finally:
await self.save_parts(live, True)
async with self.sessions.begin() as db:
row = await db.get(AIRun, run_id)
row.status, row.error, row.updated_at = status, error, now()
row.usage = {
**asdict(usage),
"token_usage_known": token_usage_known and status in ("completed", "waiting_approval"),
}
row.elapsed_ms += round((time.monotonic() - started) * 1000)
if status != "waiting_approval":
await db.execute(
update(AIToolCall)
.where(AIToolCall.run_id == run_id, AIToolCall.status == "pending")
.values(status="cancelled")
)
if live.text_id:
await live.emit({"type": "text-end", "id": live.text_id})
await live.emit(
{
"type": "data-run",
"data": await self.snapshot(run_id),
"transient": True,
}
)
await live.emit({"type": "finish"})
async with live.condition:
live.done = True
live.condition.notify_all()
if self.live.get(run_id) is live:
del self.live[run_id]
async def decision(self, approval_id, approved, token):
await self.authorize(token)
async with self.lock:
async with self.sessions.begin() as db:
call = await db.scalar(
select(AIToolCall)
.where(AIToolCall.id == approval_id, AIToolCall.admin_id == 1)
.with_for_update()
)
if not call:
raise HTTPException(404, "确认记录不存在")
run = await db.get(AIRun, call.run_id)
if call.status != "pending":
return run.id
if run.status != "waiting_approval":
raise HTTPException(409, "当前执行尚未准备好确认或已经结束")
await self.config(db)
config = await db.get(AISettings, 1)
if config.revision != run.settings_revision:
raise HTTPException(409, "模型配置已变化,请停止本轮并重新提出操作")
if approved:
try:
# Nested transaction rolls back partial bulk mutations but preserves the failed audit.
async with db.begin_nested():
capability = CAPABILITIES.get(call.name)
if capability is None or not capability.requires_confirmation:
raise HTTPException(409, "原操作已不可用,请重新提出请求")
args = capability.schema.model_validate(call.arguments)
result = await capability.execute(
ToolContext(
Business(
db,
{
"conversation_id": run.conversation_id,
"ai_run_id": run.id,
},
),
self.runner.client,
),
args,
call.preview,
)
call.result, call.status = jsonable_encoder(result), "completed"
except HTTPException as exc:
call.result, call.status = {"error": exc.detail}, "failed"
except (ValueError, ValidationError):
call.result, call.status = {"error": "原操作参数已不符合契约,请重新预览"}, "failed"
else:
call.result, call.status = (
{"denied": True, "message": "用户拒绝了此操作,不得重新提出相同操作"},
"denied",
)
await db.flush()
pending = await db.scalar(
select(AIToolCall.id).where(AIToolCall.run_id == run.id, AIToolCall.status == "pending")
)
if not pending:
run.status = "running"
run_id, result, complete = run.id, call.result, call.status == "completed"
if complete and capability.after_commit:
try:
await capability.after_commit(self.runner, result)
except Exception:
# A notification failure cannot undo a committed operation or
# strand its chat. Persist the distinction without replaying it.
logger.warning("AI tool %s committed but runner notification failed", approval_id)
async with self.sessions.begin() as db:
call = await db.get(AIToolCall, approval_id)
call.result = {
**result,
"_warning": "操作已保存,但后台通知失败;请查看任务状态,不要重复执行。",
}
if not pending:
await self.launch(run_id, token)
return run_id
async def cancel(self, run_id):
async with self.lock:
async with self.sessions.begin() as db:
row = await db.get(AIRun, run_id)
if not row:
raise HTTPException(404, "执行不存在")
if row.status == "waiting_approval":
row.status, row.updated_at = "cancelled", now()
await db.execute(
update(AIToolCall)
.where(AIToolCall.run_id == run_id, AIToolCall.status == "pending")
.values(status="cancelled")
)
live = self.live.get(run_id)
if live and live.task:
live.task.cancel()
await asyncio.gather(live.task, return_exceptions=True)
async def snapshot(self, run_id):
async with self.sessions() as db:
row = await db.get(AIRun, run_id)
if not row:
raise HTTPException(404, "执行不存在")
calls = (
await db.scalars(
select(AIToolCall).where(AIToolCall.run_id == run_id).order_by(AIToolCall.created_at)
)
).all()
return jsonable_encoder(
{
"id": row.id,
"conversation_id": row.conversation_id,
"status": row.status,
"error": row.error,
"model": row.model,
"usage": row.usage,
"elapsed_ms": row.elapsed_ms,
"tools": [
{
"id": c.id,
"name": c.name,
"status": c.status,
"presentation": presentation(c.name),
"preview": c.preview,
"result": c.result,
}
for c in calls
],
}
)
async def events(self, run_id):
"""Replay this connection's in-memory stream, or persisted message parts after completion."""
live = self.live.get(run_id)
if live:
index = 0
while True:
async with live.condition:
if index == len(live.events) and not live.done:
try:
await asyncio.wait_for(live.condition.wait(), timeout=10)
except TimeoutError:
pass
events = live.events[index:]
index = len(live.events)
done = live.done
if not events and not done:
yield ": keep-alive\n\n"
for event in events:
yield "data: " + json.dumps(event, ensure_ascii=False) + "\n\n"
if done:
break
else:
async with self.sessions() as db:
message = await db.scalar(
select(AIMessage)
.where(AIMessage.run_id == run_id, AIMessage.role == "assistant")
.order_by(AIMessage.created_at.desc())
.limit(1)
)
events = [{"type": "start", "messageId": message.id if message else uid()}]
for part in message.parts if message else []:
if part["type"] == "text":
part_id = uid()
events.extend(
[
{"type": "text-start", "id": part_id},
{"type": "text-delta", "id": part_id, "delta": part["text"]},
{"type": "text-end", "id": part_id},
]
)
else:
events.append(part)
events.extend(
[
{"type": "data-run", "data": await self.snapshot(run_id), "transient": True},
{"type": "finish"},
]
)
for event in events:
yield "data: " + json.dumps(event, ensure_ascii=False) + "\n\n"
yield "data: [DONE]\n\n"