659 lines
30 KiB
Python
659 lines
30 KiB
Python
"""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"
|