212 lines
8.0 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
"""代理事件账本写路径的离线测试(不连库、不连网)。
用假连接捕获 SQL 与参数,固定 AgentTraceWriter 与 persist_agent_evidence 的
证据形状:事件闭集校验、序号单调、usage 归一、单事务原子性与密钥拦截。
"""
from __future__ import annotations
import json
import pathlib
import sys
import unittest
PROJECT_ROOT = pathlib.Path(__file__).resolve().parents[3]
SCRIPT_DIR = PROJECT_ROOT / "muse" / "authority" / "evidence" / "skills" / "record-run-evidence" / "scripts"
if str(SCRIPT_DIR) not in sys.path:
sys.path.insert(0, str(SCRIPT_DIR))
import agent_trace # noqa: E402
from agent_trace import AgentTraceWriter, model_ids_match, persist_agent_evidence # noqa: E402
class FakeCursor:
def __init__(self, log: list) -> None:
self._log = log
def execute(self, sql, params=None):
self._log.append((sql, params))
return self
def fetchone(self):
# RETURNING id 模拟:每次自增。
FakeConn.next_id += 1
return (FakeConn.next_id,)
def fetchall(self):
return []
def commit(self):
self._log.append(("COMMIT", None))
def rollback(self):
self._log.append(("ROLLBACK", None))
class FakeConn:
next_id = 1000
def __init__(self, log: list) -> None:
self._log = log
def execute(self, sql, params=None):
return FakeCursor(self._log).execute(sql, params)
def commit(self):
self._log.append(("COMMIT", None))
def rollback(self):
self._log.append(("ROLLBACK", None))
class FakeConnect:
"""返回上下文管理器形态的假连接,记录全部语句。"""
def __init__(self) -> None:
self.log: list = []
def __call__(self, *args, **kwargs):
outer = self
class _Ctx:
def __enter__(self):
return FakeConn(outer.log)
def __exit__(self, *exc):
return False
return _Ctx()
def sqls(self):
return [entry[0] for entry in self.log]
def params_of(self, sql_head: str):
for sql, params in self.log:
if sql.startswith(sql_head):
return params
raise AssertionError(f"未找到语句: {sql_head}")
class ModelMatchTest(unittest.TestCase):
def test_model_match_normalizes_provider_resolution(self):
"""框架把模式解析成完整 ID 时,比较模型叶名而不是误报漂移。"""
self.assertTrue(model_ids_match("claude-opus-5", "catproxy-anthropic/claude-opus-5"))
self.assertTrue(model_ids_match("catproxy-anthropic/claude-opus-5", "claude-opus-5"))
self.assertFalse(model_ids_match("catproxy-anthropic/claude-opus-5", "other/claude-opus-5"))
self.assertTrue(model_ids_match("gpt-5.6-sol", "GPT-5.6-Sol"))
self.assertTrue(model_ids_match("same", "same"))
self.assertFalse(model_ids_match("claude-opus-5", "claude-haiku-4-5"))
self.assertFalse(model_ids_match("", "x"))
self.assertFalse(model_ids_match(None, "x"))
class AgentTraceWriterTest(unittest.TestCase):
def test_emit_inserts_with_monotonic_seq_and_normalized_usage(self):
conn = FakeConnect()
writer = AgentTraceWriter(run_id="r1", framework="pi", agent_role="planner", connect=conn)
writer.emit("run.started", status="ok", requested_model_id="m-a", details={"a": 1})
writer.emit(
"model.completed",
status="ok",
requested_model_id="m-a",
actual_model_id="prov/m-a",
usage={"input": 10, "cacheRead": 5, "output": 7, "reasoning": 3},
cost_usd=0.5,
)
self.assertEqual(writer.seq, 2)
insert = next(sql for sql in conn.sqls() if sql.startswith("INSERT INTO example_agent_event"))
rows = [params for sql, params in conn.log if sql.startswith("INSERT INTO example_agent_event")]
self.assertEqual(rows[0][1], 1)
self.assertEqual(rows[1][1], 2)
self.assertEqual(rows[1][9], 15) # input_tokens = input + cacheRead
self.assertEqual(rows[1][10], 7) # output_tokens
self.assertEqual(rows[1][11], 5) # cached_tokens
self.assertEqual(rows[1][12], 0.5) # cost_usd
self.assertTrue(insert)
def test_emit_rejects_unknown_type_and_incomplete_model_event(self):
writer = AgentTraceWriter(run_id="r1", framework="pi", agent_role="writer", connect=FakeConnect())
with self.assertRaises(ValueError):
writer.emit("made.up")
with self.assertRaises(ValueError):
writer.emit("model.completed", status="ok", actual_model_id=None)
with self.assertRaises(ValueError):
writer.emit("run.started", status="maybe")
class PersistAgentEvidenceTest(unittest.TestCase):
BASE = dict(
run_id="r-evidence",
agent_role="planner",
system_prompt="ROLE PROMPT",
user_message="TASK + INPUT",
final_message='{"ok": true}',
transcript='{"type":"agent_start"}\n',
requested_model_id="m-a",
)
CALLS = [
{
"actual_model_id": "prov/m-a",
"usage": {"input": 3, "output": 4, "cacheRead": 5},
"stop_reason": "stop",
"cost_usd": 0.25,
}
]
def test_success_writes_lease_contents_and_llm_calls_in_one_txn(self):
conn = FakeConnect()
result = persist_agent_evidence(connect=conn, creator="dispatch-agent-task", model_calls=self.CALLS, **self.BASE)
self.assertEqual(result["status"], "written")
self.assertIn("COMMIT", conn.sqls())
lease_params = next(
params for sql, params in conn.log if sql.startswith("INSERT INTO example_raw_lease")
)
self.assertEqual(set(json.loads(lease_params[2])), {"prompt", "response", "supplier"})
contents = [p for sql, p in conn.log if sql.startswith("INSERT INTO example_raw_content")]
kinds = {row[1] for row in contents}
self.assertEqual(kinds, {"prompt", "response", "supplier"})
prompt_text = next(row[5] for sql, row in conn.log if sql.startswith("INSERT INTO example_raw_content") and row[1] == "prompt")
self.assertEqual(json.loads(prompt_text), {"system": "ROLE PROMPT", "user": "TASK + INPUT"})
calls = [p for sql, p in conn.log if sql.startswith("INSERT INTO example_llm_call")]
self.assertEqual(len(calls), 1)
self.assertIsNone(calls[0][0]) # window_key 不是 run_id 的替代列
self.assertEqual(calls[0][1], "r-evidence")
self.assertEqual(calls[0][4], "prov/m-a")
self.assertTrue(calls[0][5]) # model_match
self.assertEqual(calls[0][6], 8) # in = 3+5
self.assertEqual(calls[0][7], 5) # cached
self.assertEqual(calls[0][8], 4) # out
def test_missing_final_message_allowed_and_bad_input_fail_closed(self):
conn = FakeConnect()
base = dict(self.BASE)
base["final_message"] = None
result = persist_agent_evidence(connect=conn, model_calls=self.CALLS, **base)
self.assertEqual(result["status"], "written")
with self.assertRaises(ValueError):
persist_agent_evidence(
connect=conn, model_calls=self.CALLS, **{**self.BASE, "system_prompt": ""}
)
with self.assertRaises(ValueError):
persist_agent_evidence(connect=conn, model_calls=[{"actual_model_id": ""}], **self.BASE)
def test_secret_like_content_rejected(self):
conn = FakeConnect()
bad = dict(self.BASE)
bad["system_prompt"] = "api_key = sk-abcdef0123456789abcdef012"
with self.assertRaises(ValueError):
persist_agent_evidence(connect=conn, model_calls=self.CALLS, **bad)
def test_dry_run_rolls_back(self):
conn = FakeConnect()
result = persist_agent_evidence(
connect=conn, dry_run=True, model_calls=self.CALLS, **self.BASE
)
self.assertEqual(result["status"], "dry_run_ok")
self.assertIn("ROLLBACK", conn.sqls())
if __name__ == "__main__":
unittest.main(verbosity=2)