212 lines
8.0 KiB
Python
212 lines
8.0 KiB
Python
#!/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)
|