372 lines
15 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

#!/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",
]