Merge pull request 'feat(agent): 持久化 Agent Trace 并支持 SSE 断点恢复' (#7) from feat/agent-trace-persistence into main

Reviewed-on: #7
This commit is contained in:
2026-09-01 10:36:43 +08:00
15 changed files with 982 additions and 76 deletions
+16 -24
View File
@@ -10,7 +10,7 @@
NotesAgent/ NotesAgent/
├── frontend/ Vue 3 + TypeScript + Vite 前端 ├── frontend/ Vue 3 + TypeScript + Vite 前端
├── backend/ FastAPI + Pydantic 后端 ├── backend/ FastAPI + Pydantic 后端
├── docs/ 分工与技术栈说明 ├── docs/ 架构、契约、开发说明、协作规范与问题复盘
└── server sync/ 云同步服务预留目录,当前未实现 └── server sync/ 云同步服务预留目录,当前未实现
``` ```
@@ -36,7 +36,7 @@ python --version
uv --version uv --version
``` ```
当前 Web 联调不需要 Rust 和 Tauri。开始桌面端集成后,再按照 `docs/AI笔记软件技术栈说明-团队版-v2.3.md` 安装 Rust Toolchain 与 Tauri CLI。 当前 Web 联调不需要 Rust 和 Tauri。开始桌面端集成后,再按照 `docs/architecture/AI笔记软件技术栈说明-团队版-v2.3.md` 安装 Rust Toolchain 与 Tauri CLI。
## 首次初始化 ## 首次初始化
@@ -118,7 +118,7 @@ cd frontend
pnpm test pnpm test
``` ```
当前回归基线为后端 76 项测试、前端 26 项测试,且生产构建通过。测试数量会随功能增长,以本地实际输出和 CI 为准。 当前回归基线为后端 81 项测试、前端 27 项测试,且生产构建通过。测试数量会随功能增长,以本地实际输出和 CI 为准。
构建产物位于 `frontend/dist`,该目录不提交到 Git。 构建产物位于 `frontend/dist`,该目录不提交到 Git。
@@ -126,24 +126,15 @@ pnpm test
| 文档 | 用途 | | 文档 | 用途 |
| --- | --- | | --- | --- |
| [技术栈说明](docs/AI笔记软件技术栈说明-团队版-v2.3.md) | 目标架构、第二阶段技术边界与模块依赖 | | [文档总索引](docs/README.md) | 文档分类、阅读顺序和维护规则 |
| [第一阶段分工表](docs/第一阶段分工表.md) | 成员职责、协作关系与当前交付状态 | | [技术栈说明](docs/architecture/AI笔记软件技术栈说明-团队版-v2.3.md) | 目标架构、第二阶段技术边界与模块依赖 |
| [第二阶段分工表](docs/第二阶段团队分工表.md) | 第二阶段人员职责、任务顺序、协作关系与验收项 | | [第二阶段分工表](docs/architecture/第二阶段团队分工表.md) | 第二阶段人员职责、任务顺序、协作关系与验收项 |
| [第一阶段测试验证操作手册](docs/第一阶段测试验证操作手册.md) | 自动化测试、接口主链路、前端人工验收与记录模板 | | [后端接口契约](docs/contracts/后端接口契约-开发版.md) | HTTP/SSE 接口、错误和当前实现状态 |
| [后端接口契约](docs/后端接口契约-开发版.md) | HTTP/SSE 接口、错误和当前实现状态 | | [第二阶段接口契约](docs/contracts/第二阶段接口契约-开发版.md) | 第二阶段公共 DTO、计划接口、SSE、错误码与联调顺序 |
| [第二阶段接口契约](docs/第二阶段接口契约-开发版.md) | 第二阶段公共 DTO、计划接口、SSE、错误码与联调顺序 | | [AI Core 与 Agent Core](docs/development/AI-Core与Agent-Core开发说明.md) | Provider、Agent、Tool、Permission 与 Extension Core |
| [AI Core 与 Agent Core](docs/AI-Core与Agent-Core开发说明.md) | Provider、Agent、Tool、Permission 与 Extension Core | | [Git 使用细则](docs/guides/Git使用细则-团队开发版.md) | 分支、提交、PR、Review 与合并流程 |
| [Knowledge 与 Retrieval Core](docs/Knowledge与Retrieval-Core开发说明.md) | Block、索引、混合检索和 Citation | | [CI/CD 细则](docs/guides/CI-CD细则-团队开发版.md) | Gitea 流水线、质量门禁、产物、发布与回滚规则 |
| [模型提供商与模型发现](docs/模型提供商与模型发现开发说明.md) | Provider 预设、模型发现和凭据边界 | | [Agent Trace 复盘](docs/retrospectives/Agent-Core第二阶段问题与修复复盘.md) | Agent 持久化、SSE 恢复、事件契约与脱敏问题复盘 |
| [前端页面需求](docs/前端页面需求说明-开发版.md) | 页面、交互、状态与验收基线 |
| [前端实现说明](docs/前端壳子与接口层开发说明.md) | 当前前端目录、Service、SSE 和运行边界 |
| [前端写作体验](docs/前端写作体验优化开发说明.md) | Milkdown、CodeMirror、格式栏和 Shiki |
| [前端视觉与轻量动效](docs/前端视觉与轻量动效优化开发说明.md) | Design Token、页面美化、性能边界与主题注入约定 |
| [Git 使用细则](docs/Git使用细则-团队开发版.md) | 分支、提交、PR、Review 与合并流程 |
| [代码注释与 TODO 约定](docs/代码注释与TODO约定.md) | 注释原则、TODO 格式、领域标签与当前待办索引 |
| [后端审阅复盘](docs/后端全面审阅问题与修复复盘.md) | 后端问题原因、后果与修复方案 |
| [Knowledge/Retrieval 复盘](docs/Knowledge与Retrieval-Core问题与修复复盘.md) | 检索与事务问题复盘 |
| [前端审阅复盘](docs/前端合并审阅问题与修复复盘.md) | 前端工程、契约和交互问题复盘 |
## 日常开发注意事项 ## 日常开发注意事项
@@ -153,6 +144,7 @@ pnpm test
- API 默认监听 `127.0.0.1:8000`,前端默认监听 `127.0.0.1:5173` - API 默认监听 `127.0.0.1:8000`,前端默认监听 `127.0.0.1:5173`
- 后端附件目录默认是 `backend/data/attachments`,可通过 `APP_ATTACHMENTS_PATH` 覆盖;该目录由桌面 Host 管理。 - 后端附件目录默认是 `backend/data/attachments`,可通过 `APP_ATTACHMENTS_PATH` 覆盖;该目录由桌面 Host 管理。
- 跨模块接口发生变化时,需要同步更新前后端类型和 `docs` 中的接口说明。 - 跨模块接口发生变化时,需要同步更新前后端类型和 `docs` 中的接口说明。
- 当前已实现接口见 `docs/后端接口契约-开发版.md`,第二阶段规划接口见 `docs/第二阶段接口契约-开发版.md`;已实现能力以 `/openapi.json` 为准。 - 当前已实现接口见 `docs/contracts/后端接口契约-开发版.md`,第二阶段规划接口见 `docs/contracts/第二阶段接口契约-开发版.md`;已实现能力以 `/openapi.json` 为准。
- 前端页面、交互、状态管理和第一阶段验收要求见 `docs/前端页面需求说明-开发版.md` - 前端页面、交互、状态管理和第一阶段验收要求见 `docs/contracts/前端页面需求说明-开发版.md`
- 分支、提交、Pull Request、Review 和冲突处理规范见 `docs/Git使用细则-团队开发版.md` - 分支、提交、Pull Request、Review 和冲突处理规范见 `docs/guides/Git使用细则-团队开发版.md`
- CI 检查、产物、发布和回滚规范见 `docs/guides/CI-CD细则-团队开发版.md`
+4 -4
View File
@@ -23,10 +23,10 @@ uv run uvicorn app.main:app --reload --host 127.0.0.1 --port 8000
uv run pytest uv run pytest
``` ```
当前基线为 71 项测试通过。Provider API Key 可通过前端设置页写入,也可用 `OPENAI_API_KEY``DEEPSEEK_API_KEY``AINOTE_CREDENTIAL_<ID>` 注入;不要把真实密钥写入仓库。 当前基线为 81 项测试通过。Provider API Key 可通过前端设置页写入,也可用 `OPENAI_API_KEY``DEEPSEEK_API_KEY``AINOTE_CREDENTIAL_<ID>` 注入;不要把真实密钥写入仓库。
团队接口清单见 `../docs/后端接口契约-开发版.md`,机器可读契约以运行时的 `/openapi.json` 为准。 团队接口清单见 `../docs/contracts/后端接口契约-开发版.md`,机器可读契约以运行时的 `/openapi.json` 为准。
AI Core 与 Agent Core 的模块边界、Mock Provider 和 Tool Calling 调试方式见 `../docs/AI-Core与Agent-Core开发说明.md` AI Core 与 Agent Core 的模块边界、Mock Provider 和 Tool Calling 调试方式见 `../docs/development/AI-Core与Agent-Core开发说明.md`
Knowledge Core 与 Retrieval Core 的模块边界、数据模型、接口与检索流程见 `../docs/Knowledge与Retrieval-Core开发说明.md` Knowledge Core 与 Retrieval Core 的模块边界、数据模型、接口与检索流程见 `../docs/development/Knowledge与Retrieval-Core开发说明.md`
+5
View File
@@ -104,6 +104,11 @@ class PermissionManager:
ticket.future.set_result(decision) ticket.future.set_result(decision)
return True return True
def get_ticket(self, run_id: str, request_id: str) -> PermissionTicket | None:
"""只读返回待确认票据,供 Trace 记录权限类型;不暴露 Future 给接口层。"""
return self._pending.get((run_id, request_id))
def cancel_run(self, run_id: str) -> None: def cancel_run(self, run_id: str) -> None:
for key, ticket in list(self._pending.items()): for key, ticket in list(self._pending.items()):
if ticket.run_id == run_id: if ticket.run_id == run_id:
+170 -27
View File
@@ -7,17 +7,20 @@ import json
from collections.abc import AsyncIterator from collections.abc import AsyncIterator
from dataclasses import dataclass, field from dataclasses import dataclass, field
from datetime import datetime, timezone from datetime import datetime, timezone
from time import perf_counter
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from uuid import uuid4 from uuid import uuid4
from app.agent.permissions import PermissionManager, PermissionMode from app.agent.permissions import PermissionManager, PermissionMode
from app.agent.tools import ToolExecutionContext, ToolNotFoundError, ToolRegistry from app.agent.tools import ToolExecutionContext, ToolNotFoundError, ToolRegistry
from app.agent.trace_repository import AgentTraceRepository, sanitize_trace_value
from app.contracts import ( from app.contracts import (
AgentEvent, AgentEvent,
AgentEventType, AgentEventType,
AgentRun, AgentRun,
AgentRunCreateRequest, AgentRunCreateRequest,
AgentRunStatus, AgentRunStatus,
AgentTraceResponse,
Citation, Citation,
Message, Message,
MessageRole, MessageRole,
@@ -61,6 +64,7 @@ class RunRecord:
events: list[AgentEvent] = field(default_factory=list) events: list[AgentEvent] = field(default_factory=list)
subscribers: set[asyncio.Queue[AgentEvent]] = field(default_factory=set) subscribers: set[asyncio.Queue[AgentEvent]] = field(default_factory=set)
task: asyncio.Task[None] | None = None task: asyncio.Task[None] | None = None
next_sequence: int = 0
class AgentRuntime: class AgentRuntime:
@@ -72,11 +76,13 @@ class AgentRuntime:
tools: ToolRegistry, tools: ToolRegistry,
permissions: PermissionManager, permissions: PermissionManager,
skills: SkillRuntime | None = None, skills: SkillRuntime | None = None,
trace_repository: AgentTraceRepository | None = None,
) -> None: ) -> None:
self.providers = providers self.providers = providers
self.tools = tools self.tools = tools
self.permissions = permissions self.permissions = permissions
self.skills = skills self.skills = skills
self.trace_repository = trace_repository or AgentTraceRepository()
self._records: dict[str, RunRecord] = {} self._records: dict[str, RunRecord] = {}
async def create_run(self, request: AgentRunCreateRequest) -> AgentRun: async def create_run(self, request: AgentRunCreateRequest) -> AgentRun:
@@ -116,22 +122,38 @@ class AgentRuntime:
skill_config=skill_config, skill_config=skill_config,
allowed_tools=allowed_tools, allowed_tools=allowed_tools,
) )
self.trace_repository.create_run(
run,
request,
self._config_snapshot(record),
)
self._records[run.run_id] = record self._records[run.run_id] = record
record.task = asyncio.create_task(self._execute(record), name=run.run_id) record.task = asyncio.create_task(self._execute(record), name=run.run_id)
return run.model_copy(deep=True) return run.model_copy(deep=True)
def get_run(self, run_id: str) -> AgentRun: def get_run(self, run_id: str) -> AgentRun:
return self._get_record(run_id).run.model_copy(deep=True) record = self._records.get(run_id)
if record is not None:
return record.run.model_copy(deep=True)
run = self.trace_repository.recover_interrupted(run_id)
if run is None:
raise AgentRunNotFoundError(run_id)
return run.model_copy(deep=True)
def list_runs(self, limit: int, offset: int) -> tuple[list[AgentRun], int]: def list_runs(self, limit: int, offset: int) -> tuple[list[AgentRun], int]:
records = sorted( items, total = self.trace_repository.list_runs(limit=limit, offset=offset)
self._records.values(), key=lambda item: item.run.created_at, reverse=True recovered = [
) self.trace_repository.recover_interrupted(item.run_id) or item
items = [item.run.model_copy(deep=True) for item in records[offset : offset + limit]] if item.run_id not in self._records
return items, len(records) else self._records[item.run_id].run.model_copy(deep=True)
for item in items
]
return recovered, total
async def cancel(self, run_id: str) -> AgentRun: async def cancel(self, run_id: str) -> AgentRun:
record = self._get_record(run_id) record = self._records.get(run_id)
if record is None:
return self.get_run(run_id)
if record.run.status in TERMINAL_STATUSES: if record.run.status in TERMINAL_STATUSES:
return record.run.model_copy(deep=True) return record.run.model_copy(deep=True)
record.run.cancelled = True record.run.cancelled = True
@@ -144,23 +166,53 @@ class AgentRuntime:
return record.run.model_copy(deep=True) return record.run.model_copy(deep=True)
def resolve_permission(self, run_id: str, request_id: str, decision: str) -> bool: def resolve_permission(self, run_id: str, request_id: str, decision: str) -> bool:
self._get_record(run_id) record = self._records.get(run_id)
return self.permissions.resolve(run_id, request_id, decision) if record is None:
return False
ticket = self.permissions.get_ticket(run_id, request_id)
resolved = self.permissions.resolve(run_id, request_id, decision)
if resolved:
self._publish(
record,
AgentEventType.permission_resolved,
{
"request_id": request_id,
"permission": ticket.permission if ticket else None,
"decision": decision,
},
)
return resolved
async def events(self, run_id: str) -> AsyncIterator[AgentEvent]: async def events(
record = self._get_record(run_id) self, run_id: str, *, after_sequence: int = -1
# 先回放快照再订阅实时事件,使晚加入的 SSE 客户端也能恢复界面状态。 ) -> AsyncIterator[AgentEvent]:
# TODO(agent): 持久化事件并支持 Last-Event-ID,进程重启后仍可续传。 record = self._records.get(run_id)
run = self.get_run(run_id)
if record is None:
for event in self.trace_repository.list_events(
run_id, after_sequence=after_sequence
):
yield event
return
# 先注册订阅再读持久化历史;同一事件循环内没有 await,不会丢失交界事件。
queue: asyncio.Queue[AgentEvent] = asyncio.Queue() queue: asyncio.Queue[AgentEvent] = asyncio.Queue()
record.subscribers.add(queue) record.subscribers.add(queue)
history = [event.model_copy(deep=True) for event in record.events] history = self.trace_repository.list_events(
run_id, after_sequence=after_sequence
)
last_sequence = after_sequence
try: try:
for event in history: for event in history:
last_sequence = event.sequence
yield event yield event
if record.run.status in TERMINAL_STATUSES: if run.status in TERMINAL_STATUSES:
return return
while True: while True:
event = await queue.get() event = await queue.get()
if event.sequence <= last_sequence:
continue
last_sequence = event.sequence
yield event.model_copy(deep=True) yield event.model_copy(deep=True)
if event.event in { if event.event in {
AgentEventType.run_completed, AgentEventType.run_completed,
@@ -172,7 +224,9 @@ class AgentRuntime:
record.subscribers.discard(queue) record.subscribers.discard(queue)
async def wait(self, run_id: str) -> AgentRun: async def wait(self, run_id: str) -> AgentRun:
record = self._get_record(run_id) record = self._records.get(run_id)
if record is None:
return self.get_run(run_id)
if record.task: if record.task:
try: try:
await asyncio.shield(record.task) await asyncio.shield(record.task)
@@ -180,6 +234,17 @@ class AgentRuntime:
pass pass
return record.run.model_copy(deep=True) return record.run.model_copy(deep=True)
def get_trace(
self, run_id: str, *, after_sequence: int, limit: int
) -> AgentTraceResponse:
self.get_run(run_id)
trace = self.trace_repository.get_trace(
run_id, after_sequence=after_sequence, limit=limit
)
if trace is None:
raise AgentRunNotFoundError(run_id)
return trace
async def _execute(self, record: RunRecord) -> None: async def _execute(self, record: RunRecord) -> None:
try: try:
async with asyncio.timeout(record.request.run_timeout_seconds): async with asyncio.timeout(record.request.run_timeout_seconds):
@@ -210,6 +275,19 @@ class AgentRuntime:
for step in range(1, record.request.max_steps + 1): for step in range(1, record.request.max_steps + 1):
record.run.current_step = step record.run.current_step = step
record.run.updated_at = datetime.now(timezone.utc) record.run.updated_at = datetime.now(timezone.utc)
model_call_id = f"model_call_{uuid4().hex}"
started_at = perf_counter()
self._publish(
record,
AgentEventType.model_call_started,
{
"model_call_id": model_call_id,
"step": step,
"provider_id": record.request.provider_id,
"model": record.request.model,
},
)
try:
turn = await provider.complete( turn = await provider.complete(
ModelRequest( ModelRequest(
provider_id=record.request.provider_id, provider_id=record.request.provider_id,
@@ -220,6 +298,29 @@ class AgentRuntime:
metadata=self._request_metadata(record), metadata=self._request_metadata(record),
) )
) )
except Exception as exc:
self._publish(
record,
AgentEventType.model_call_failed,
{
"model_call_id": model_call_id,
"duration_ms": int((perf_counter() - started_at) * 1000),
"error_code": getattr(exc, "code", type(exc).__name__),
},
)
raise
self._publish(
record,
AgentEventType.model_call_completed,
{
"model_call_id": model_call_id,
"duration_ms": int((perf_counter() - started_at) * 1000),
"finish_reason": "tool_calls" if turn.tool_calls else "stop",
"input_tokens": turn.input_tokens,
"output_tokens": turn.output_tokens,
"tool_call_count": len(turn.tool_calls),
},
)
record.run.token_usage += turn.input_tokens + turn.output_tokens record.run.token_usage += turn.input_tokens + turn.output_tokens
self._publish( self._publish(
record, record,
@@ -257,7 +358,7 @@ class AgentRuntime:
async def execute(call: ToolCall) -> ToolResult: async def execute(call: ToolCall) -> ToolResult:
async with semaphore: async with semaphore:
return await self._execute_tool(record, call) return await self._execute_tool(record, call, model_call_id)
results = await asyncio.gather(*(execute(call) for call in calls)) results = await asyncio.gather(*(execute(call) for call in calls))
for call, result in zip(calls, results): for call, result in zip(calls, results):
@@ -290,8 +391,13 @@ class AgentRuntime:
self._fail(record, "MAX_STEPS_EXCEEDED", "Agent reached its maximum step count.") self._fail(record, "MAX_STEPS_EXCEEDED", "Agent reached its maximum step count.")
async def _execute_tool(self, record: RunRecord, call: ToolCall) -> ToolResult: async def _execute_tool(
self._publish(record, AgentEventType.tool_call, call.model_dump(mode="json")) self, record: RunRecord, call: ToolCall, parent_model_call_id: str
) -> ToolResult:
started_at = perf_counter()
call_data = call.model_dump(mode="json")
call_data["parent_model_call_id"] = parent_model_call_id
self._publish(record, AgentEventType.tool_call, call_data)
try: try:
registered = self.tools.get(call.name) registered = self.tools.get(call.name)
except ToolNotFoundError: except ToolNotFoundError:
@@ -305,7 +411,9 @@ class AgentRuntime:
error_code="TOOL_NOT_ALLOWED", error_code="TOOL_NOT_ALLOWED",
error_message="Tool is not included in allowed_tools.", error_message="Tool is not included in allowed_tools.",
) )
self._publish(record, AgentEventType.tool_result, result.model_dump(mode="json")) self._publish_tool_result(
record, result, parent_model_call_id, started_at
)
return result return result
permission = registered.definition.permission if registered else None permission = registered.definition.permission if registered else None
@@ -317,7 +425,9 @@ class AgentRuntime:
error_code="NETWORK_NOT_ALLOWED", error_code="NETWORK_NOT_ALLOWED",
error_message="Agent run does not allow network tools.", error_message="Agent run does not allow network tools.",
) )
self._publish(record, AgentEventType.tool_result, result.model_dump(mode="json")) self._publish_tool_result(
record, result, parent_model_call_id, started_at
)
return result return result
mode = self.permissions.mode_for(permission) mode = self.permissions.mode_for(permission)
if mode == PermissionMode.deny: if mode == PermissionMode.deny:
@@ -348,11 +458,13 @@ class AgentRuntime:
error_code="PERMISSION_TIMEOUT", error_code="PERMISSION_TIMEOUT",
error_message="Tool permission confirmation timed out.", error_message="Tool permission confirmation timed out.",
) )
self._publish( self._publish_tool_result(
record, AgentEventType.tool_result, result.model_dump(mode="json") record, result, parent_model_call_id, started_at
) )
return result return result
record.run.status = AgentRunStatus.running record.run.status = AgentRunStatus.running
record.run.updated_at = datetime.now(timezone.utc)
self.trace_repository.save_run(record.run)
result = ( result = (
await self._invoke_tool(record, call) await self._invoke_tool(record, call)
if decision in {"allow_once", "allow_session"} if decision in {"allow_once", "allow_session"}
@@ -361,9 +473,21 @@ class AgentRuntime:
else: else:
result = await self._invoke_tool(record, call) result = await self._invoke_tool(record, call)
self._publish(record, AgentEventType.tool_result, result.model_dump(mode="json")) self._publish_tool_result(record, result, parent_model_call_id, started_at)
return result return result
def _publish_tool_result(
self,
record: RunRecord,
result: ToolResult,
parent_model_call_id: str,
started_at: float,
) -> None:
data = result.model_dump(mode="json")
data["parent_model_call_id"] = parent_model_call_id
data["duration_ms"] = int((perf_counter() - started_at) * 1000)
self._publish(record, AgentEventType.tool_result, data)
async def _invoke_tool(self, record: RunRecord, call: ToolCall) -> ToolResult: async def _invoke_tool(self, record: RunRecord, call: ToolCall) -> ToolResult:
try: try:
return await asyncio.wait_for( return await asyncio.wait_for(
@@ -411,15 +535,19 @@ class AgentRuntime:
def _publish( def _publish(
self, record: RunRecord, event_type: AgentEventType, data: dict[str, object] self, record: RunRecord, event_type: AgentEventType, data: dict[str, object]
) -> None: ) -> None:
sanitized = sanitize_trace_value(data)
assert isinstance(sanitized, dict)
event = AgentEvent( event = AgentEvent(
event=event_type, event=event_type,
run_id=record.run.run_id, run_id=record.run.run_id,
sequence=len(record.events), sequence=record.next_sequence,
data=data, data=sanitized,
timestamp=datetime.now(timezone.utc), timestamp=datetime.now(timezone.utc),
) )
record.next_sequence += 1
record.events.append(event) record.events.append(event)
# 内存事件只保留最近窗口;完整审计轨迹应由后续持久化层承担。 self.trace_repository.append_event(record.run, event)
# 内存只保留实时订阅窗口;完整审计轨迹由 SQLite 保存。
if len(record.events) > MAX_EVENTS_PER_RUN: if len(record.events) > MAX_EVENTS_PER_RUN:
del record.events[: len(record.events) - MAX_EVENTS_PER_RUN] del record.events[: len(record.events) - MAX_EVENTS_PER_RUN]
for queue in record.subscribers: for queue in record.subscribers:
@@ -433,6 +561,21 @@ class AgentRuntime:
metadata["retrieval"] = record.skill_config.retrieval.model_dump(mode="json") metadata["retrieval"] = record.skill_config.retrieval.model_dump(mode="json")
return metadata return metadata
def _config_snapshot(self, record: RunRecord) -> dict[str, object]:
provider = self.providers.get(record.request.provider_id).config
return {
"provider_id": record.request.provider_id,
"provider_type": provider.provider_type.value,
"model": record.request.model,
"capabilities": [item.value for item in provider.capabilities],
"skill_id": record.request.skill_id,
"allowed_tools": list(record.allowed_tools),
"max_steps": record.request.max_steps,
"token_budget": record.request.token_budget,
"allow_network": record.request.allow_network,
"metadata": record.request.metadata,
}
def _collect_citations(self, record: RunRecord, result: ToolResult) -> None: def _collect_citations(self, record: RunRecord, result: ToolResult) -> None:
if not result.success or not isinstance(result.output, dict): if not result.success or not isinstance(result.output, dict):
return return
+372
View File
@@ -0,0 +1,372 @@
"""Agent Run/Event 持久化与 Trace 查询。
SQLite 中的事件是 SSE、前端 Trace 和 Benchmark 的共同事实来源。写入前统一脱敏和
限长,避免 Secret 或无限大的 Tool Result 进入审计数据。
"""
from __future__ import annotations
import json
import re
from datetime import datetime, timezone
from typing import Any
from app.contracts import (
AgentEvent,
AgentEventType,
AgentRun,
AgentRunCreateRequest,
AgentRunStatus,
AgentTraceResponse,
AgentTraceSummary,
)
from app.database.db import connect, transaction
MAX_TRACE_STRING = 4_096
MAX_TRACE_COLLECTION = 100
MAX_TRACE_DEPTH = 8
_SECRET_KEYS = {
"api_key",
"apikey",
"authorization",
"access_token",
"refresh_token",
"client_secret",
"password",
"secret",
"token",
}
_SECRET_KEY_SUFFIXES = ("_api_key", "_password", "_secret")
_TERMINAL_VALUES = {
AgentRunStatus.completed.value,
AgentRunStatus.failed.value,
AgentRunStatus.cancelled.value,
}
_BEARER_PATTERN = re.compile(r"(?i)\bBearer\s+[^\s,;]+")
_API_KEY_PATTERN = re.compile(r"\bsk-[A-Za-z0-9_-]{8,}\b")
def sanitize_trace_value(
value: Any, *, depth: int = 0, apply_limits: bool = True
) -> Any:
"""递归净化持久化数据;可按审计用途限制体积,Secret 始终脱敏。"""
if apply_limits and depth >= MAX_TRACE_DEPTH:
return "[MAX_DEPTH]"
if isinstance(value, dict):
sanitized: dict[str, Any] = {}
for index, (key, item) in enumerate(value.items()):
if apply_limits and index >= MAX_TRACE_COLLECTION:
sanitized["__truncated__"] = True
break
normalized = str(key).casefold().replace("-", "_")
sanitized[str(key)] = (
"[REDACTED]"
if normalized in _SECRET_KEYS
or normalized.endswith(_SECRET_KEY_SUFFIXES)
else sanitize_trace_value(
item, depth=depth + 1, apply_limits=apply_limits
)
)
return sanitized
if isinstance(value, (list, tuple)):
source_items = value[:MAX_TRACE_COLLECTION] if apply_limits else value
items = [
sanitize_trace_value(
item, depth=depth + 1, apply_limits=apply_limits
)
for item in source_items
]
if apply_limits and len(value) > MAX_TRACE_COLLECTION:
items.append("[TRUNCATED]")
return items
if isinstance(value, str):
value = _BEARER_PATTERN.sub("Bearer [REDACTED]", value)
value = _API_KEY_PATTERN.sub("[REDACTED]", value)
if apply_limits and len(value) > MAX_TRACE_STRING:
return f"{value[:MAX_TRACE_STRING]}...[TRUNCATED]"
return value
if value is None or isinstance(value, (str, int, float, bool)):
return value
return sanitize_trace_value(
str(value), depth=depth + 1, apply_limits=apply_limits
)
class AgentTraceRepository:
def create_run(
self,
run: AgentRun,
request: AgentRunCreateRequest,
config_snapshot: dict[str, Any],
) -> None:
conn = connect()
try:
with transaction(conn):
conn.execute(
"""
INSERT INTO agent_runs(
run_id, status, run_json, request_json, config_snapshot_json,
created_at, updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?)
""",
(
run.run_id,
run.status.value,
self._serialize_run(run),
json.dumps(
sanitize_trace_value(request.model_dump(mode="json")),
ensure_ascii=False,
),
json.dumps(
sanitize_trace_value(config_snapshot), ensure_ascii=False
),
run.created_at.isoformat(),
run.updated_at.isoformat(),
),
)
finally:
conn.close()
def save_run(self, run: AgentRun) -> None:
conn = connect()
try:
with transaction(conn):
self._update_run(conn, run)
finally:
conn.close()
def append_event(self, run: AgentRun, event: AgentEvent) -> None:
"""在同一事务中保存最新 Run 和事件;复写同一序号时保持幂等。"""
conn = connect()
try:
with transaction(conn):
self._update_run(conn, run)
conn.execute(
"""
INSERT INTO agent_events(run_id, sequence, event, data_json, timestamp)
VALUES (?, ?, ?, ?, ?)
ON CONFLICT(run_id, sequence) DO NOTHING
""",
(
event.run_id,
event.sequence,
event.event.value,
json.dumps(event.data, ensure_ascii=False),
event.timestamp.isoformat(),
),
)
finally:
conn.close()
def get_run(self, run_id: str) -> AgentRun | None:
conn = connect()
try:
row = conn.execute(
"SELECT run_json FROM agent_runs WHERE run_id = ?", (run_id,)
).fetchone()
return AgentRun.model_validate_json(row["run_json"]) if row else None
finally:
conn.close()
def list_runs(self, limit: int, offset: int) -> tuple[list[AgentRun], int]:
conn = connect()
try:
total = int(conn.execute("SELECT COUNT(*) FROM agent_runs").fetchone()[0])
rows = conn.execute(
"""
SELECT run_json FROM agent_runs
ORDER BY created_at DESC LIMIT ? OFFSET ?
""",
(limit, offset),
).fetchall()
return [AgentRun.model_validate_json(row["run_json"]) for row in rows], total
finally:
conn.close()
def list_events(
self, run_id: str, *, after_sequence: int = -1, limit: int | None = None
) -> list[AgentEvent]:
conn = connect()
try:
sql = """
SELECT event, sequence, data_json, timestamp
FROM agent_events
WHERE run_id = ? AND sequence > ?
ORDER BY sequence
"""
params: tuple[Any, ...] = (run_id, after_sequence)
if limit is not None:
sql += " LIMIT ?"
params += (limit,)
return [self._event_from_row(run_id, row) for row in conn.execute(sql, params)]
finally:
conn.close()
def get_trace(
self, run_id: str, *, after_sequence: int, limit: int
) -> AgentTraceResponse | None:
conn = connect()
try:
row = conn.execute(
"""
SELECT run_json, config_snapshot_json
FROM agent_runs WHERE run_id = ?
""",
(run_id,),
).fetchone()
if row is None:
return None
run = AgentRun.model_validate_json(row["run_json"])
event_rows = conn.execute(
"""
SELECT event, sequence, data_json, timestamp
FROM agent_events
WHERE run_id = ? AND sequence > ?
ORDER BY sequence LIMIT ?
""",
(run_id, after_sequence, limit + 1),
).fetchall()
has_more = len(event_rows) > limit
items = [
self._event_from_row(run_id, item) for item in event_rows[:limit]
]
counts = {
item["event"]: int(item["count"])
for item in conn.execute(
"""
SELECT event, COUNT(*) AS count
FROM agent_events WHERE run_id = ? GROUP BY event
""",
(run_id,),
)
}
tool_errors = int(
conn.execute(
"""
SELECT COUNT(*) FROM agent_events
WHERE run_id = ? AND event = 'ToolResult'
AND json_extract(data_json, '$.success') = 0
""",
(run_id,),
).fetchone()[0]
)
errors = (
counts.get(AgentEventType.run_failed.value, 0)
+ counts.get(AgentEventType.model_call_failed.value, 0)
+ tool_errors
)
duration_ms = max(
0, int((run.updated_at - run.created_at).total_seconds() * 1000)
)
return AgentTraceResponse(
run_id=run_id,
status=run.status,
items=items,
next_sequence=items[-1].sequence if items else after_sequence,
has_more=has_more,
summary=AgentTraceSummary(
model_calls=counts.get(AgentEventType.model_call_started.value, 0),
tool_calls=counts.get(AgentEventType.tool_call.value, 0),
duration_ms=duration_ms,
token_usage=run.token_usage,
errors=errors,
),
config_snapshot=json.loads(row["config_snapshot_json"]),
)
finally:
conn.close()
def recover_interrupted(self, run_id: str) -> AgentRun | None:
"""把上个进程遗留的非终态 Run 收束为失败,并追加可回放终止事件。"""
conn = connect()
try:
with transaction(conn):
row = conn.execute(
"SELECT run_json, status FROM agent_runs WHERE run_id = ?", (run_id,)
).fetchone()
if row is None:
return None
run = AgentRun.model_validate_json(row["run_json"])
if row["status"] in _TERMINAL_VALUES:
return run
run.status = AgentRunStatus.failed
run.error_code = "AGENT_PROCESS_RESTARTED"
run.error_message = "Agent process restarted before the run completed."
run.updated_at = datetime.now(timezone.utc)
next_sequence = int(
conn.execute(
"""
SELECT COALESCE(MAX(sequence), -1) + 1
FROM agent_events WHERE run_id = ?
""",
(run_id,),
).fetchone()[0]
)
event = AgentEvent(
event=AgentEventType.run_failed,
run_id=run_id,
sequence=next_sequence,
data={
"code": run.error_code,
"message": run.error_message,
},
timestamp=run.updated_at,
)
self._update_run(conn, run)
conn.execute(
"""
INSERT INTO agent_events(run_id, sequence, event, data_json, timestamp)
VALUES (?, ?, ?, ?, ?)
""",
(
run_id,
next_sequence,
event.event.value,
json.dumps(event.data, ensure_ascii=False),
event.timestamp.isoformat(),
),
)
return run
finally:
conn.close()
@staticmethod
def _update_run(conn, run: AgentRun) -> None:
cursor = conn.execute(
"""
UPDATE agent_runs
SET status = ?, run_json = ?, updated_at = ?
WHERE run_id = ?
""",
(
run.status.value,
AgentTraceRepository._serialize_run(run),
run.updated_at.isoformat(),
run.run_id,
),
)
if cursor.rowcount != 1:
raise LookupError(run.run_id)
@staticmethod
def _event_from_row(run_id: str, row) -> AgentEvent:
return AgentEvent(
event=AgentEventType(row["event"]),
run_id=run_id,
sequence=int(row["sequence"]),
data=json.loads(row["data_json"]),
timestamp=datetime.fromisoformat(row["timestamp"]),
)
@staticmethod
def _serialize_run(run: AgentRun) -> str:
# Run 是重启后 GET/list 的完整事实;只做 Secret 脱敏,不套用 Trace 摘要限长。
return json.dumps(
sanitize_trace_value(
run.model_dump(mode="json"), apply_limits=False
),
ensure_ascii=False,
)
+22
View File
@@ -329,6 +329,10 @@ class AgentEventType(str, Enum):
permission_required = "PermissionRequired" permission_required = "PermissionRequired"
usage = "Usage" usage = "Usage"
citation = "Citation" citation = "Citation"
model_call_started = "ModelCallStarted"
model_call_completed = "ModelCallCompleted"
model_call_failed = "ModelCallFailed"
permission_resolved = "PermissionResolved"
run_completed = "RunCompleted" run_completed = "RunCompleted"
run_failed = "RunFailed" run_failed = "RunFailed"
run_cancelled = "RunCancelled" run_cancelled = "RunCancelled"
@@ -342,6 +346,24 @@ class AgentEvent(Contract):
timestamp: datetime timestamp: datetime
class AgentTraceSummary(Contract):
model_calls: int = 0
tool_calls: int = 0
duration_ms: int = 0
token_usage: int = 0
errors: int = 0
class AgentTraceResponse(Contract):
run_id: str
status: AgentRunStatus
items: list[AgentEvent] = Field(default_factory=list)
next_sequence: int
has_more: bool = False
summary: AgentTraceSummary = Field(default_factory=AgentTraceSummary)
config_snapshot: dict[str, Any] = Field(default_factory=dict)
class PermissionDecisionRequest(Contract): class PermissionDecisionRequest(Contract):
decision: Literal["allow_once", "allow_session", "deny"] decision: Literal["allow_once", "allow_session", "deny"]
+27
View File
@@ -69,6 +69,33 @@ MIGRATIONS: list[str] = [
); );
CREATE INDEX IF NOT EXISTS idx_tasks_status_due ON tasks(status, due_at); CREATE INDEX IF NOT EXISTS idx_tasks_status_due ON tasks(status, due_at);
""", """,
# v3: 第二阶段 Agent TraceRun 与事件事实持久化,供 SSE 恢复和 Benchmark 复用。
"""
CREATE TABLE IF NOT EXISTS agent_runs (
run_id TEXT PRIMARY KEY,
status TEXT NOT NULL,
run_json TEXT NOT NULL,
request_json TEXT NOT NULL,
config_snapshot_json TEXT NOT NULL DEFAULT '{}',
created_at TEXT NOT NULL,
updated_at TEXT NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_agent_runs_created
ON agent_runs(created_at DESC);
CREATE INDEX IF NOT EXISTS idx_agent_runs_status
ON agent_runs(status, updated_at DESC);
CREATE TABLE IF NOT EXISTS agent_events (
run_id TEXT NOT NULL REFERENCES agent_runs(run_id) ON DELETE CASCADE,
sequence INTEGER NOT NULL,
event TEXT NOT NULL,
data_json TEXT NOT NULL DEFAULT '{}',
timestamp TEXT NOT NULL,
PRIMARY KEY (run_id, sequence)
);
CREATE INDEX IF NOT EXISTS idx_agent_events_type
ON agent_events(run_id, event, sequence);
""",
] ]
+57 -6
View File
@@ -2,13 +2,14 @@ from collections.abc import AsyncIterator
from datetime import datetime, timezone from datetime import datetime, timezone
from uuid import uuid4 from uuid import uuid4
from fastapi import APIRouter, Query from fastapi import APIRouter, Header, Query
from fastapi.responses import StreamingResponse from fastapi.responses import StreamingResponse
from app.contracts import ( from app.contracts import (
AgentRun, AgentRun,
AgentRunCreateRequest, AgentRunCreateRequest,
AgentRunListResponse, AgentRunListResponse,
AgentTraceResponse,
ChatRequest, ChatRequest,
CredentialStatus, CredentialStatus,
CredentialWriteRequest, CredentialWriteRequest,
@@ -81,8 +82,9 @@ def utc_now() -> datetime:
return datetime.now(timezone.utc) return datetime.now(timezone.utc)
def as_sse(event: str, payload: str) -> str: def as_sse(event: str, payload: str, *, event_id: int | None = None) -> str:
return f"event: {event}\ndata: {payload}\n\n" id_line = f"id: {event_id}\n" if event_id is not None else ""
return f"{id_line}event: {event}\ndata: {payload}\n\n"
def provider_or_404(provider_id: str): def provider_or_404(provider_id: str):
@@ -314,16 +316,65 @@ async def cancel_agent_run(run_id: str) -> OperationResponse:
}, },
tags=["Agent"], tags=["Agent"],
) )
async def agent_events(run_id: str) -> StreamingResponse: async def agent_events(
run_id: str,
after_sequence: int | None = Query(default=None, ge=-1),
last_event_id: str | None = Header(default=None, alias="Last-Event-ID"),
) -> StreamingResponse:
agent_run_or_404(run_id) agent_run_or_404(run_id)
cursor = after_sequence
if cursor is None and last_event_id is not None:
try:
cursor = int(last_event_id)
except ValueError as exc:
raise ApiError(
400,
"TRACE_CURSOR_INVALID",
"Last-Event-ID must be an integer sequence.",
{"last_event_id": last_event_id},
) from exc
if cursor < -1:
raise ApiError(
400,
"TRACE_CURSOR_INVALID",
"Last-Event-ID must be greater than or equal to -1.",
)
cursor = cursor if cursor is not None else -1
async def stream() -> AsyncIterator[str]: async def stream() -> AsyncIterator[str]:
async for event in container.agent.events(run_id): async for event in container.agent.events(run_id, after_sequence=cursor):
yield as_sse(event.event.value, event.model_dump_json()) yield as_sse(
event.event.value,
event.model_dump_json(),
event_id=event.sequence,
)
return StreamingResponse(stream(), media_type="text/event-stream") return StreamingResponse(stream(), media_type="text/event-stream")
@router.get(
"/agent/runs/{run_id}/trace",
response_model=AgentTraceResponse,
tags=["Agent"],
)
async def get_agent_trace(
run_id: str,
after_sequence: int = Query(default=-1, ge=-1),
limit: int = Query(default=200, ge=1, le=500),
) -> AgentTraceResponse:
try:
return container.agent.get_trace(
run_id, after_sequence=after_sequence, limit=limit
)
except AgentRunNotFoundError as exc:
raise ApiError(
404,
"AGENT_RUN_NOT_FOUND",
f"Agent run does not exist: {run_id}",
{"run_id": run_id},
) from exc
@router.post( @router.post(
"/agent/runs/{run_id}/permissions/{request_id}", "/agent/runs/{run_id}/permissions/{request_id}",
response_model=OperationResponse, response_model=OperationResponse,
+200
View File
@@ -1,10 +1,18 @@
import asyncio import asyncio
from datetime import datetime, timezone
import pytest
from app.agent.trace_repository import AgentTraceRepository
from app.agent.permissions import PermissionMode from app.agent.permissions import PermissionMode
from app.agent.tools import ToolExecutionContext from app.agent.tools import ToolExecutionContext
from app.container import build_container from app.container import build_container
from app.database.db import connect
from app.errors import ApiError
from app.routes import agent_events
from app.contracts import ( from app.contracts import (
AgentEventType, AgentEventType,
AgentRun,
AgentRunCreateRequest, AgentRunCreateRequest,
AgentRunStatus, AgentRunStatus,
ToolCall, ToolCall,
@@ -108,8 +116,200 @@ def test_permission_confirmation_resumes_agent() -> None:
created.run_id, request_id, "allow_once" created.run_id, request_id, "allow_once"
) )
completed = await container.agent.wait(created.run_id) completed = await container.agent.wait(created.run_id)
events = [event async for event in container.agent.events(created.run_id)]
assert completed.status == AgentRunStatus.completed assert completed.status == AgentRunStatus.completed
assert completed.tool_results[0].success is True assert completed.tool_results[0].success is True
assert AgentEventType.permission_resolved in {event.event for event in events}
run(scenario())
def test_agent_trace_persists_and_replays_from_sequence() -> None:
async def scenario() -> None:
first = build_container()
created = await first.agent.create_run(
AgentRunCreateRequest(
input="persistent trace",
provider_id="mock",
model="mock-1",
metadata={"suite": "agent-benchmark-v1"},
)
)
completed = await first.agent.wait(created.run_id)
restarted = build_container()
restored = restarted.agent.get_run(created.run_id)
first_page = restarted.agent.get_trace(
created.run_id, after_sequence=-1, limit=2
)
second_page = restarted.agent.get_trace(
created.run_id,
after_sequence=first_page.next_sequence,
limit=100,
)
replay = [
event
async for event in restarted.agent.events(
created.run_id, after_sequence=first_page.next_sequence
)
]
assert completed.status == restored.status == AgentRunStatus.completed
assert first_page.has_more is True
assert [item.sequence for item in first_page.items] == [0, 1]
assert second_page.items[0].sequence == 2
assert replay == second_page.items
assert first_page.summary.model_calls == 1
assert first_page.summary.token_usage == completed.token_usage
assert first_page.config_snapshot["metadata"] == {
"suite": "agent-benchmark-v1"
}
assert second_page.items[-1].event == AgentEventType.run_completed
run(scenario())
def test_interrupted_persisted_run_is_closed_after_restart() -> None:
now = datetime.now(timezone.utc)
request = AgentRunCreateRequest(
input="interrupted",
provider_id="mock",
model="mock-1",
)
persisted = AgentRun(
run_id="run_interrupted",
status=AgentRunStatus.running,
input=request.input,
provider_id=request.provider_id,
model=request.model,
max_steps=request.max_steps,
created_at=now,
updated_at=now,
)
AgentTraceRepository().create_run(persisted, request, {"model": "mock-1"})
restarted = build_container()
recovered = restarted.agent.get_run(persisted.run_id)
events = run(
_collect_events(restarted.agent.events(persisted.run_id, after_sequence=-1))
)
assert recovered.status == AgentRunStatus.failed
assert recovered.error_code == "AGENT_PROCESS_RESTARTED"
assert events[-1].event == AgentEventType.run_failed
assert events[-1].sequence == 0
def test_trace_redacts_secrets_and_truncates_large_values() -> None:
async def scenario() -> None:
container = build_container()
secret = "sk-should-not-be-stored"
created = await container.agent.create_run(
AgentRunCreateRequest(
input=f'/tool system.echo {{"text":"{"x" * 4200}","api_key":"{secret}"}}',
provider_id="mock",
model="mock-1",
allowed_tools=["system.echo"],
metadata={"authorization": secret},
)
)
await container.agent.wait(created.run_id)
trace = container.agent.get_trace(
created.run_id, after_sequence=-1, limit=100
)
tool_call = next(
item for item in trace.items if item.event == AgentEventType.tool_call
)
assert tool_call.data["arguments"]["api_key"] == "[REDACTED]"
assert str(tool_call.data["arguments"]["text"]).endswith("...[TRUNCATED]")
assert trace.config_snapshot["metadata"]["authorization"] == "[REDACTED]"
assert secret not in trace.model_dump_json()
conn = connect()
try:
stored_row = conn.execute(
"""
SELECT run_json, request_json, config_snapshot_json
FROM agent_runs WHERE run_id = ?
""",
(created.run_id,),
).fetchone()
stored = "\n".join(str(value) for value in stored_row)
finally:
conn.close()
assert secret not in stored
run(scenario())
def test_persisted_agent_run_preserves_long_input_and_output() -> None:
"""审计事件可以限长,但重启后读取的 AgentRun 不能丢失正文。"""
now = datetime.now(timezone.utc)
long_input = "输入" * 2_500
long_output = "输出" * 2_500
request = AgentRunCreateRequest(
input=long_input,
provider_id="mock",
model="mock-1",
)
persisted = AgentRun(
run_id="run_long_content",
status=AgentRunStatus.completed,
input=long_input,
output=long_output,
provider_id=request.provider_id,
model=request.model,
max_steps=request.max_steps,
created_at=now,
updated_at=now,
)
repository = AgentTraceRepository()
repository.create_run(persisted, request, {"model": request.model})
restored = repository.get_run(persisted.run_id)
assert restored is not None
assert restored.input == long_input
assert restored.output == long_output
async def _collect_events(iterator):
return [event async for event in iterator]
def test_agent_sse_uses_last_event_id_and_emits_event_ids(monkeypatch) -> None:
async def scenario() -> None:
test_container = build_container()
monkeypatch.setattr("app.routes.container", test_container)
created = await test_container.agent.create_run(
AgentRunCreateRequest(
input="resume sse",
provider_id="mock",
model="mock-1",
)
)
await test_container.agent.wait(created.run_id)
response = await agent_events(
created.run_id, after_sequence=None, last_event_id="1"
)
chunks = [chunk async for chunk in response.body_iterator]
body = "".join(
chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk
for chunk in chunks
)
assert "id: 0\n" not in body
assert "id: 1\n" not in body
assert "id: 2\n" in body
assert "event: RunCompleted" in body
with pytest.raises(ApiError) as error:
await agent_events(
created.run_id, after_sequence=None, last_event_id="invalid"
)
assert error.value.code == "TRACE_CURSOR_INVALID"
run(scenario()) run(scenario())
+1
View File
@@ -89,6 +89,7 @@ def test_openapi_contains_documented_frontend_interfaces() -> None:
"/api/agent/runs", "/api/agent/runs",
"/api/agent/runs/{run_id}/cancel", "/api/agent/runs/{run_id}/cancel",
"/api/agent/runs/{run_id}/events", "/api/agent/runs/{run_id}/events",
"/api/agent/runs/{run_id}/trace",
"/api/skills", "/api/skills",
"/api/plugins", "/api/plugins",
"/api/plugins/install", "/api/plugins/install",
+22
View File
@@ -143,6 +143,10 @@ export type AgentEventType =
| 'PermissionRequired' | 'PermissionRequired'
| 'Usage' | 'Usage'
| 'Citation' | 'Citation'
| 'ModelCallStarted'
| 'ModelCallCompleted'
| 'ModelCallFailed'
| 'PermissionResolved'
| 'RunCompleted' | 'RunCompleted'
| 'RunFailed' | 'RunFailed'
| 'RunCancelled' | 'RunCancelled'
@@ -155,6 +159,24 @@ export interface AgentEvent {
timestamp: string timestamp: string
} }
export interface AgentTraceSummary {
model_calls: number
tool_calls: number
duration_ms: number
token_usage: number
errors: number
}
export interface AgentTraceResponse {
run_id: string
status: AgentRunStatus
items: AgentEvent[]
next_sequence: number
has_more: boolean
summary: AgentTraceSummary
config_snapshot: Record<string, unknown>
}
export interface ToolCall { export interface ToolCall {
tool_call_id: string tool_call_id: string
name: string name: string
+8
View File
@@ -18,6 +18,10 @@ const eventLabels: Record<AgentEventType, string> = {
PermissionRequired: '请求权限', PermissionRequired: '请求权限',
Usage: '用量统计', Usage: '用量统计',
Citation: '引用来源', Citation: '引用来源',
ModelCallStarted: '模型调用开始',
ModelCallCompleted: '模型调用完成',
ModelCallFailed: '模型调用失败',
PermissionResolved: '权限已处理',
RunCompleted: '运行完成', RunCompleted: '运行完成',
RunFailed: '运行失败', RunFailed: '运行失败',
RunCancelled: '运行取消', RunCancelled: '运行取消',
@@ -86,6 +90,10 @@ const detailLabels: Record<string, string> = {
total_tokens: '令牌总数', total_tokens: '令牌总数',
status: '状态', status: '状态',
duration_ms: '耗时(毫秒)', duration_ms: '耗时(毫秒)',
model_call_id: '模型调用 ID',
parent_model_call_id: '上级模型调用 ID',
finish_reason: '结束原因',
decision: '授权决定',
} }
export function runStatusLabel(status?: AgentRunStatus): string { export function runStatusLabel(status?: AgentRunStatus): string {
+12 -3
View File
@@ -1,6 +1,6 @@
import apiClient from './apiClient' import apiClient from './apiClient'
import { SseClient } from './sseClient' import { SseClient } from './sseClient'
import type { AgentRun, AgentEvent, ApiAgentRun, OperationResponse, PageMeta, ToolDefinition, PermissionRequest } from '@/contracts' import type { AgentRun, AgentEvent, AgentTraceResponse, ApiAgentRun, OperationResponse, PageMeta, ToolDefinition, PermissionRequest } from '@/contracts'
function toAgentRun(run: ApiAgentRun): AgentRun { function toAgentRun(run: ApiAgentRun): AgentRun {
// API 的 token_usage 是累计值,UI 模型预留了输入/输出拆分字段。 // API 的 token_usage 是累计值,UI 模型预留了输入/输出拆分字段。
@@ -54,6 +54,13 @@ export async function cancelAgentRun(runId: string): Promise<OperationResponse>
return apiClient.post(`/api/agent/runs/${runId}/cancel`) return apiClient.post(`/api/agent/runs/${runId}/cancel`)
} }
export async function getAgentTrace(
runId: string,
params?: { after_sequence?: number; limit?: number },
): Promise<AgentTraceResponse> {
return apiClient.get(`/api/agent/runs/${runId}/trace`, { params })
}
export async function listTools(): Promise<ToolDefinition[]> { export async function listTools(): Promise<ToolDefinition[]> {
const response = await apiClient.get<{ items: ToolDefinition[] }>('/api/tools') const response = await apiClient.get<{ items: ToolDefinition[] }>('/api/tools')
return response.items return response.items
@@ -66,12 +73,14 @@ export function streamAgentEvents(
onError?: (error: Error) => void onError?: (error: Error) => void
onDone?: () => void onDone?: () => void
onOpen?: () => void onOpen?: () => void
} },
afterSequence = -1,
): SseClient { ): SseClient {
// 将通用 SSE 包装成领域事件,Store 无需了解传输层 envelope。 // 将通用 SSE 包装成领域事件,Store 无需了解传输层 envelope。
const client = new SseClient({ const client = new SseClient({
url: `/api/agent/runs/${runId}/events`, url: `/api/agent/runs/${runId}/events?after_sequence=${afterSequence}`,
method: 'GET', method: 'GET',
lastEventId: afterSequence >= 0 ? String(afterSequence) : undefined,
onEvent: (eventName, data) => { onEvent: (eventName, data) => {
handlers.onEvent?.({ handlers.onEvent?.({
event: eventName as AgentEvent['event'], event: eventName as AgentEvent['event'],
+42
View File
@@ -0,0 +1,42 @@
import { afterEach, describe, expect, it, vi } from 'vitest'
import { SseClient } from './sseClient'
afterEach(() => {
vi.unstubAllGlobals()
vi.restoreAllMocks()
})
describe('SseClient resumable event transport', () => {
it('sends Last-Event-ID and exposes the returned SSE id', async () => {
const fetchMock = vi.fn().mockResolvedValue(
new Response(
'id: 3\nevent: ModelCallCompleted\ndata: {"sequence":3,"data":{"duration_ms":12}}\n\n',
{ status: 200, headers: { 'Content-Type': 'text/event-stream' } },
),
)
vi.stubGlobal('fetch', fetchMock)
const received = vi.fn()
const client = new SseClient({
url: '/api/agent/runs/run-1/events?after_sequence=2',
method: 'GET',
lastEventId: '2',
onEvent: received,
})
await client.connect()
expect(fetchMock).toHaveBeenCalledWith(
'/api/agent/runs/run-1/events?after_sequence=2',
expect.objectContaining({
method: 'GET',
headers: expect.objectContaining({ 'Last-Event-ID': '2' }),
}),
)
expect(received).toHaveBeenCalledWith(
'ModelCallCompleted',
{ sequence: 3, data: { duration_ms: 12 } },
'3',
)
})
})
+16 -4
View File
@@ -1,12 +1,17 @@
import { resolveApiUrl } from './apiClient' import { resolveApiUrl } from './apiClient'
export type SseEventHandler = (event: string, data: Record<string, unknown>) => void export type SseEventHandler = (
event: string,
data: Record<string, unknown>,
eventId?: string,
) => void
export interface SseClientOptions { export interface SseClientOptions {
url: string url: string
method?: string method?: string
body?: unknown body?: unknown
token?: string token?: string
lastEventId?: string
onEvent?: SseEventHandler onEvent?: SseEventHandler
onError?: (error: Error) => void onError?: (error: Error) => void
onOpen?: () => void onOpen?: () => void
@@ -26,7 +31,7 @@ export class SseClient {
} }
async connect() { async connect() {
const { url, method = 'POST', body, token, onEvent, onError, onOpen, onDone } = this.options const { url, method = 'POST', body, token, lastEventId, onEvent, onError, onOpen, onDone } = this.options
try { try {
const headers: Record<string, string> = { const headers: Record<string, string> = {
@@ -38,6 +43,9 @@ export class SseClient {
if (token) { if (token) {
headers['Authorization'] = `Bearer ${token}` headers['Authorization'] = `Bearer ${token}`
} }
if (lastEventId !== undefined) {
headers['Last-Event-ID'] = lastEventId
}
const resp = await fetch(resolveApiUrl(url), { const resp = await fetch(resolveApiUrl(url), {
method, method,
@@ -57,17 +65,19 @@ export class SseClient {
// 一个 UTF-8 字符或 SSE 行可能横跨多个网络分片,必须累积后再按空行派发。 // 一个 UTF-8 字符或 SSE 行可能横跨多个网络分片,必须累积后再按空行派发。
const decoder = new TextDecoder('utf-8') const decoder = new TextDecoder('utf-8')
let eventName = 'message' let eventName = 'message'
let eventId: string | undefined
let dataLines: string[] = [] let dataLines: string[] = []
let doneNotified = false let doneNotified = false
const dispatchEvent = () => { const dispatchEvent = () => {
if (!dataLines.length) { if (!dataLines.length) {
eventName = 'message' eventName = 'message'
eventId = undefined
return return
} }
try { try {
const data = JSON.parse(dataLines.join('\n')) as Record<string, unknown> const data = JSON.parse(dataLines.join('\n')) as Record<string, unknown>
onEvent?.(eventName, data) onEvent?.(eventName, data, eventId)
if (!doneNotified && ['Done', 'RunCompleted', 'RunFailed', 'RunCancelled'].includes(eventName)) { if (!doneNotified && ['Done', 'RunCompleted', 'RunFailed', 'RunCancelled'].includes(eventName)) {
doneNotified = true doneNotified = true
onDone?.() onDone?.()
@@ -76,6 +86,7 @@ export class SseClient {
onError?.(error instanceof Error ? error : new Error('Malformed SSE data')) onError?.(error instanceof Error ? error : new Error('Malformed SSE data'))
} }
eventName = 'message' eventName = 'message'
eventId = undefined
dataLines = [] dataLines = []
} }
@@ -87,6 +98,7 @@ export class SseClient {
let fieldValue = separator === -1 ? '' : line.slice(separator + 1) let fieldValue = separator === -1 ? '' : line.slice(separator + 1)
if (fieldValue.startsWith(' ')) fieldValue = fieldValue.slice(1) if (fieldValue.startsWith(' ')) fieldValue = fieldValue.slice(1)
if (field === 'event') eventName = fieldValue if (field === 'event') eventName = fieldValue
if (field === 'id') eventId = fieldValue
if (field === 'data') dataLines.push(fieldValue) if (field === 'data') dataLines.push(fieldValue)
} }
@@ -118,7 +130,7 @@ export class SseClient {
this.controller.abort() this.controller.abort()
} }
// TODO(streaming): Agent 事件持久化后,增加 Last-Event-ID 与指数退避重连。 // TODO(streaming): 桌面网络策略确定后,在 Store 层增加有上限的指数退避重连。
isConnected() { isConnected() {
return this.connected return this.connected