diff --git a/README.md b/README.md index 45454e6..7bfa242 100644 --- a/README.md +++ b/README.md @@ -10,7 +10,7 @@ NotesAgent/ ├── frontend/ Vue 3 + TypeScript + Vite 前端 ├── backend/ FastAPI + Pydantic 后端 -├── docs/ 分工与技术栈说明 +├── docs/ 架构、契约、开发说明、协作规范与问题复盘 └── server sync/ 云同步服务预留目录,当前未实现 ``` @@ -36,7 +36,7 @@ python --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 ``` -当前回归基线为后端 76 项测试、前端 26 项测试,且生产构建通过。测试数量会随功能增长,以本地实际输出和 CI 为准。 +当前回归基线为后端 81 项测试、前端 27 项测试,且生产构建通过。测试数量会随功能增长,以本地实际输出和 CI 为准。 构建产物位于 `frontend/dist`,该目录不提交到 Git。 @@ -126,24 +126,15 @@ pnpm test | 文档 | 用途 | | --- | --- | -| [技术栈说明](docs/AI笔记软件技术栈说明-团队版-v2.3.md) | 目标架构、第二阶段技术边界与模块依赖 | -| [第一阶段分工表](docs/第一阶段分工表.md) | 成员职责、协作关系与当前交付状态 | -| [第二阶段分工表](docs/第二阶段团队分工表.md) | 第二阶段人员职责、任务顺序、协作关系与验收项 | -| [第一阶段测试验证操作手册](docs/第一阶段测试验证操作手册.md) | 自动化测试、接口主链路、前端人工验收与记录模板 | -| [后端接口契约](docs/后端接口契约-开发版.md) | HTTP/SSE 接口、错误和当前实现状态 | -| [第二阶段接口契约](docs/第二阶段接口契约-开发版.md) | 第二阶段公共 DTO、计划接口、SSE、错误码与联调顺序 | -| [AI Core 与 Agent Core](docs/AI-Core与Agent-Core开发说明.md) | Provider、Agent、Tool、Permission 与 Extension Core | -| [Knowledge 与 Retrieval Core](docs/Knowledge与Retrieval-Core开发说明.md) | Block、索引、混合检索和 Citation | -| [模型提供商与模型发现](docs/模型提供商与模型发现开发说明.md) | Provider 预设、模型发现和凭据边界 | -| [前端页面需求](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) | 前端工程、契约和交互问题复盘 | +| [文档总索引](docs/README.md) | 文档分类、阅读顺序和维护规则 | +| [技术栈说明](docs/architecture/AI笔记软件技术栈说明-团队版-v2.3.md) | 目标架构、第二阶段技术边界与模块依赖 | +| [第二阶段分工表](docs/architecture/第二阶段团队分工表.md) | 第二阶段人员职责、任务顺序、协作关系与验收项 | +| [后端接口契约](docs/contracts/后端接口契约-开发版.md) | HTTP/SSE 接口、错误和当前实现状态 | +| [第二阶段接口契约](docs/contracts/第二阶段接口契约-开发版.md) | 第二阶段公共 DTO、计划接口、SSE、错误码与联调顺序 | +| [AI Core 与 Agent Core](docs/development/AI-Core与Agent-Core开发说明.md) | Provider、Agent、Tool、Permission 与 Extension Core | +| [Git 使用细则](docs/guides/Git使用细则-团队开发版.md) | 分支、提交、PR、Review 与合并流程 | +| [CI/CD 细则](docs/guides/CI-CD细则-团队开发版.md) | Gitea 流水线、质量门禁、产物、发布与回滚规则 | +| [Agent Trace 复盘](docs/retrospectives/Agent-Core第二阶段问题与修复复盘.md) | Agent 持久化、SSE 恢复、事件契约与脱敏问题复盘 | ## 日常开发注意事项 @@ -153,6 +144,7 @@ pnpm test - API 默认监听 `127.0.0.1:8000`,前端默认监听 `127.0.0.1:5173`。 - 后端附件目录默认是 `backend/data/attachments`,可通过 `APP_ATTACHMENTS_PATH` 覆盖;该目录由桌面 Host 管理。 - 跨模块接口发生变化时,需要同步更新前后端类型和 `docs` 中的接口说明。 -- 当前已实现接口见 `docs/后端接口契约-开发版.md`,第二阶段规划接口见 `docs/第二阶段接口契约-开发版.md`;已实现能力以 `/openapi.json` 为准。 -- 前端页面、交互、状态管理和第一阶段验收要求见 `docs/前端页面需求说明-开发版.md`。 -- 分支、提交、Pull Request、Review 和冲突处理规范见 `docs/Git使用细则-团队开发版.md`。 +- 当前已实现接口见 `docs/contracts/后端接口契约-开发版.md`,第二阶段规划接口见 `docs/contracts/第二阶段接口契约-开发版.md`;已实现能力以 `/openapi.json` 为准。 +- 前端页面、交互、状态管理和第一阶段验收要求见 `docs/contracts/前端页面需求说明-开发版.md`。 +- 分支、提交、Pull Request、Review 和冲突处理规范见 `docs/guides/Git使用细则-团队开发版.md`。 +- CI 检查、产物、发布和回滚规范见 `docs/guides/CI-CD细则-团队开发版.md`。 diff --git a/backend/README.md b/backend/README.md index dc18dfb..c211fc1 100644 --- a/backend/README.md +++ b/backend/README.md @@ -23,10 +23,10 @@ uv run uvicorn app.main:app --reload --host 127.0.0.1 --port 8000 uv run pytest ``` -当前基线为 71 项测试通过。Provider API Key 可通过前端设置页写入,也可用 `OPENAI_API_KEY`、`DEEPSEEK_API_KEY` 或 `AINOTE_CREDENTIAL_` 注入;不要把真实密钥写入仓库。 +当前基线为 81 项测试通过。Provider API Key 可通过前端设置页写入,也可用 `OPENAI_API_KEY`、`DEEPSEEK_API_KEY` 或 `AINOTE_CREDENTIAL_` 注入;不要把真实密钥写入仓库。 -团队接口清单见 `../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`。 diff --git a/backend/app/agent/permissions.py b/backend/app/agent/permissions.py index aee818f..b0d7c8c 100644 --- a/backend/app/agent/permissions.py +++ b/backend/app/agent/permissions.py @@ -104,6 +104,11 @@ class PermissionManager: ticket.future.set_result(decision) 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: for key, ticket in list(self._pending.items()): if ticket.run_id == run_id: diff --git a/backend/app/agent/runtime.py b/backend/app/agent/runtime.py index 78af004..3daf87d 100644 --- a/backend/app/agent/runtime.py +++ b/backend/app/agent/runtime.py @@ -7,17 +7,20 @@ import json from collections.abc import AsyncIterator from dataclasses import dataclass, field from datetime import datetime, timezone +from time import perf_counter from typing import TYPE_CHECKING from uuid import uuid4 from app.agent.permissions import PermissionManager, PermissionMode from app.agent.tools import ToolExecutionContext, ToolNotFoundError, ToolRegistry +from app.agent.trace_repository import AgentTraceRepository, sanitize_trace_value from app.contracts import ( AgentEvent, AgentEventType, AgentRun, AgentRunCreateRequest, AgentRunStatus, + AgentTraceResponse, Citation, Message, MessageRole, @@ -61,6 +64,7 @@ class RunRecord: events: list[AgentEvent] = field(default_factory=list) subscribers: set[asyncio.Queue[AgentEvent]] = field(default_factory=set) task: asyncio.Task[None] | None = None + next_sequence: int = 0 class AgentRuntime: @@ -72,11 +76,13 @@ class AgentRuntime: tools: ToolRegistry, permissions: PermissionManager, skills: SkillRuntime | None = None, + trace_repository: AgentTraceRepository | None = None, ) -> None: self.providers = providers self.tools = tools self.permissions = permissions self.skills = skills + self.trace_repository = trace_repository or AgentTraceRepository() self._records: dict[str, RunRecord] = {} async def create_run(self, request: AgentRunCreateRequest) -> AgentRun: @@ -116,22 +122,38 @@ class AgentRuntime: skill_config=skill_config, allowed_tools=allowed_tools, ) + self.trace_repository.create_run( + run, + request, + self._config_snapshot(record), + ) self._records[run.run_id] = record record.task = asyncio.create_task(self._execute(record), name=run.run_id) return run.model_copy(deep=True) 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]: - records = sorted( - self._records.values(), key=lambda item: item.run.created_at, reverse=True - ) - items = [item.run.model_copy(deep=True) for item in records[offset : offset + limit]] - return items, len(records) + items, total = self.trace_repository.list_runs(limit=limit, offset=offset) + recovered = [ + self.trace_repository.recover_interrupted(item.run_id) or item + if item.run_id not in self._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: - 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: return record.run.model_copy(deep=True) record.run.cancelled = True @@ -144,23 +166,53 @@ class AgentRuntime: return record.run.model_copy(deep=True) def resolve_permission(self, run_id: str, request_id: str, decision: str) -> bool: - self._get_record(run_id) - return self.permissions.resolve(run_id, request_id, decision) + record = self._records.get(run_id) + 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]: - record = self._get_record(run_id) - # 先回放快照再订阅实时事件,使晚加入的 SSE 客户端也能恢复界面状态。 - # TODO(agent): 持久化事件并支持 Last-Event-ID,进程重启后仍可续传。 + async def events( + self, run_id: str, *, after_sequence: int = -1 + ) -> AsyncIterator[AgentEvent]: + 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() 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: for event in history: + last_sequence = event.sequence yield event - if record.run.status in TERMINAL_STATUSES: + if run.status in TERMINAL_STATUSES: return while True: event = await queue.get() + if event.sequence <= last_sequence: + continue + last_sequence = event.sequence yield event.model_copy(deep=True) if event.event in { AgentEventType.run_completed, @@ -172,7 +224,9 @@ class AgentRuntime: record.subscribers.discard(queue) 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: try: await asyncio.shield(record.task) @@ -180,6 +234,17 @@ class AgentRuntime: pass 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: try: async with asyncio.timeout(record.request.run_timeout_seconds): @@ -210,15 +275,51 @@ class AgentRuntime: for step in range(1, record.request.max_steps + 1): record.run.current_step = step record.run.updated_at = datetime.now(timezone.utc) - turn = await provider.complete( - ModelRequest( - provider_id=record.request.provider_id, - model=record.request.model, - system=(record.skill_config.system_prompt if record.skill_config else None), - messages=messages, - tools=allowed_tools, - metadata=self._request_metadata(record), + 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( + ModelRequest( + provider_id=record.request.provider_id, + model=record.request.model, + system=(record.skill_config.system_prompt if record.skill_config else None), + messages=messages, + tools=allowed_tools, + 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 self._publish( @@ -257,7 +358,7 @@ class AgentRuntime: async def execute(call: ToolCall) -> ToolResult: 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)) 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.") - async def _execute_tool(self, record: RunRecord, call: ToolCall) -> ToolResult: - self._publish(record, AgentEventType.tool_call, call.model_dump(mode="json")) + async def _execute_tool( + 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: registered = self.tools.get(call.name) except ToolNotFoundError: @@ -305,7 +411,9 @@ class AgentRuntime: error_code="TOOL_NOT_ALLOWED", 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 permission = registered.definition.permission if registered else None @@ -317,7 +425,9 @@ class AgentRuntime: error_code="NETWORK_NOT_ALLOWED", 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 mode = self.permissions.mode_for(permission) if mode == PermissionMode.deny: @@ -348,11 +458,13 @@ class AgentRuntime: error_code="PERMISSION_TIMEOUT", error_message="Tool permission confirmation timed out.", ) - 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 record.run.status = AgentRunStatus.running + record.run.updated_at = datetime.now(timezone.utc) + self.trace_repository.save_run(record.run) result = ( await self._invoke_tool(record, call) if decision in {"allow_once", "allow_session"} @@ -361,9 +473,21 @@ class AgentRuntime: else: 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 + 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: try: return await asyncio.wait_for( @@ -411,15 +535,19 @@ class AgentRuntime: def _publish( self, record: RunRecord, event_type: AgentEventType, data: dict[str, object] ) -> None: + sanitized = sanitize_trace_value(data) + assert isinstance(sanitized, dict) event = AgentEvent( event=event_type, run_id=record.run.run_id, - sequence=len(record.events), - data=data, + sequence=record.next_sequence, + data=sanitized, timestamp=datetime.now(timezone.utc), ) + record.next_sequence += 1 record.events.append(event) - # 内存事件只保留最近窗口;完整审计轨迹应由后续持久化层承担。 + self.trace_repository.append_event(record.run, event) + # 内存只保留实时订阅窗口;完整审计轨迹由 SQLite 保存。 if len(record.events) > MAX_EVENTS_PER_RUN: del record.events[: len(record.events) - MAX_EVENTS_PER_RUN] for queue in record.subscribers: @@ -433,6 +561,21 @@ class AgentRuntime: metadata["retrieval"] = record.skill_config.retrieval.model_dump(mode="json") 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: if not result.success or not isinstance(result.output, dict): return diff --git a/backend/app/agent/trace_repository.py b/backend/app/agent/trace_repository.py new file mode 100644 index 0000000..e09759a --- /dev/null +++ b/backend/app/agent/trace_repository.py @@ -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, + ) diff --git a/backend/app/contracts.py b/backend/app/contracts.py index e7c9be8..bb8bf90 100644 --- a/backend/app/contracts.py +++ b/backend/app/contracts.py @@ -329,6 +329,10 @@ class AgentEventType(str, Enum): permission_required = "PermissionRequired" usage = "Usage" citation = "Citation" + model_call_started = "ModelCallStarted" + model_call_completed = "ModelCallCompleted" + model_call_failed = "ModelCallFailed" + permission_resolved = "PermissionResolved" run_completed = "RunCompleted" run_failed = "RunFailed" run_cancelled = "RunCancelled" @@ -342,6 +346,24 @@ class AgentEvent(Contract): 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): decision: Literal["allow_once", "allow_session", "deny"] diff --git a/backend/app/database/migrations.py b/backend/app/database/migrations.py index 771e69a..02cdb33 100644 --- a/backend/app/database/migrations.py +++ b/backend/app/database/migrations.py @@ -69,6 +69,33 @@ MIGRATIONS: list[str] = [ ); CREATE INDEX IF NOT EXISTS idx_tasks_status_due ON tasks(status, due_at); """, + # v3: 第二阶段 Agent Trace;Run 与事件事实持久化,供 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); + """, ] diff --git a/backend/app/routes.py b/backend/app/routes.py index 31ff060..e5e2d07 100644 --- a/backend/app/routes.py +++ b/backend/app/routes.py @@ -2,13 +2,14 @@ from collections.abc import AsyncIterator from datetime import datetime, timezone from uuid import uuid4 -from fastapi import APIRouter, Query +from fastapi import APIRouter, Header, Query from fastapi.responses import StreamingResponse from app.contracts import ( AgentRun, AgentRunCreateRequest, AgentRunListResponse, + AgentTraceResponse, ChatRequest, CredentialStatus, CredentialWriteRequest, @@ -81,8 +82,9 @@ def utc_now() -> datetime: return datetime.now(timezone.utc) -def as_sse(event: str, payload: str) -> str: - return f"event: {event}\ndata: {payload}\n\n" +def as_sse(event: str, payload: str, *, event_id: int | None = None) -> str: + 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): @@ -314,16 +316,65 @@ async def cancel_agent_run(run_id: str) -> OperationResponse: }, 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) + 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 for event in container.agent.events(run_id): - yield as_sse(event.event.value, event.model_dump_json()) + async for event in container.agent.events(run_id, after_sequence=cursor): + yield as_sse( + event.event.value, + event.model_dump_json(), + event_id=event.sequence, + ) 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( "/agent/runs/{run_id}/permissions/{request_id}", response_model=OperationResponse, diff --git a/backend/tests/test_agent_core.py b/backend/tests/test_agent_core.py index 0aa843b..2cba7b7 100644 --- a/backend/tests/test_agent_core.py +++ b/backend/tests/test_agent_core.py @@ -1,10 +1,18 @@ 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.tools import ToolExecutionContext 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 ( AgentEventType, + AgentRun, AgentRunCreateRequest, AgentRunStatus, ToolCall, @@ -108,8 +116,200 @@ def test_permission_confirmation_resumes_agent() -> None: created.run_id, request_id, "allow_once" ) 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.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()) diff --git a/backend/tests/test_api.py b/backend/tests/test_api.py index d7d76c4..f62496e 100644 --- a/backend/tests/test_api.py +++ b/backend/tests/test_api.py @@ -89,6 +89,7 @@ def test_openapi_contains_documented_frontend_interfaces() -> None: "/api/agent/runs", "/api/agent/runs/{run_id}/cancel", "/api/agent/runs/{run_id}/events", + "/api/agent/runs/{run_id}/trace", "/api/skills", "/api/plugins", "/api/plugins/install", diff --git a/frontend/src/contracts/index.ts b/frontend/src/contracts/index.ts index afcdcde..0fcfcfd 100644 --- a/frontend/src/contracts/index.ts +++ b/frontend/src/contracts/index.ts @@ -143,6 +143,10 @@ export type AgentEventType = | 'PermissionRequired' | 'Usage' | 'Citation' + | 'ModelCallStarted' + | 'ModelCallCompleted' + | 'ModelCallFailed' + | 'PermissionResolved' | 'RunCompleted' | 'RunFailed' | 'RunCancelled' @@ -155,6 +159,24 @@ export interface AgentEvent { 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 +} + export interface ToolCall { tool_call_id: string name: string diff --git a/frontend/src/features/agent/labels.ts b/frontend/src/features/agent/labels.ts index 871e8d4..c6646d6 100644 --- a/frontend/src/features/agent/labels.ts +++ b/frontend/src/features/agent/labels.ts @@ -18,6 +18,10 @@ const eventLabels: Record = { PermissionRequired: '请求权限', Usage: '用量统计', Citation: '引用来源', + ModelCallStarted: '模型调用开始', + ModelCallCompleted: '模型调用完成', + ModelCallFailed: '模型调用失败', + PermissionResolved: '权限已处理', RunCompleted: '运行完成', RunFailed: '运行失败', RunCancelled: '运行取消', @@ -86,6 +90,10 @@ const detailLabels: Record = { total_tokens: '令牌总数', status: '状态', duration_ms: '耗时(毫秒)', + model_call_id: '模型调用 ID', + parent_model_call_id: '上级模型调用 ID', + finish_reason: '结束原因', + decision: '授权决定', } export function runStatusLabel(status?: AgentRunStatus): string { diff --git a/frontend/src/services/agentService.ts b/frontend/src/services/agentService.ts index d6a5f41..5c95728 100644 --- a/frontend/src/services/agentService.ts +++ b/frontend/src/services/agentService.ts @@ -1,6 +1,6 @@ import apiClient from './apiClient' 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 { // API 的 token_usage 是累计值,UI 模型预留了输入/输出拆分字段。 @@ -54,6 +54,13 @@ export async function cancelAgentRun(runId: string): Promise return apiClient.post(`/api/agent/runs/${runId}/cancel`) } +export async function getAgentTrace( + runId: string, + params?: { after_sequence?: number; limit?: number }, +): Promise { + return apiClient.get(`/api/agent/runs/${runId}/trace`, { params }) +} + export async function listTools(): Promise { const response = await apiClient.get<{ items: ToolDefinition[] }>('/api/tools') return response.items @@ -66,12 +73,14 @@ export function streamAgentEvents( onError?: (error: Error) => void onDone?: () => void onOpen?: () => void - } + }, + afterSequence = -1, ): SseClient { // 将通用 SSE 包装成领域事件,Store 无需了解传输层 envelope。 const client = new SseClient({ - url: `/api/agent/runs/${runId}/events`, + url: `/api/agent/runs/${runId}/events?after_sequence=${afterSequence}`, method: 'GET', + lastEventId: afterSequence >= 0 ? String(afterSequence) : undefined, onEvent: (eventName, data) => { handlers.onEvent?.({ event: eventName as AgentEvent['event'], diff --git a/frontend/src/services/sseClient.spec.ts b/frontend/src/services/sseClient.spec.ts new file mode 100644 index 0000000..ba4fa7a --- /dev/null +++ b/frontend/src/services/sseClient.spec.ts @@ -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', + ) + }) +}) diff --git a/frontend/src/services/sseClient.ts b/frontend/src/services/sseClient.ts index 30389b1..34db218 100644 --- a/frontend/src/services/sseClient.ts +++ b/frontend/src/services/sseClient.ts @@ -1,12 +1,17 @@ import { resolveApiUrl } from './apiClient' -export type SseEventHandler = (event: string, data: Record) => void +export type SseEventHandler = ( + event: string, + data: Record, + eventId?: string, +) => void export interface SseClientOptions { url: string method?: string body?: unknown token?: string + lastEventId?: string onEvent?: SseEventHandler onError?: (error: Error) => void onOpen?: () => void @@ -26,7 +31,7 @@ export class SseClient { } 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 { const headers: Record = { @@ -38,6 +43,9 @@ export class SseClient { if (token) { headers['Authorization'] = `Bearer ${token}` } + if (lastEventId !== undefined) { + headers['Last-Event-ID'] = lastEventId + } const resp = await fetch(resolveApiUrl(url), { method, @@ -57,17 +65,19 @@ export class SseClient { // 一个 UTF-8 字符或 SSE 行可能横跨多个网络分片,必须累积后再按空行派发。 const decoder = new TextDecoder('utf-8') let eventName = 'message' + let eventId: string | undefined let dataLines: string[] = [] let doneNotified = false const dispatchEvent = () => { if (!dataLines.length) { eventName = 'message' + eventId = undefined return } try { const data = JSON.parse(dataLines.join('\n')) as Record - onEvent?.(eventName, data) + onEvent?.(eventName, data, eventId) if (!doneNotified && ['Done', 'RunCompleted', 'RunFailed', 'RunCancelled'].includes(eventName)) { doneNotified = true onDone?.() @@ -76,6 +86,7 @@ export class SseClient { onError?.(error instanceof Error ? error : new Error('Malformed SSE data')) } eventName = 'message' + eventId = undefined dataLines = [] } @@ -87,6 +98,7 @@ export class SseClient { let fieldValue = separator === -1 ? '' : line.slice(separator + 1) if (fieldValue.startsWith(' ')) fieldValue = fieldValue.slice(1) if (field === 'event') eventName = fieldValue + if (field === 'id') eventId = fieldValue if (field === 'data') dataLines.push(fieldValue) } @@ -118,7 +130,7 @@ export class SseClient { this.controller.abort() } - // TODO(streaming): Agent 事件持久化后,增加 Last-Event-ID 与指数退避重连。 + // TODO(streaming): 桌面网络策略确定后,在 Store 层增加有上限的指数退避重连。 isConnected() { return this.connected