diff --git a/src/polygateway/types.py b/src/polygateway/types.py index 6b8980a..0aa2326 100644 --- a/src/polygateway/types.py +++ b/src/polygateway/types.py @@ -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: diff --git a/tests/unit/test_types.py b/tests/unit/test_types.py index f377268..1eae3b8 100644 --- a/tests/unit/test_types.py +++ b/tests/unit/test_types.py @@ -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())