Files
PolyGateway/tests/unit/test_telemetry.py
T
iomgaa 42e429eb58 fix: void the cost of rows whose usage is unavailable
失败尝试与终态失败行的 usage_source 由 estimated 改 unavailable(用量确实
不可得),并在 TelemetryEmitter 的成本换算里为 unavailable 短路记 NULL。
短路刻意插在 cache_hit 分支之后: 缓存命中未产生新调用,0.0 是事实而非未知。
附 OCR 成功行的防回归钉(仍为 measured、settle 恒 0,设计 §3.3 剔出决定)。
2026-07-30 10:37:48 -04:00

331 lines
11 KiB
Python

"""遥测子系统测试: 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.pricing import ModelPrice, PricingTable
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 = {
"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,
)
# 输出单价 8 元/百万: 改前 `unavailable` 行按兜底的 0/4000 换算恰好是 0.032
_PRICING = PricingTable({"m": ModelPrice(input_per_1m=1.0, output_per_1m=8.0)})
async def _record_minimal(recorder, call_id="c1", **overrides):
fields = {
"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")
# 失败尝试没有任何用量信息可言 → unavailable(设计 §3.2 #6)
assert row["response"] == "" and row["usage_source"] == "unavailable"
assert row["cost"] is None
async def test_terminal_failure_row_is_unavailable(self):
rec = _MemoryRecorder()
await TelemetryEmitter(rec, pricing=_PRICING).emit_terminal_failure(
request=_REQ, call_id="cid-t", latency_ms=5, error="cancelled"
)
row = rec.rows[0]
assert row["usage_source"] == "unavailable" and row["cost"] is None
@pytest.mark.parametrize(("prompt", "completion"), [(0, 0), (0, 4000)])
async def test_unavailable_success_row_has_null_cost(self, prompt, completion):
"""产生了真实调用但用量不可得 → cost 记 NULL(设计 §3.1 不变式)。
参数第二组是改前兜底写出的 `0/4000` 形态: 那时换算出 0.032 的假金额。
"""
rec = _MemoryRecorder()
await TelemetryEmitter(rec, pricing=_PRICING).emit_attempt(
request=_REQ,
source=_source(),
call_id="cid-u",
latency_ms=42,
response=_resp(
usage_source="unavailable", prompt_tokens=prompt, completion_tokens=completion
),
error=None,
)
assert rec.rows[0]["cost"] is None
async def test_measured_row_still_priced(self):
"""对照组: 同一价格表下 measured 行照常换算,证明 None 不是价格表没接上。"""
rec = _MemoryRecorder()
await TelemetryEmitter(rec, pricing=_PRICING).emit_attempt(
request=_REQ,
source=_source(),
call_id="cid-m",
latency_ms=42,
response=_resp(prompt_tokens=0, completion_tokens=4000),
error=None,
)
assert rec.rows[0]["cost"] == pytest.approx(0.032)
async def test_cache_hit_keeps_zero_cost_even_when_unavailable(self):
"""缓存命中未产生新调用,0.0 是事实而非未知 → 短路必须排在 cache_hit 之后。"""
rec = _MemoryRecorder()
await TelemetryEmitter(rec, pricing=_PRICING).emit_cache_hit(
request=_REQ,
response=_resp(cache_hit=True, usage_source="unavailable", completion_tokens=4000),
)
assert rec.rows[0]["cache_hit"] is True and rec.rows[0]["cost"] == 0.0
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", "telemetry/postgres.py"))
]
assert callers == ["src/polygateway/middleware/telemetry.py"]