Files
Video-Tree-TRM5/tests/unit/test_telemetry.py
T
iomgaa 5a91f392f0 fix(telemetry): INSERT OR IGNORE + WAL + try/except 三层防御加固
根治遥测写入主键冲突(UNIQUE constraint)和并发写锁(database is locked)
导致的异常冒泡,遥测侧信道错误不再污染 LLM 重试链。

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-07-09 00:11:54 -04:00

142 lines
4.5 KiB
Python

"""adapters/telemetry.py 单元测试 — SQLiteTelemetryRecorder。"""
from __future__ import annotations
import sqlite3
import uuid
import pytest
from adapters.telemetry import SQLiteTelemetryRecorder
from core.protocols import TelemetryRecorder
# ---------------------------------------------------------------------------
# Fixtures
# ---------------------------------------------------------------------------
@pytest.fixture()
def db_path(tmp_path):
"""返回临时数据库路径。"""
return tmp_path / "telemetry_test.db"
@pytest.fixture()
def recorder(db_path):
"""构造 SQLiteTelemetryRecorder 实例。"""
return SQLiteTelemetryRecorder(db_path=db_path)
def _make_call_kwargs(*, cache_hit: bool = False, error: str | None = None):
"""构造 record_llm_call 的标准参数字典。"""
return dict(
call_id=str(uuid.uuid4()),
parent_call_id=None,
session_id="sess-001",
model_name="gpt-4o",
provider="openai",
messages='[{"role":"user","content":"hi"}]',
response="hello",
thinking="",
prompt_tokens=10,
completion_tokens=5,
latency_ms=120,
ttft_ms=45.2,
max_inter_token_ms=12.3,
cache_hit=cache_hit,
error=error,
)
# ---------------------------------------------------------------------------
# Tests
# ---------------------------------------------------------------------------
def test_satisfies_protocol(recorder):
"""SQLiteTelemetryRecorder 满足 TelemetryRecorder Protocol。"""
assert isinstance(recorder, TelemetryRecorder)
@pytest.mark.asyncio
async def test_record_creates_table_and_inserts(recorder, db_path):
"""首次写入应懒创建表并成功插入一条记录。"""
kwargs = _make_call_kwargs()
await recorder.record_llm_call(**kwargs)
conn = sqlite3.connect(str(db_path))
conn.row_factory = sqlite3.Row
rows = conn.execute("SELECT * FROM llm_calls").fetchall()
conn.close()
assert len(rows) == 1
row = rows[0]
assert row["call_id"] == kwargs["call_id"]
assert row["model_name"] == "gpt-4o"
assert row["prompt_tokens"] == 10
assert row["completion_tokens"] == 5
assert row["cache_hit"] == 0 # False → INTEGER 0
assert row["error"] is None
assert row["created_at"] is not None
@pytest.mark.asyncio
async def test_record_with_error(recorder, db_path):
"""error 字段非 None 时应正确存储。"""
kwargs = _make_call_kwargs(error="RateLimitError: 429")
await recorder.record_llm_call(**kwargs)
conn = sqlite3.connect(str(db_path))
conn.row_factory = sqlite3.Row
row = conn.execute("SELECT error FROM llm_calls WHERE call_id = ?", (kwargs["call_id"],)).fetchone()
conn.close()
assert row["error"] == "RateLimitError: 429"
@pytest.mark.asyncio
async def test_record_cache_hit(recorder, db_path):
"""cache_hit=True 时应存储为 INTEGER 1。"""
kwargs = _make_call_kwargs(cache_hit=True)
await recorder.record_llm_call(**kwargs)
conn = sqlite3.connect(str(db_path))
conn.row_factory = sqlite3.Row
row = conn.execute("SELECT cache_hit FROM llm_calls WHERE call_id = ?", (kwargs["call_id"],)).fetchone()
conn.close()
assert row["cache_hit"] == 1
@pytest.mark.asyncio
async def test_duplicate_call_id_does_not_raise(recorder, db_path):
"""重复 call_id 写入应静默忽略(INSERT OR IGNORE),不抛异常。"""
kwargs = _make_call_kwargs()
await recorder.record_llm_call(**kwargs)
await recorder.record_llm_call(**kwargs)
conn = sqlite3.connect(str(db_path))
rows = conn.execute("SELECT COUNT(*) FROM llm_calls").fetchone()
conn.close()
assert rows[0] == 1
@pytest.mark.asyncio
async def test_db_error_does_not_propagate(tmp_path):
"""SQLite 写入失败时 record_llm_call 应静默降级,不抛异常。"""
bad_recorder = SQLiteTelemetryRecorder(db_path=tmp_path / "nonexistent_dir" / "bad.db")
kwargs = _make_call_kwargs()
await bad_recorder.record_llm_call(**kwargs)
@pytest.mark.asyncio
async def test_concurrent_writes_no_lock_error(recorder, db_path):
"""16 路并发 record_llm_call 应全部成功,无 database is locked 错误。"""
import asyncio
tasks = []
for _ in range(16):
kwargs = _make_call_kwargs()
tasks.append(recorder.record_llm_call(**kwargs))
await asyncio.gather(*tasks)
conn = sqlite3.connect(str(db_path))
count = conn.execute("SELECT COUNT(*) FROM llm_calls").fetchone()[0]
conn.close()
assert count == 16