feat: record sampling parameters in telemetry (port 20 to 21 fields)

Each of the three emitter entry points has a pinned meaning: only the
attempt path has an effective source, so only it merges extra_body.
This commit is contained in:
2026-07-31 21:30:45 -04:00
parent b6e4cc3f3b
commit 4516761dbe
6 changed files with 141 additions and 16 deletions
+12
View File
@@ -18,6 +18,7 @@ from loguru import logger
from polygateway.errors import GatewayUnavailableError, GovernanceBackendError from polygateway.errors import GatewayUnavailableError, GovernanceBackendError
from polygateway.middleware.cache import digest_messages from polygateway.middleware.cache import digest_messages
from polygateway.types import canonical_sampling_json, merge_sampling
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import Callable from collections.abc import Callable
@@ -63,6 +64,10 @@ class TelemetryEmitter:
error=error, error=error,
cached_prompt_tokens=response.cached_prompt_tokens if response else None, cached_prompt_tokens=response.cached_prompt_tokens if response else None,
model_reported=response.model_reported if response else None, model_reported=response.model_reported if response else None,
# 唯一有"生效源"的入口,故是唯一能并上 extra_body 的(设计决策 D)
sampling=canonical_sampling_json(
merge_sampling(source.extra_body, request.sampling)
),
) )
async def emit_cache_hit(self, *, request: ChatRequest, response: LLMResponse) -> None: async def emit_cache_hit(self, *, request: ChatRequest, response: LLMResponse) -> None:
@@ -87,6 +92,9 @@ class TelemetryEmitter:
# 统计供应商缓存命中率必须带 WHERE cache_hit = false,否则重复计数。 # 统计供应商缓存命中率必须带 WHERE cache_hit = false,否则重复计数。
cached_prompt_tokens=response.cached_prompt_tokens, cached_prompt_tokens=response.cached_prompt_tokens,
model_reported=response.model_reported, model_reported=response.model_reported,
# 由最外层 TelemetryMW 调用,手上没有 source。缓存命中行无损:
# sampling 已进缓存 key,能命中即意味调用级参数与历史那次逐字相同
sampling=canonical_sampling_json(request.sampling),
) )
async def emit_terminal_failure( async def emit_terminal_failure(
@@ -111,6 +119,8 @@ class TelemetryEmitter:
error=error, error=error,
cached_prompt_tokens=None, cached_prompt_tokens=None,
model_reported=None, model_reported=None,
# 无具体源,与 model/provider/source_name 置空同一先例(设计决策 D)
sampling=canonical_sampling_json(request.sampling),
) )
async def _record( async def _record(
@@ -133,6 +143,7 @@ class TelemetryEmitter:
error: str | None, error: str | None,
cached_prompt_tokens: int | None, cached_prompt_tokens: int | None,
model_reported: str | None, model_reported: str | None,
sampling: str | None,
) -> None: ) -> None:
try: try:
# 成本换算(M2 §6): 成功行按单价换算;缓存命中 0.0(未产生新调用); # 成本换算(M2 §6): 成功行按单价换算;缓存命中 0.0(未产生新调用);
@@ -172,6 +183,7 @@ class TelemetryEmitter:
cost=cost, cost=cost,
cached_prompt_tokens=cached_prompt_tokens, cached_prompt_tokens=cached_prompt_tokens,
model_reported=model_reported, model_reported=model_reported,
sampling=sampling,
) )
except asyncio.CancelledError: except asyncio.CancelledError:
raise raise
+1
View File
@@ -274,4 +274,5 @@ class TelemetryRecorder(Protocol):
cost: float | None, cost: float | None,
cached_prompt_tokens: int | None, cached_prompt_tokens: int | None,
model_reported: str | None, model_reported: str | None,
sampling: str | None,
) -> None: ... ) -> None: ...
+4 -1
View File
@@ -41,7 +41,8 @@ CREATE TABLE IF NOT EXISTS llm_calls (
cost DOUBLE PRECISION, cost DOUBLE PRECISION,
created_at TIMESTAMPTZ NOT NULL DEFAULT now(), created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
cached_prompt_tokens INTEGER, cached_prompt_tokens INTEGER,
model_reported TEXT model_reported TEXT,
sampling TEXT
); );
""" """
@@ -49,6 +50,7 @@ CREATE TABLE IF NOT EXISTS llm_calls (
_BACKFILL = ( _BACKFILL = (
("cached_prompt_tokens", "ALTER TABLE llm_calls ADD COLUMN cached_prompt_tokens INTEGER"), ("cached_prompt_tokens", "ALTER TABLE llm_calls ADD COLUMN cached_prompt_tokens INTEGER"),
("model_reported", "ALTER TABLE llm_calls ADD COLUMN model_reported TEXT"), ("model_reported", "ALTER TABLE llm_calls ADD COLUMN model_reported TEXT"),
("sampling", "ALTER TABLE llm_calls ADD COLUMN sampling TEXT"),
) )
# 探测现有列;尊重 search_path(to_regclass 按当前 search_path 解析) # 探测现有列;尊重 search_path(to_regclass 按当前 search_path 解析)
@@ -78,6 +80,7 @@ _COLUMNS = (
"cost", "cost",
"cached_prompt_tokens", "cached_prompt_tokens",
"model_reported", "model_reported",
"sampling",
) )
_INSERT = ( _INSERT = (
+9 -3
View File
@@ -36,13 +36,18 @@ CREATE TABLE IF NOT EXISTS llm_calls (
cost REAL, cost REAL,
created_at TEXT NOT NULL DEFAULT (datetime('now')), created_at TEXT NOT NULL DEFAULT (datetime('now')),
cached_prompt_tokens INTEGER, cached_prompt_tokens INTEGER,
model_reported TEXT model_reported TEXT,
sampling TEXT
); );
""" """
# 新列必须排在 created_at 之后: 旧表只能经 ALTER 追加到末尾,新建库若把它们 # 新列必须排在 created_at 之后: 旧表只能经 ALTER 追加到末尾,新建库若把它们
# 插在前面,两条路径的物理列序会分叉(列序断言测试无合规修法)。 # 插在前面,两条路径的物理列序会分叉(列序断言测试无合规修法)。
_BACKFILL_COLUMNS = (("cached_prompt_tokens", "INTEGER"), ("model_reported", "TEXT")) _BACKFILL_COLUMNS = (
("cached_prompt_tokens", "INTEGER"),
("model_reported", "TEXT"),
("sampling", "TEXT"),
)
_COLUMNS = ( _COLUMNS = (
"call_id", "call_id",
@@ -65,6 +70,7 @@ _COLUMNS = (
"cost", "cost",
"cached_prompt_tokens", "cached_prompt_tokens",
"model_reported", "model_reported",
"sampling",
) )
_INSERT = ( _INSERT = (
@@ -120,7 +126,7 @@ class SQLiteRecorder:
logger.warning("SQLite 遥测补列失败(写入将逐行降级): {}", exc) logger.warning("SQLite 遥测补列失败(写入将逐行降级): {}", exc)
async def record_llm_call(self, **fields: object) -> None: async def record_llm_call(self, **fields: object) -> None:
"""写一行遥测;字段集合即 20 字段冻结签名(ports.TelemetryRecorder)。""" """写一行遥测;字段集合即 21 字段冻结签名(ports.TelemetryRecorder)。"""
if self._conn is None: if self._conn is None:
return return
row = tuple(fields[col] for col in _COLUMNS) row = tuple(fields[col] for col in _COLUMNS)
@@ -41,6 +41,7 @@ _EXPECTED_COLUMNS = [
"created_at", "created_at",
"cached_prompt_tokens", "cached_prompt_tokens",
"model_reported", "model_reported",
"sampling",
] ]
# run 级前缀: 同库并存的其他运行(迁移批跑/另一开发机)互不可见 # run 级前缀: 同库并存的其他运行(迁移批跑/另一开发机)互不可见
@@ -104,6 +105,7 @@ async def _record_minimal(
"cost": None, "cost": None,
"cached_prompt_tokens": None, "cached_prompt_tokens": None,
"model_reported": None, "model_reported": None,
"sampling": None,
} }
fields.update(overrides) fields.update(overrides)
await recorder.record_llm_call(**fields) await recorder.record_llm_call(**fields)
+113 -12
View File
@@ -1,6 +1,7 @@
"""遥测子系统测试: SQLiteRecorder(20 列)+ TelemetryEmitter 单一 helper + TelemetryMW。""" """遥测子系统测试: SQLiteRecorder(21 列)+ TelemetryEmitter 单一 helper + TelemetryMW。"""
import asyncio import asyncio
import json
import sqlite3 import sqlite3
import subprocess import subprocess
from pathlib import Path from pathlib import Path
@@ -37,6 +38,7 @@ _EXPECTED_COLUMNS = [
"created_at", "created_at",
"cached_prompt_tokens", "cached_prompt_tokens",
"model_reported", "model_reported",
"sampling",
] ]
@@ -60,15 +62,17 @@ def _resp(**overrides):
return LLMResponse(**base) return LLMResponse(**base)
def _source(): def _source(**overrides):
return SourceConfig( base = {
name="s1", "name": "s1",
provider="p", "provider": "p",
base_url="https://gw.example/v1", "base_url": "https://gw.example/v1",
api_key="sk", "api_key": "sk",
model="m", "model": "m",
timeout_s=10.0, "timeout_s": 10.0,
) }
base.update(overrides)
return SourceConfig(**base)
# 输出单价 8 元/百万: 改前 `unavailable` 行按兜底的 0/4000 换算恰好是 0.032 # 输出单价 8 元/百万: 改前 `unavailable` 行按兜底的 0/4000 换算恰好是 0.032
@@ -97,6 +101,7 @@ async def _record_minimal(recorder, call_id="c1", **overrides):
"cost": None, "cost": None,
"cached_prompt_tokens": None, "cached_prompt_tokens": None,
"model_reported": None, "model_reported": None,
"sampling": None,
} }
fields.update(overrides) fields.update(overrides)
await recorder.record_llm_call(**fields) await recorder.record_llm_call(**fields)
@@ -153,6 +158,20 @@ class TestSQLiteRecorder:
assert rows["c-zero"] == 0 # 真实零命中,读回仍是 0 而非 NULL assert rows["c-zero"] == 0 # 真实零命中,读回仍是 0 而非 NULL
assert rows["c-none"] is None assert rows["c-none"] is None
async def test_sampling_column_round_trips(self, tmp_path):
"""issue #4: 采样参数落库,否则事后无法证明某批数据跑在什么温度下。"""
recorder = SQLiteRecorder(tmp_path / "t.db")
await _record_minimal(recorder, call_id="c-s", sampling='{"seed": 42, "temperature": 0}')
await _record_minimal(recorder, call_id="c-plain")
recorder.close()
rows = dict(
sqlite3.connect(tmp_path / "t.db")
.execute("SELECT call_id, sampling FROM llm_calls")
.fetchall()
)
assert json.loads(rows["c-s"]) == {"seed": 42, "temperature": 0}
assert rows["c-plain"] is None # 无采样参数为 NULL,便于 SQL 过滤
class TestSQLiteColumnBackfill: class TestSQLiteColumnBackfill:
"""issue #3: 已存在的 18 列旧表必须自动补列,否则每行写入都被丢弃。""" """issue #3: 已存在的 18 列旧表必须自动补列,否则每行写入都被丢弃。"""
@@ -258,7 +277,14 @@ class TestPostgresBackfillDiscipline:
"""PG 补列必须与 SQLite 侧对称: 失败只逐行降级,且稳态不抢排他锁(issue #3)。""" """PG 补列必须与 SQLite 侧对称: 失败只逐行降级,且稳态不抢排他锁(issue #3)。"""
_LEGACY = ["call_id", "cost", "created_at"] _LEGACY = ["call_id", "cost", "created_at"]
_CURRENT = ["call_id", "cost", "created_at", "cached_prompt_tokens", "model_reported"] _CURRENT = [
"call_id",
"cost",
"created_at",
"cached_prompt_tokens",
"model_reported",
"sampling",
]
def _recorder(self, conn): def _recorder(self, conn):
from polygateway.telemetry.postgres import PostgresRecorder from polygateway.telemetry.postgres import PostgresRecorder
@@ -286,8 +312,10 @@ class TestPostgresBackfillDiscipline:
async def test_missing_columns_are_added_once(self): async def test_missing_columns_are_added_once(self):
conn = _FakePgConn(self._LEGACY) conn = _FakePgConn(self._LEGACY)
await _record_minimal(self._recorder(conn)) await _record_minimal(self._recorder(conn))
from polygateway.telemetry.postgres import _BACKFILL
altered = [s for s in conn.statements if s.startswith("ALTER TABLE")] altered = [s for s in conn.statements if s.startswith("ALTER TABLE")]
assert len(altered) == 2 assert len(altered) == len(_BACKFILL) # 旧表缺全部补列,故一列一条 ALTER
assert all("IF NOT EXISTS" not in s for s in altered) # 探测已确认缺列,无需再判 assert all("IF NOT EXISTS" not in s for s in altered) # 探测已确认缺列,无需再判
@@ -394,6 +422,79 @@ class TestEmitterObservabilityFields:
assert rec.rows[0]["model_reported"] is None assert rec.rows[0]["model_reported"] is None
class TestEmitterSamplingColumn:
"""issue #4: sampling 列在三个入口的口径(设计决策 D 表格)。
列语义 = 「调用方采样意图 ⊎ 生效源 extra_body」,**不含**结构化注入的
response_format(列名是采样参数,schema 不是;且数 KB schema 逐行落库会让
审计表无谓膨胀)。三入口若各读各的层,同一列在不同行含义就不同。
"""
_SAMPLED = ChatRequest(
messages=[{"role": "user", "content": "hi"}],
sampling={"seed": 42},
overlay={"seed": 42, "response_format": {"type": "json_object"}},
)
async def test_attempt_merges_source_extra_body(self):
rec = _MemoryRecorder()
await TelemetryEmitter(rec).emit_attempt(
request=self._SAMPLED,
source=_source(extra_body={"temperature": 0}),
call_id="c",
latency_ms=1,
response=_resp(),
error=None,
)
assert json.loads(rec.rows[0]["sampling"]) == {"seed": 42, "temperature": 0}
async def test_response_format_never_leaks_into_the_column(self):
"""三行都不得出现 response_format——它不是采样参数。"""
rec = _MemoryRecorder()
emitter = TelemetryEmitter(rec)
await emitter.emit_attempt(
request=self._SAMPLED,
source=_source(),
call_id="c",
latency_ms=1,
response=_resp(),
error=None,
)
await emitter.emit_cache_hit(request=self._SAMPLED, response=_resp())
await emitter.emit_terminal_failure(
request=self._SAMPLED, call_id="c", latency_ms=1, error="dead"
)
assert len(rec.rows) == 3
for row in rec.rows:
assert "response_format" not in row["sampling"]
@pytest.mark.parametrize("emit", ["cache_hit", "terminal_failure"])
async def test_sourceless_entries_record_call_level_only(self, emit):
"""两个最外层入口没有"生效源"可言,与 model/source_name 置空同一先例。"""
rec = _MemoryRecorder()
emitter = TelemetryEmitter(rec)
if emit == "cache_hit":
await emitter.emit_cache_hit(request=self._SAMPLED, response=_resp())
else:
await emitter.emit_terminal_failure(
request=self._SAMPLED, call_id="c", latency_ms=1, error="dead"
)
assert json.loads(rec.rows[0]["sampling"]) == {"seed": 42}
async def test_absent_sampling_is_null(self):
"""无采样参数时为 NULL,而非空字符串或 "{}"——便于 SQL 过滤。"""
rec = _MemoryRecorder()
await TelemetryEmitter(rec).emit_attempt(
request=_REQ,
source=_source(),
call_id="c",
latency_ms=1,
response=_resp(),
error=None,
)
assert rec.rows[0]["sampling"] is None
class TestCostWithCachedTier: class TestCostWithCachedTier:
"""issue #3: 命中部分按缓存单价计费,避免 cost 系统性高估。""" """issue #3: 命中部分按缓存单价计费,避免 cost 系统性高估。"""