feat: add sqlite telemetry with single-emitter discipline
This commit is contained in:
@@ -0,0 +1,203 @@
|
||||
"""遥测子系统测试: SQLiteRecorder(18 列)+ TelemetryEmitter 单一 helper + TelemetryMW。"""
|
||||
|
||||
import asyncio
|
||||
import sqlite3
|
||||
import subprocess
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from polygateway.errors import CircuitOpenError, RequestRejectedError
|
||||
from polygateway.middleware.telemetry import TelemetryEmitter, TelemetryMW
|
||||
from polygateway.telemetry.sqlite import SQLiteRecorder
|
||||
from polygateway.types import ChatRequest, LLMResponse, SourceConfig
|
||||
|
||||
_REQ = ChatRequest(messages=[{"role": "user", "content": "hi"}], session_id="sess-1")
|
||||
|
||||
_EXPECTED_COLUMNS = [
|
||||
"call_id", "parent_call_id", "session_id", "model", "provider", "source_name",
|
||||
"messages", "response", "thinking", "prompt_tokens", "completion_tokens",
|
||||
"usage_source", "latency_ms", "ttft_ms", "max_inter_token_ms", "cache_hit",
|
||||
"error", "cost", "created_at",
|
||||
]
|
||||
|
||||
|
||||
def _resp(**overrides):
|
||||
base = dict(
|
||||
content="ok", thinking="", model="m", provider="p", prompt_tokens=1,
|
||||
completion_tokens=2, latency_ms=30, ttft_ms=None, max_inter_token_ms=None,
|
||||
cache_hit=False, call_id="cid-1", source_name="s1", usage_source="measured",
|
||||
)
|
||||
base.update(overrides)
|
||||
return LLMResponse(**base)
|
||||
|
||||
|
||||
def _source():
|
||||
return SourceConfig(
|
||||
name="s1", provider="p", base_url="https://gw.example/v1",
|
||||
api_key="sk", model="m", timeout_s=10.0,
|
||||
)
|
||||
|
||||
|
||||
async def _record_minimal(recorder, call_id="c1", **overrides):
|
||||
fields = dict(
|
||||
call_id=call_id, parent_call_id=None, session_id="sess-1", model="m",
|
||||
provider="p", source_name="s1", messages="[]", response="ok", thinking="",
|
||||
prompt_tokens=1, completion_tokens=2, usage_source="measured", latency_ms=10,
|
||||
ttft_ms=None, max_inter_token_ms=None, cache_hit=False, error=None, cost=None,
|
||||
)
|
||||
fields.update(overrides)
|
||||
await recorder.record_llm_call(**fields)
|
||||
|
||||
|
||||
class TestSQLiteRecorder:
|
||||
async def test_schema_has_frozen_columns(self, tmp_path):
|
||||
recorder = SQLiteRecorder(tmp_path / "t.db")
|
||||
await _record_minimal(recorder)
|
||||
recorder.close()
|
||||
cols = [r[1] for r in sqlite3.connect(tmp_path / "t.db").execute(
|
||||
"PRAGMA table_info(llm_calls)"
|
||||
)]
|
||||
assert cols == _EXPECTED_COLUMNS
|
||||
|
||||
async def test_call_id_idempotent(self, tmp_path):
|
||||
recorder = SQLiteRecorder(tmp_path / "t.db")
|
||||
await _record_minimal(recorder, call_id="dup")
|
||||
await _record_minimal(recorder, call_id="dup", response="second")
|
||||
recorder.close()
|
||||
rows = sqlite3.connect(tmp_path / "t.db").execute(
|
||||
"SELECT response FROM llm_calls WHERE call_id='dup'"
|
||||
).fetchall()
|
||||
assert rows == [("ok",)] # INSERT OR IGNORE: 第二次静默忽略
|
||||
|
||||
async def test_concurrent_writes_all_land(self, tmp_path):
|
||||
recorder = SQLiteRecorder(tmp_path / "t.db")
|
||||
await asyncio.gather(*(_record_minimal(recorder, call_id=f"c{i}") for i in range(50)))
|
||||
recorder.close()
|
||||
(count,) = sqlite3.connect(tmp_path / "t.db").execute(
|
||||
"SELECT COUNT(*) FROM llm_calls"
|
||||
).fetchone()
|
||||
assert count == 50
|
||||
|
||||
async def test_unwritable_path_degrades_silently(self):
|
||||
recorder = SQLiteRecorder(Path("/nonexistent-root/deep/t.db"))
|
||||
await _record_minimal(recorder) # 不抛
|
||||
recorder.close()
|
||||
|
||||
|
||||
class _MemoryRecorder:
|
||||
def __init__(self):
|
||||
self.rows = []
|
||||
|
||||
async def record_llm_call(self, **fields):
|
||||
self.rows.append(fields)
|
||||
|
||||
|
||||
class TestEmitter:
|
||||
async def test_attempt_success_row(self):
|
||||
rec = _MemoryRecorder()
|
||||
emitter = TelemetryEmitter(rec)
|
||||
await emitter.emit_attempt(
|
||||
request=_REQ, source=_source(), call_id="cid-1", latency_ms=42,
|
||||
response=_resp(), error=None,
|
||||
)
|
||||
row = rec.rows[0]
|
||||
assert row["call_id"] == "cid-1" and row["error"] is None
|
||||
assert row["session_id"] == "sess-1" and row["source_name"] == "s1"
|
||||
assert row["response"] == "ok" and row["cost"] is None
|
||||
|
||||
async def test_attempt_failure_row(self):
|
||||
rec = _MemoryRecorder()
|
||||
emitter = TelemetryEmitter(rec)
|
||||
await emitter.emit_attempt(
|
||||
request=_REQ, source=_source(), call_id="cid-2", latency_ms=7,
|
||||
response=None, error="TransientError: boom",
|
||||
)
|
||||
row = rec.rows[0]
|
||||
assert row["error"].startswith("TransientError")
|
||||
assert row["response"] == "" and row["usage_source"] == "estimated"
|
||||
|
||||
async def test_multimodal_messages_digested_before_storage(self):
|
||||
rec = _MemoryRecorder()
|
||||
emitter = TelemetryEmitter(rec)
|
||||
big = "data:image/png;base64," + "A" * 100_000
|
||||
req = ChatRequest(messages=[{"role": "user", "content": [
|
||||
{"type": "image_url", "image_url": {"url": big}},
|
||||
]}])
|
||||
await emitter.emit_attempt(
|
||||
request=req, source=_source(), call_id="c", latency_ms=1,
|
||||
response=None, error="x",
|
||||
)
|
||||
assert len(rec.rows[0]["messages"]) < 500 # base64 不整段进库(VT R12)
|
||||
|
||||
async def test_recorder_failure_swallowed(self):
|
||||
class Broken:
|
||||
async def record_llm_call(self, **fields):
|
||||
raise OSError("disk full")
|
||||
|
||||
emitter = TelemetryEmitter(Broken())
|
||||
await emitter.emit_attempt(
|
||||
request=_REQ, source=_source(), call_id="c", latency_ms=1,
|
||||
response=_resp(), error=None,
|
||||
) # 不抛(降级不冒泡)
|
||||
|
||||
|
||||
class TestTelemetryMW:
|
||||
async def test_cache_hit_recorded(self):
|
||||
rec = _MemoryRecorder()
|
||||
mw = TelemetryMW(TelemetryEmitter(rec))
|
||||
|
||||
async def terminal(request):
|
||||
return _resp(cache_hit=True, latency_ms=0, call_id="cache-cid")
|
||||
|
||||
resp = await mw(_REQ, terminal)
|
||||
assert resp.cache_hit
|
||||
assert len(rec.rows) == 1
|
||||
assert rec.rows[0]["cache_hit"] is True and rec.rows[0]["latency_ms"] == 0
|
||||
|
||||
async def test_normal_success_not_double_recorded(self):
|
||||
"""成功尝试由 RetryMW 逐次记录;最外层不得重复记。"""
|
||||
rec = _MemoryRecorder()
|
||||
mw = TelemetryMW(TelemetryEmitter(rec))
|
||||
|
||||
async def terminal(request):
|
||||
return _resp(cache_hit=False)
|
||||
|
||||
await mw(_REQ, terminal)
|
||||
assert rec.rows == []
|
||||
|
||||
async def test_scope_level_failure_recorded(self):
|
||||
rec = _MemoryRecorder()
|
||||
mw = TelemetryMW(TelemetryEmitter(rec))
|
||||
|
||||
async def terminal(request):
|
||||
raise CircuitOpenError(scope="llm", retry_after_s=30.0)
|
||||
|
||||
with pytest.raises(CircuitOpenError):
|
||||
await mw(_REQ, terminal)
|
||||
assert len(rec.rows) == 1 and "circuit_open" in rec.rows[0]["error"]
|
||||
|
||||
async def test_attempt_level_failure_not_double_recorded(self):
|
||||
"""RequestRejected 已被 RetryMW 逐次记录 → 最外层跳过。"""
|
||||
rec = _MemoryRecorder()
|
||||
mw = TelemetryMW(TelemetryEmitter(rec))
|
||||
|
||||
async def terminal(request):
|
||||
raise RequestRejectedError("400")
|
||||
|
||||
with pytest.raises(RequestRejectedError):
|
||||
await mw(_REQ, terminal)
|
||||
assert rec.rows == []
|
||||
|
||||
|
||||
def test_single_emitter_discipline():
|
||||
"""铁律执法: record_llm_call 在 src/ 的调用点只允许出现在 telemetry emitter。"""
|
||||
out = subprocess.run(
|
||||
["grep", "-rln", "record_llm_call(", "src/polygateway"],
|
||||
capture_output=True, text=True, cwd=Path(__file__).resolve().parents[2],
|
||||
).stdout.splitlines()
|
||||
callers = [
|
||||
p for p in out
|
||||
if not p.endswith(("ports.py", "telemetry/sqlite.py"))
|
||||
]
|
||||
assert callers == ["src/polygateway/middleware/telemetry.py"]
|
||||
Reference in New Issue
Block a user