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:
@@ -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
|
||||||
|
|||||||
@@ -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: ...
|
||||||
|
|||||||
@@ -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 = (
|
||||||
|
|||||||
@@ -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
@@ -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 系统性高估。"""
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user