372 lines
15 KiB
Python
372 lines
15 KiB
Python
#!/usr/bin/env python3
|
||
"""代理事件账本与框架证据的写路径(07-Agent与Skill领域 §2 框架派发)。
|
||
|
||
Agent 框架适配器(如 dispatch-agent-task/pi_runner)把框架原生事件流归一后,
|
||
经 ``AgentTraceWriter`` 逐条追加进 ``example_agent_event``;运行结束后由
|
||
``persist_agent_evidence`` 把 system prompt、任务输入、最终输出、全量转录和
|
||
逐回合模型调用投影原子落库。本模块只做被动留痕:业务调用方不需要、也不能
|
||
决定"是否留痕";除幂等哈希与安全摘要外不携带任何正文原文。
|
||
|
||
数据库连接可注入(``connect=None`` 时懒加载 muse_db.connect),离线测试用假连接。
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import json
|
||
import math
|
||
from typing import Any, Callable, Mapping, Sequence
|
||
|
||
EVENT_TYPES = frozenset(
|
||
{
|
||
"run.started",
|
||
"agent.started",
|
||
"model.completed",
|
||
"tool.started",
|
||
"tool.completed",
|
||
"agent.completed",
|
||
"agent.failed",
|
||
"run.completed",
|
||
"run.failed",
|
||
}
|
||
)
|
||
CREATOR = "agent-trace"
|
||
|
||
# pi/anthropic 风格 usage 到账本列的归一口径;同义字段取首个,避免重复计数。
|
||
_USAGE_IN_KEYS = ("input", "input_tokens", "prompt_tokens")
|
||
_USAGE_CACHE_READ_KEYS = ("cacheRead", "cache_read_input_tokens", "cached_tokens")
|
||
_USAGE_CACHE_WRITE_KEYS = ("cacheWrite", "cache_creation_input_tokens")
|
||
_USAGE_OUT_KEYS = ("output", "output_tokens", "completion_tokens")
|
||
|
||
|
||
def _connect_factory(connect: Callable[..., Any] | None) -> Callable[..., Any]:
|
||
if connect is not None:
|
||
return connect
|
||
from muse_db import connect as muse_connect
|
||
|
||
return muse_connect
|
||
|
||
|
||
def _usage_int(data: Mapping[str, Any], keys: tuple[str, ...]) -> int:
|
||
"""读取第一种存在的 usage 字段;脏值与负值按 0 记账。"""
|
||
|
||
for key in keys:
|
||
if key not in data:
|
||
continue
|
||
try:
|
||
return max(0, int(data.get(key) or 0))
|
||
except (TypeError, ValueError, OverflowError):
|
||
return 0
|
||
return 0
|
||
|
||
|
||
def _tokens(usage: Mapping[str, Any] | None) -> tuple[int, int, int]:
|
||
"""把框架 usage 归一为 (input, output, cached);input 含 cache 读写。"""
|
||
|
||
data = usage if isinstance(usage, Mapping) else {}
|
||
cached = _usage_int(data, _USAGE_CACHE_READ_KEYS)
|
||
cache_write = _usage_int(data, _USAGE_CACHE_WRITE_KEYS)
|
||
value = _usage_int(data, _USAGE_IN_KEYS) + cached + cache_write
|
||
output = _usage_int(data, _USAGE_OUT_KEYS)
|
||
return value, output, cached
|
||
|
||
|
||
class AgentTraceWriter:
|
||
"""把归一事件逐条追加进 ``example_agent_event``(每条短事务,崩溃可审计)。"""
|
||
|
||
def __init__(
|
||
self,
|
||
*,
|
||
run_id: str,
|
||
framework: str,
|
||
agent_role: str,
|
||
connect: Callable[..., Any] | None = None,
|
||
creator: str = CREATOR,
|
||
) -> None:
|
||
if not run_id or len(run_id) > 64:
|
||
raise ValueError("run_id 不能为空且不超过 64 字符")
|
||
if not framework or len(framework) > 32:
|
||
raise ValueError("framework 不能为空且不超过 32 字符")
|
||
if not agent_role or len(agent_role) > 32:
|
||
raise ValueError("agent_role 不能为空且不超过 32 字符")
|
||
self.run_id = run_id
|
||
self.framework = framework
|
||
self.agent_role = agent_role
|
||
self.creator = creator
|
||
self._connect = _connect_factory(connect)
|
||
self._seq = 0
|
||
|
||
@property
|
||
def seq(self) -> int:
|
||
"""已写入的事件数(下一个序号 = seq + 1)。"""
|
||
|
||
return self._seq
|
||
|
||
def emit(
|
||
self,
|
||
event_type: str,
|
||
*,
|
||
status: str | None = None,
|
||
tool_name: str | None = None,
|
||
requested_model_id: str | None = None,
|
||
actual_model_id: str | None = None,
|
||
usage: Mapping[str, Any] | None = None,
|
||
cost_usd: float | None = None,
|
||
raw_ref: int | None = None,
|
||
details: Mapping[str, Any] | None = None,
|
||
) -> int:
|
||
"""追加一条事件;事件类型与字段合法性在本层失败关闭。"""
|
||
|
||
if event_type not in EVENT_TYPES:
|
||
raise ValueError(f"未知代理事件类型: {event_type}")
|
||
if status is not None and status not in ("ok", "error"):
|
||
raise ValueError("status 只能是 ok/error")
|
||
if event_type == "model.completed" and not actual_model_id:
|
||
raise ValueError("model.completed 必须携带 actual_model_id")
|
||
for field, value, limit in (
|
||
("tool_name", tool_name, 64),
|
||
("requested_model_id", requested_model_id, 64),
|
||
("actual_model_id", actual_model_id, 64),
|
||
):
|
||
if value is not None and (not isinstance(value, str) or len(value) > limit):
|
||
raise ValueError(f"{field} 非法或超过 {limit} 字符")
|
||
if cost_usd is not None:
|
||
try:
|
||
numeric_cost = float(cost_usd)
|
||
except (TypeError, ValueError, OverflowError) as exc:
|
||
raise ValueError("cost_usd 必须是非负有限数或 NULL") from exc
|
||
if not math.isfinite(numeric_cost) or numeric_cost < 0:
|
||
raise ValueError("cost_usd 必须是非负有限数或 NULL")
|
||
in_tokens, out_tokens, cached_tokens = _tokens(usage)
|
||
payload = json.dumps(details or {}, ensure_ascii=False, default=str)
|
||
from persist_raw import _check_no_secrets
|
||
|
||
_check_no_secrets(payload)
|
||
next_seq = self._seq + 1
|
||
with self._connect() as conn:
|
||
try:
|
||
conn.execute(
|
||
"INSERT INTO example_agent_event(run_id, seq, event_type, framework, agent_role, "
|
||
"tool_name, status, requested_model_id, actual_model_id, input_tokens, output_tokens, "
|
||
"cached_tokens, cost_usd, raw_ref, details, creator) "
|
||
"VALUES (%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,%s::jsonb,%s)",
|
||
(
|
||
self.run_id,
|
||
next_seq,
|
||
event_type,
|
||
self.framework,
|
||
self.agent_role,
|
||
tool_name,
|
||
status,
|
||
requested_model_id,
|
||
actual_model_id,
|
||
in_tokens,
|
||
out_tokens,
|
||
cached_tokens,
|
||
cost_usd,
|
||
raw_ref,
|
||
payload,
|
||
self.creator,
|
||
),
|
||
)
|
||
conn.commit()
|
||
self._seq = next_seq
|
||
except Exception:
|
||
conn.rollback()
|
||
raise
|
||
return self._seq
|
||
|
||
|
||
def model_ids_match(requested: str | None, actual: str | None) -> bool:
|
||
"""匹配 pi 的模型模式解析:完整 ID 精确匹配,单边省略 provider 时比较模型叶名。"""
|
||
|
||
if not isinstance(requested, str) or not isinstance(actual, str):
|
||
return False
|
||
req, act = requested.strip().lower(), actual.strip().lower()
|
||
if not req or not act:
|
||
return False
|
||
if "/" in req and "/" in act:
|
||
return req == act
|
||
return req.rsplit("/", 1)[-1] == act.rsplit("/", 1)[-1]
|
||
|
||
|
||
def persist_agent_evidence(
|
||
*,
|
||
run_id: str,
|
||
agent_role: str,
|
||
system_prompt: str,
|
||
user_message: str,
|
||
final_message: str | None,
|
||
transcript: str,
|
||
model_calls: Sequence[Mapping[str, Any]],
|
||
requested_model_id: str,
|
||
connect: Callable[..., Any] | None = None,
|
||
creator: str = "dispatch-agent-task",
|
||
dry_run: bool = False,
|
||
) -> dict[str, Any]:
|
||
"""把一次框架派发的全部证据原子落库。
|
||
|
||
一个事务内写入:raw 租约(purpose=agent_task)+ prompt/response/supplier 三份
|
||
raw 全文 + 每个模型回合一条 ``example_llm_call`` 投影。任一步失败整体回滚,
|
||
调用方拿不到看似成功却缺证据的结果。
|
||
"""
|
||
|
||
import sys
|
||
from pathlib import Path
|
||
|
||
here = Path(__file__).resolve().parent
|
||
if str(here) not in sys.path:
|
||
sys.path.insert(0, str(here))
|
||
from persist_raw import _bare_sha256, _check_no_secrets # noqa: E402
|
||
|
||
if not isinstance(run_id, str) or not run_id or len(run_id) > 64:
|
||
raise ValueError("agent 证据 run_id 为空或超过 64 字符")
|
||
if not isinstance(agent_role, str) or not agent_role or len(agent_role) > 32:
|
||
raise ValueError("agent 证据 agent_role 为空或超过 32 字符")
|
||
if not isinstance(system_prompt, str) or not system_prompt:
|
||
raise ValueError("agent 证据缺 system prompt")
|
||
if not isinstance(user_message, str) or not user_message:
|
||
raise ValueError("agent 证据缺任务输入")
|
||
if final_message is not None and not isinstance(final_message, str):
|
||
raise ValueError("agent 证据 final_message 必须是字符串或 NULL")
|
||
if not isinstance(transcript, str) or not transcript:
|
||
raise ValueError("agent 证据缺框架转录")
|
||
_check_no_secrets(system_prompt)
|
||
_check_no_secrets(user_message)
|
||
_check_no_secrets(transcript)
|
||
if final_message:
|
||
_check_no_secrets(final_message)
|
||
if not isinstance(requested_model_id, str) or not requested_model_id or len(requested_model_id) > 64:
|
||
raise ValueError("agent 证据 requested_model_id 为空或超过 64 字符")
|
||
if not isinstance(model_calls, (list, tuple)) or not all(
|
||
isinstance(call, Mapping) for call in model_calls
|
||
):
|
||
raise ValueError("agent 证据 model_calls 必须是对象数组")
|
||
|
||
prompt_request = json.dumps(
|
||
{"system": system_prompt, "user": user_message},
|
||
ensure_ascii=False,
|
||
sort_keys=True,
|
||
separators=(",", ":"),
|
||
)
|
||
prompt_sha = _bare_sha256(prompt_request)
|
||
final_sha = _bare_sha256(final_message) if final_message else None
|
||
transcript_sha = _bare_sha256(transcript)
|
||
# lease 的哈希清单按 raw kind 计数;system/user 哈希已封装在 prompt 内容内,
|
||
# 不另造不会对应 raw 行的清单项,保证轮次封存不变量可机械核对。
|
||
content_hashes = {"prompt": prompt_sha, "supplier": transcript_sha}
|
||
if final_sha is not None:
|
||
content_hashes["response"] = final_sha
|
||
|
||
connect_fn = _connect_factory(connect)
|
||
with connect_fn() as conn:
|
||
try:
|
||
lease_id = conn.execute(
|
||
"INSERT INTO example_raw_lease(run_id, source_version, content_hashes, purpose, status, creator) "
|
||
"VALUES (%s,%s,%s::jsonb,%s,%s,%s) RETURNING id",
|
||
(
|
||
run_id,
|
||
None,
|
||
json.dumps(content_hashes, ensure_ascii=False),
|
||
"agent_task",
|
||
"closed",
|
||
creator,
|
||
),
|
||
).fetchone()[0]
|
||
|
||
def _content(kind: str, text: str, role: str) -> int:
|
||
row = conn.execute(
|
||
"INSERT INTO example_raw_content(lease_id, kind, run_id, role, content_sha256, content, creator) "
|
||
"VALUES (%s,%s,%s,%s,%s,%s,%s) ON CONFLICT (lease_id, content_sha256) DO NOTHING "
|
||
"RETURNING id",
|
||
(lease_id, kind, run_id, role, _bare_sha256(text), text, creator),
|
||
).fetchone()
|
||
if row:
|
||
return row[0]
|
||
row = conn.execute(
|
||
"SELECT id FROM example_raw_content WHERE lease_id=%s AND content_sha256=%s",
|
||
(lease_id, _bare_sha256(text)),
|
||
).fetchone()
|
||
if not row:
|
||
raise RuntimeError(f"raw {kind} 幂等回读失败")
|
||
return row[0]
|
||
|
||
prompt_id = _content("prompt", prompt_request, agent_role)
|
||
response_id = _content("response", final_message, agent_role) if final_message else None
|
||
transcript_id = _content("supplier", transcript, agent_role)
|
||
|
||
llm_call_ids: list[int] = []
|
||
for index, call in enumerate(model_calls, start=1):
|
||
actual = call.get("actual_model_id")
|
||
if not isinstance(actual, str) or not actual or len(actual) > 64:
|
||
raise ValueError(f"第 {index} 个模型回合 actual_model_id 为空或超过 64 字符")
|
||
usage = call.get("usage") or {}
|
||
if not isinstance(usage, Mapping):
|
||
raise ValueError(f"第 {index} 个模型回合 usage 必须是对象")
|
||
cost = call.get("cost_usd")
|
||
if cost is not None:
|
||
try:
|
||
numeric_cost = float(cost)
|
||
except (TypeError, ValueError, OverflowError) as exc:
|
||
raise ValueError(f"第 {index} 个模型回合 cost_usd 非法") from exc
|
||
if not math.isfinite(numeric_cost) or numeric_cost < 0:
|
||
raise ValueError(f"第 {index} 个模型回合 cost_usd 非法")
|
||
duration = call.get("duration_ms")
|
||
if duration is not None:
|
||
if isinstance(duration, bool) or not isinstance(duration, int) or duration < 0:
|
||
raise ValueError(f"第 {index} 个模型回合 duration_ms 非法")
|
||
stop_reason = call.get("stop_reason")
|
||
if stop_reason is not None and not isinstance(stop_reason, str):
|
||
raise ValueError(f"第 {index} 个模型回合 stop_reason 非法")
|
||
in_tokens, out_tokens, cached_tokens = _tokens(usage)
|
||
row = conn.execute(
|
||
"INSERT INTO example_llm_call(window_key, run_id, caller, requested_model_id, "
|
||
"actual_model_id, model_match, in_tokens, cached_tokens, out_tokens, cost_usd, "
|
||
"stop_reason, duration_ms, prompt_sha256, raw_content_id, creator) "
|
||
"VALUES (%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,%s) RETURNING id",
|
||
(
|
||
None,
|
||
run_id,
|
||
creator,
|
||
requested_model_id,
|
||
actual,
|
||
model_ids_match(requested_model_id, actual),
|
||
in_tokens,
|
||
cached_tokens,
|
||
out_tokens,
|
||
cost if cost is not None else 0,
|
||
stop_reason[:32] if stop_reason else None,
|
||
duration,
|
||
prompt_sha,
|
||
transcript_id,
|
||
creator,
|
||
),
|
||
).fetchone()
|
||
llm_call_ids.append(row[0])
|
||
|
||
result = {
|
||
"status": "written",
|
||
"leaseId": lease_id,
|
||
"promptId": prompt_id,
|
||
"responseId": response_id,
|
||
"transcriptId": transcript_id,
|
||
"llmCallIds": llm_call_ids,
|
||
}
|
||
if dry_run:
|
||
conn.rollback()
|
||
result["status"] = "dry_run_ok"
|
||
result["note"] = "试跑已回滚,未落库"
|
||
else:
|
||
conn.commit()
|
||
return result
|
||
except Exception:
|
||
conn.rollback()
|
||
raise
|
||
|
||
|
||
__all__ = [
|
||
"AgentTraceWriter",
|
||
"CREATOR",
|
||
"EVENT_TYPES",
|
||
"model_ids_match",
|
||
"persist_agent_evidence",
|
||
]
|