feat(agent): 持久化Trace并支持SSE恢复

This commit is contained in:
2026-09-01 00:39:37 +08:00
parent 8da75d4420
commit 3cb197aafe
20 changed files with 949 additions and 70 deletions
+5
View File
@@ -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:
+178 -35
View File
@@ -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
+359
View File
@@ -0,0 +1,359 @@
"""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) -> Any:
"""递归净化 Trace 数据;键名疑似 Secret 时不保留原值。"""
if depth >= MAX_TRACE_DEPTH:
return "[MAX_DEPTH]"
if isinstance(value, dict):
sanitized: dict[str, Any] = {}
for index, (key, item) in enumerate(value.items()):
if 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)
)
return sanitized
if isinstance(value, (list, tuple)):
items = [
sanitize_trace_value(item, depth=depth + 1)
for item in value[:MAX_TRACE_COLLECTION]
]
if 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 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)
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:
return json.dumps(
sanitize_trace_value(run.model_dump(mode="json")), ensure_ascii=False
)