From 0aa7202c87a17c0ed052448eedf612b578c20b5d Mon Sep 17 00:00:00 2001 From: iomgaa Date: Fri, 31 Jul 2026 07:51:47 -0400 Subject: [PATCH] feat: collect provider cache tokens and reported model in transport --- src/polygateway/transports/openai_compat.py | 33 ++++++ tests/unit/test_openai_compat.py | 122 +++++++++++++++++++- 2 files changed, 154 insertions(+), 1 deletion(-) diff --git a/src/polygateway/transports/openai_compat.py b/src/polygateway/transports/openai_compat.py index 80f6615..117a069 100644 --- a/src/polygateway/transports/openai_compat.py +++ b/src/polygateway/transports/openai_compat.py @@ -45,6 +45,9 @@ def _sse_delta(chunk: dict[str, Any], usage_sink: dict[str, Any]) -> tuple[bool, """从 chunk 提取增量: (True, content) 或 (False, reasoning);usage 帧旁路进 sink。""" if chunk.get("usage"): usage_sink["usage"] = chunk["usage"] + if "model" not in usage_sink and chunk.get("model") is not None: + # 首次写入即固定: 末帧的异常值不得覆盖首帧报的真实版本(issue #3) + usage_sink["model"] = chunk["model"] choices = chunk.get("choices") or [] if not choices: return None @@ -152,6 +155,32 @@ def _resolve_usage(usage: dict[str, Any]) -> tuple[int, int, str]: return 0, 0, "unavailable" +def _coerce_cached_tokens(usage: Any) -> int | None: + """取 usage.prompt_tokens_details.cached_tokens(issue #3);形态异常一律 None。 + + `0` 与 `None` 必须可区分: 前者是"该源上报了一次真实零命中",后者是"该源 + 不报这个数",下游对两者的处置不同(后者不可做缓存成本校正)。故只把 + **负数与非整数**归 None,`0` 如实保留。`bool` 显式排除——isinstance(True, int) + 在 Python 里为真,放行会把 `True` 记成 1 个命中 token。 + """ + if not isinstance(usage, dict): + return None + details = usage.get("prompt_tokens_details") + if not isinstance(details, dict): + return None + cached = details.get("cached_tokens") + if isinstance(cached, bool) or not isinstance(cached, int) or cached < 0: + return None + return cached + + +def _coerce_model_reported(value: Any) -> str | None: + """取响应体的 model 字段(issue #3);非 str 或空白串一律 None。""" + if not isinstance(value, str) or not value.strip(): + return None + return value + + def _resolve_stream_usage(sink: dict[str, Any], salvaged: bool) -> tuple[int, int, str]: """流式用量口径: 打捞路径把 measured 降级为 estimated,unavailable 原样保留。 @@ -360,6 +389,8 @@ class OpenAICompatTransport: ttft_ms=ttft_ms, max_inter_token_ms=(max_gap if ttft_ms is not None else None), raw={"usage": sink.get("usage")}, + cached_prompt_tokens=_coerce_cached_tokens(sink.get("usage")), + model_reported=_coerce_model_reported(sink.get("model")), ) def _check_done( @@ -442,6 +473,8 @@ class OpenAICompatTransport: ttft_ms=None, max_inter_token_ms=None, raw={"usage": body.get("usage")}, + cached_prompt_tokens=_coerce_cached_tokens(body.get("usage")), + model_reported=_coerce_model_reported(body.get("model")), ) async def aclose(self) -> None: diff --git a/tests/unit/test_openai_compat.py b/tests/unit/test_openai_compat.py index 3058437..f76b046 100644 --- a/tests/unit/test_openai_compat.py +++ b/tests/unit/test_openai_compat.py @@ -36,7 +36,7 @@ def _source(**overrides): return SourceConfig(**base) -def _chunk(content=None, reasoning=None, usage=None): +def _chunk(content=None, reasoning=None, usage=None, model=None): delta = {} if content is not None: delta["content"] = content @@ -45,6 +45,8 @@ def _chunk(content=None, reasoning=None, usage=None): body = {"choices": [{"delta": delta}]} if (delta or usage is None) else {"choices": []} if usage is not None: body["usage"] = usage + if model is not None: + body["model"] = model return f"data: {json.dumps(body)}\n\n" @@ -261,6 +263,124 @@ class TestEmptyCompletion: await _complete(_transport_for(handler), _source()) +class TestObservabilityFields: + """issue #3: 供应商 prompt cache 命中数与 API 实际返回的模型版本串。 + + 网关报文一律不可信: 形态异常只归 None,绝不因一个可观测字段打断调用。 + """ + + def _cached_usage(self, cached): + return {**_USAGE, "prompt_tokens_details": {"cached_tokens": cached}} + + async def test_stream_reads_cached_tokens(self): + def handler(request): + return _sse_stream(_chunk(content="ok"), _chunk(usage=self._cached_usage(128))) + + result = await _complete(_transport_for(handler), _source()) + assert result.cached_prompt_tokens == 128 + + async def test_non_stream_reads_cached_tokens(self): + def handler(request): + return httpx.Response( + 200, + json={ + "choices": [{"message": {"content": "42"}}], + "usage": self._cached_usage(128), + }, + ) + + result = await _complete(_transport_for(handler), _source(), stream=False) + assert result.cached_prompt_tokens == 128 + + async def test_zero_cached_tokens_is_a_real_zero(self): + """0(真实零命中)与 None(该源未上报)必须可区分——issue #3 的核心诉求。""" + + def handler(request): + return _sse_stream(_chunk(content="ok"), _chunk(usage=self._cached_usage(0))) + + result = await _complete(_transport_for(handler), _source()) + assert result.cached_prompt_tokens == 0 + + async def test_usage_without_details_is_none(self): + def handler(request): + return _sse_stream(_chunk(content="ok"), _chunk(usage=_USAGE)) + + result = await _complete(_transport_for(handler), _source()) + assert result.cached_prompt_tokens is None + + async def test_missing_usage_frame_is_none(self): + def handler(request): + return _sse_stream(_chunk(content="ok")) + + result = await _complete(_transport_for(handler), _source()) + assert result.cached_prompt_tokens is None + + @pytest.mark.parametrize("bad", ["abc", -1, True, 1.5, None, [], {"x": 1}]) + async def test_malformed_cached_tokens_degrade_to_none(self, bad): + """`True` 必须排除: Python 里 isinstance(True, int) 为真。""" + + def handler(request): + return _sse_stream(_chunk(content="ok"), _chunk(usage=self._cached_usage(bad))) + + result = await _complete(_transport_for(handler), _source()) + assert result.cached_prompt_tokens is None + + async def test_details_not_a_dict_is_none(self): + def handler(request): + usage = {**_USAGE, "prompt_tokens_details": "oops"} + return _sse_stream(_chunk(content="ok"), _chunk(usage=usage)) + + result = await _complete(_transport_for(handler), _source()) + assert result.cached_prompt_tokens is None + + async def test_non_stream_reads_reported_model(self): + def handler(request): + return httpx.Response( + 200, + json={ + "choices": [{"message": {"content": "42"}}], + "usage": _USAGE, + "model": "MiniMax-Text-01-250321", + }, + ) + + result = await _complete(_transport_for(handler), _source(), stream=False) + assert result.model_reported == "MiniMax-Text-01-250321" + + async def test_stream_keeps_the_first_reported_model(self): + """末帧异常值不得覆盖首帧: 首次写入即固定。""" + + def handler(request): + return _sse_stream( + _chunk(content="a", model="MiniMax-Text-01-250321"), + _chunk(content="b", model="something-else"), + _chunk(usage=_USAGE), + ) + + result = await _complete(_transport_for(handler), _source()) + assert result.model_reported == "MiniMax-Text-01-250321" + + @pytest.mark.parametrize("bad", [None, "", " ", 123, {}]) + async def test_missing_or_malformed_model_is_none(self, bad): + def handler(request): + body = {"choices": [{"message": {"content": "42"}}], "usage": _USAGE} + if bad is not None: + body["model"] = bad + return httpx.Response(200, json=body) + + result = await _complete(_transport_for(handler), _source(), stream=False) + assert result.model_reported is None + + async def test_raw_payload_is_unchanged(self): + """新字段是独立格子,不改动 raw 的既有内容。""" + + def handler(request): + return _sse_stream(_chunk(content="ok"), _chunk(usage=self._cached_usage(5))) + + result = await _complete(_transport_for(handler), _source()) + assert set(result.raw) == {"usage"} + + class TestNonStreamFastPath: async def test_non_stream_parses_message(self): def handler(request):