196 lines
6.5 KiB
Python
196 lines
6.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
|
||
|
||
|
||
def test_high_concurrency_writes_zero_loss(recorder, db_path):
|
||
"""64 路线程并发直压 _write 应零丢失——复现生产 concurrency 下 database is locked 丢失。
|
||
|
||
直接压同步 _write(不经 to_thread 排队),最大化并发连接数以逼出锁竞争;
|
||
单连接 + threading.Lock 串行化模式下应全部落库、零丢失(对齐 HarnessLog)。
|
||
"""
|
||
import concurrent.futures
|
||
|
||
n = 64
|
||
kwargs_list = [_make_call_kwargs() for _ in range(n)]
|
||
with concurrent.futures.ThreadPoolExecutor(max_workers=n) as executor:
|
||
list(executor.map(lambda kw: recorder._write(**kw), kwargs_list))
|
||
|
||
conn = sqlite3.connect(str(db_path))
|
||
count = conn.execute("SELECT COUNT(*) FROM llm_calls").fetchone()[0]
|
||
conn.close()
|
||
assert count == n, f"并发写丢失 {n - count} 条(database is locked 降级丢弃): 落库 {count}/{n}"
|
||
|
||
|
||
def test_uses_single_persistent_connection(db_path, monkeypatch):
|
||
"""对齐 HarnessLog:单持久连接(构造时建一次),写入复用而非每次新建。
|
||
|
||
每次写新建连接是并发锁竞争根源(多连接争 SQLite 写锁,撑爆 busy_timeout);
|
||
单连接 + 进程内 Lock 串行化把并发控制拉到进程内,消除 SQLite 层锁竞争。
|
||
"""
|
||
connect_calls = {"n": 0}
|
||
real_connect = sqlite3.connect
|
||
|
||
def _counting_connect(*args, **kwargs):
|
||
connect_calls["n"] += 1
|
||
return real_connect(*args, **kwargs)
|
||
|
||
monkeypatch.setattr(sqlite3, "connect", _counting_connect)
|
||
|
||
recorder = SQLiteTelemetryRecorder(db_path=db_path)
|
||
after_init = connect_calls["n"]
|
||
for _ in range(10):
|
||
recorder._write(**_make_call_kwargs())
|
||
after_writes = connect_calls["n"]
|
||
|
||
assert after_init >= 1, "构造时应建立持久连接(对齐 HarnessLog)"
|
||
assert after_writes == after_init, (
|
||
f"写入期间新建了 {after_writes - after_init} 个连接(应复用单持久连接,"
|
||
"每次新连接是并发锁竞争根源)"
|
||
)
|