feat: add sampling overlay validation and source extra_body
Three pure helpers in the innermost layer plus ChatRequest.sampling as a cross-layer snapshot, so cache keys and telemetry read one stable value.
This commit is contained in:
@@ -4,11 +4,24 @@
|
||||
fake,字段顺序即公共承诺;新增字段只增不删且必带默认值。
|
||||
"""
|
||||
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass, field
|
||||
from types import MappingProxyType
|
||||
from typing import Any
|
||||
|
||||
_MISSING_DONE_DOMAIN = frozenset({"retry", "salvage"})
|
||||
|
||||
_PROTECTED_OVERLAY_KEYS: Mapping[str, str] = MappingProxyType(
|
||||
{
|
||||
"model": "会让遥测记录的 model 与实际请求分叉,成本按错单价换算",
|
||||
"messages": "会同时破坏缓存 key 与遥测的 messages 口径",
|
||||
"stream": "会绕过流式活性看门狗(TTFT/inter-token 超时全部失效)",
|
||||
"stream_options": "会丢 usage 帧,导致成本遥测归零、TPM 闸按预扣量结算失准",
|
||||
}
|
||||
)
|
||||
"""禁止出现在采样参数覆盖层里的键: 它们由治理层拥有,被覆盖即击穿治理。"""
|
||||
|
||||
USAGE_SOURCES = frozenset({"measured", "estimated", "unavailable"})
|
||||
"""usage_source 值域;仅约束库内生产侧取值,不在 frozen dataclass 上做运行时校验。"""
|
||||
|
||||
@@ -16,6 +29,47 @@ _EST_TOKENS_QUOTA_DIVISOR = 60
|
||||
"""未显式配置时的预扣量除数: 假定一次调用约占一秒钟的 TPM 配额份额。"""
|
||||
|
||||
|
||||
def validate_request_overlay(overlay: Mapping[str, Any], *, origin: str) -> dict[str, Any]:
|
||||
"""校验采样参数覆盖层并返回浅拷贝;origin 用于把错误指回配置/调用点。
|
||||
|
||||
两类校验缺一不可(issue #4 设计决策 B):保护键会击穿治理;不可 JSON
|
||||
序列化的值会在 `CacheMW` 的降级 try **之外**抛裸 `TypeError`——那条路径
|
||||
不属错误四分类、`TelemetryMW` 也不捕,结果是一行遥测都没有就崩了。
|
||||
两者都在进洋葱之前收口,故抛裸 `ValueError`(调用方编程错误,不可重试)。
|
||||
"""
|
||||
# Phase 1: 键形态——必须先于序列化试探,否则非 str 键会因 sort_keys 的
|
||||
# 比较失败被误报成"值不可序列化",把人指向错误的方向
|
||||
for key in overlay:
|
||||
if not isinstance(key, str):
|
||||
raise ValueError(f"{origin} 的键必须是 str: {key!r}(canonical JSON 要求)")
|
||||
# Phase 2: 保护键
|
||||
for key, reason in _PROTECTED_OVERLAY_KEYS.items():
|
||||
if key in overlay:
|
||||
raise ValueError(f"{origin} 不得覆盖 {key!r}: {reason}")
|
||||
# Phase 3: 值可序列化(缓存 key 与遥测列都要 json.dumps)
|
||||
try:
|
||||
json.dumps(dict(overlay), sort_keys=True, ensure_ascii=False)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise ValueError(
|
||||
f"{origin} 的值必须可 JSON 序列化(如 numpy 标量请先转 float/int): {exc}"
|
||||
) from exc
|
||||
return dict(overlay)
|
||||
|
||||
|
||||
def merge_sampling(
|
||||
extra_body: Mapping[str, Any], sampling: Mapping[str, Any]
|
||||
) -> dict[str, Any]:
|
||||
"""合并配置级与调用级采样参数;调用级优先(issue #4 设计决策 A)。"""
|
||||
return {**extra_body, **sampling}
|
||||
|
||||
|
||||
def canonical_sampling_json(merged: Mapping[str, Any]) -> str | None:
|
||||
"""缓存 key 与遥测 sampling 列共用的序列化口径;空 mapping → None。"""
|
||||
if not merged:
|
||||
return None
|
||||
return json.dumps(dict(merged), sort_keys=True, ensure_ascii=False)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class LLMResponse:
|
||||
"""一次治理调用的统一响应(与三项目超集兼容,ARCH §5.1)。"""
|
||||
@@ -58,6 +112,12 @@ class ChatRequest:
|
||||
structured: Any | None = None
|
||||
stream: bool = True
|
||||
overlay: dict[str, Any] = field(default_factory=dict)
|
||||
sampling: Mapping[str, Any] = field(default_factory=dict)
|
||||
"""调用方采样意图的快照,库内中间件**永不修改**(issue #4 设计决策 A)。
|
||||
|
||||
与 `overlay` 分开是因为后者会被结构化中间件注入 `response_format`,在洋葱
|
||||
不同深度取值不同;缓存 key 与三个遥测入口需要一个跨层恒定的读取点,否则
|
||||
同一列在不同行口径分叉。"""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@@ -118,11 +178,18 @@ class SourceConfig:
|
||||
enable_thinking: bool | None = None
|
||||
missing_done: str = "retry"
|
||||
trust_env: bool = True
|
||||
extra_body: Mapping[str, Any] = field(default_factory=dict)
|
||||
"""本源恒定的采样参数(如 `temperature=0`),并入请求体(issue #4)。
|
||||
|
||||
优先级低于调用级 overlay。注: 本字段令 `SourceConfig` 不再 hashable
|
||||
(加任何 mapping 字段的固有代价,裸 dict 亦然),库内无以源作 key 的写法;
|
||||
要可变副本用 `dict(source.extra_body)`,要改字段用 `dataclasses.replace`。"""
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self._validate_identity()
|
||||
self._validate_gates()
|
||||
self._validate_watchdog()
|
||||
self._freeze_extra_body()
|
||||
|
||||
def effective_est_tokens(self) -> int:
|
||||
"""TPM 入场预扣量: 显式配置优先,否则按 tpm 派生(设计 §2.2)。"""
|
||||
@@ -159,6 +226,13 @@ class SourceConfig:
|
||||
):
|
||||
raise ValueError("看门狗不变式要求 0 < inter_token < ttft < timeout_s")
|
||||
|
||||
def _freeze_extra_body(self) -> None:
|
||||
"""校验后转只读视图: 装配完成的源不应再被就地改采样参数(设计决策 E)。"""
|
||||
validated = validate_request_overlay(
|
||||
self.extra_body, origin=f"SourceConfig({self.name}).extra_body"
|
||||
)
|
||||
object.__setattr__(self, "extra_body", MappingProxyType(validated))
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RetryPolicy:
|
||||
|
||||
@@ -305,3 +305,92 @@ class TestOcrTypes:
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
OcrTextResult(text="x") # 溯源件不可省略
|
||||
|
||||
|
||||
class TestSamplingValidation:
|
||||
"""采样参数覆盖层的构造期校验(issue #4 设计决策 B)。"""
|
||||
|
||||
@pytest.mark.parametrize("key", ["model", "messages", "stream", "stream_options"])
|
||||
def test_protected_keys_rejected(self, key):
|
||||
"""保护键会击穿治理: 成本算错/口径失真/绕过看门狗与 usage 帧。"""
|
||||
from polygateway.types import validate_request_overlay
|
||||
|
||||
with pytest.raises(ValueError) as exc:
|
||||
validate_request_overlay({key: "x"}, origin="chat(overlay=...)")
|
||||
assert key in str(exc.value)
|
||||
assert "chat(overlay=...)" in str(exc.value) # 信息须能定位来源
|
||||
|
||||
def test_non_str_key_reports_key_problem(self):
|
||||
"""非 str 键须报"键必须是 str",不能被 sort_keys 的比较错误误报成不可序列化。"""
|
||||
from polygateway.types import validate_request_overlay
|
||||
|
||||
with pytest.raises(ValueError, match="str"):
|
||||
validate_request_overlay({1: "a", "b": 2}, origin="test")
|
||||
|
||||
def test_unserializable_value_becomes_value_error(self):
|
||||
"""裸 TypeError 会逃出 CacheMW 的降级 try 且一行遥测都没有(设计决策 B)。"""
|
||||
from polygateway.types import validate_request_overlay
|
||||
|
||||
with pytest.raises(ValueError, match="JSON"):
|
||||
validate_request_overlay({"temperature": object()}, origin="test")
|
||||
|
||||
def test_returns_independent_copy(self):
|
||||
"""调用方逐次改 seed 复用同一 dict 是预期模式,不拷贝会有竞态(决策 E)。"""
|
||||
from polygateway.types import validate_request_overlay
|
||||
|
||||
caller_dict = {"temperature": 0, "seed": 42}
|
||||
validated = validate_request_overlay(caller_dict, origin="test")
|
||||
caller_dict["seed"] = 43
|
||||
assert validated == {"temperature": 0, "seed": 42}
|
||||
|
||||
def test_merge_prefers_call_level(self):
|
||||
"""优先级: 调用级 > 配置级(设计决策 A)。"""
|
||||
from polygateway.types import merge_sampling
|
||||
|
||||
merged = merge_sampling({"temperature": 0, "top_p": 1}, {"temperature": 1})
|
||||
assert merged == {"temperature": 1, "top_p": 1}
|
||||
|
||||
def test_canonical_json_is_key_order_stable(self):
|
||||
"""缓存 key 与遥测列共用同一序列化口径,键序不得影响结果。"""
|
||||
from polygateway.types import canonical_sampling_json
|
||||
|
||||
assert canonical_sampling_json({"b": 1, "a": 2}) == canonical_sampling_json(
|
||||
{"a": 2, "b": 1}
|
||||
)
|
||||
assert canonical_sampling_json({}) is None
|
||||
|
||||
|
||||
class TestSourceConfigExtraBody:
|
||||
"""配置级采样参数(issue #4 设计决策 A/E)。"""
|
||||
|
||||
def test_defaults_to_empty_and_is_read_only(self):
|
||||
source = _make_source()
|
||||
assert source.extra_body == {}
|
||||
with pytest.raises(TypeError):
|
||||
source.extra_body["temperature"] = 0 # MappingProxyType 只读
|
||||
|
||||
def test_protected_key_rejected_at_construction(self):
|
||||
"""装配期报错,不放到运行时才炸(CLAUDE.md §4.5)。"""
|
||||
with pytest.raises(ValueError, match="model"):
|
||||
_make_source(extra_body={"model": "sneaky"})
|
||||
|
||||
def test_accepts_sampling_params(self):
|
||||
source = _make_source(extra_body={"temperature": 0})
|
||||
assert source.extra_body["temperature"] == 0
|
||||
|
||||
def test_replace_rebuilds_proxy(self):
|
||||
"""决策 G 的剥离依赖 replace 能重跑 __post_init__ 且不递归。"""
|
||||
source = _make_source(extra_body={"temperature": 0})
|
||||
stripped = dataclasses.replace(source, extra_body={})
|
||||
assert stripped.extra_body == {}
|
||||
with pytest.raises(TypeError):
|
||||
stripped.extra_body["x"] = 1
|
||||
|
||||
def test_no_longer_hashable_is_intentional(self):
|
||||
"""加 mapping 字段的固有代价(裸 dict 亦然),库内无调用点会踩。
|
||||
|
||||
锁定为有意行为: 将来踩到的人不应把它当 bug"修"回去——要可变副本用
|
||||
dict(source.extra_body),要改字段用 dataclasses.replace(设计 Task 1)。
|
||||
"""
|
||||
with pytest.raises(TypeError):
|
||||
hash(_make_source())
|
||||
|
||||
Reference in New Issue
Block a user