diff --git a/src/polygateway/client.py b/src/polygateway/client.py index f8352a6..7ba53ab 100644 --- a/src/polygateway/client.py +++ b/src/polygateway/client.py @@ -136,6 +136,7 @@ class GatewayClient: quota_full: str = "wait", telemetry: TelemetryRecorder | None = None, pricing: PricingTable | None = None, + text_cap: int | None = None, cache: CacheBackend | None = None, cache_namespace: str | None = None, cache_ttl_s: int | None = None, @@ -146,7 +147,11 @@ class GatewayClient: sleep: Any = asyncio.sleep, rng: Any = random.random, ) -> None: - emitter = TelemetryEmitter(telemetry, pricing=pricing) if telemetry is not None else None + emitter = ( + TelemetryEmitter(telemetry, pricing=pricing, text_cap=text_cap) + if telemetry is not None + else None + ) terminal = RetryMW( scope=scope, sources=sources, diff --git a/src/polygateway/embedding.py b/src/polygateway/embedding.py index 18ad334..5b2dc56 100644 --- a/src/polygateway/embedding.py +++ b/src/polygateway/embedding.py @@ -104,6 +104,7 @@ class EmbeddingClient: quota_full: str = "wait", telemetry: TelemetryRecorder | None = None, pricing: PricingTable | None = None, + text_cap: int | None = None, batch_size: int, normalize: bool = False, expected_dim: int | None = None, @@ -128,7 +129,9 @@ class EmbeddingClient: self._retry = retry self._bp = backpressure self._quota_full = quota_full - self._emitter = TelemetryEmitter(telemetry, pricing=pricing) if telemetry else None + self._emitter = ( + TelemetryEmitter(telemetry, pricing=pricing, text_cap=text_cap) if telemetry else None + ) self._telemetry = telemetry self._pricing = pricing self._batch_size = batch_size diff --git a/src/polygateway/middleware/telemetry.py b/src/polygateway/middleware/telemetry.py index 335a495..da1e9f8 100644 --- a/src/polygateway/middleware/telemetry.py +++ b/src/polygateway/middleware/telemetry.py @@ -55,6 +55,44 @@ def _canonical_meta_json(meta: Mapping[str, Any]) -> str: return json.dumps(dict(meta), sort_keys=True, ensure_ascii=False, allow_nan=False) +def _cap_text(text: str, cap: int | None) -> str: + """超出 cap 时头部硬切并附省略标记 `…(略 N 字)`;cap 为 None 原样返回。""" + if cap is None or len(text) <= cap: + return text + return f"{text[:cap]}…(略 {len(text) - cap} 字)" + + +def _cap_part(part: Any, cap: int) -> Any: + """多模态 part 的文本截断;非 `type == "text"` 的 part 原样返回同一对象。""" + if isinstance(part, dict) and part.get("type") == "text" and isinstance(part.get("text"), str): + return {**part, "text": _cap_text(part["text"], cap)} + return part + + +def _cap_messages(messages: list[dict[str, Any]], cap: int | None) -> list[dict[str, Any]]: + """对每条消息的文本 content 与多模态 part 中 type == "text" 的 text 逐条施加 cap。 + + 非字符串 content 原样放行(外部输入形状不可控,遥测路径不得因此抛错)。 + + **只产出新对象,严禁就地修改**: `digest_messages` 对 content 非 list 的消息是 + 原样透传**同一个 dict 对象**(`cache.py:43`),多模态里非 image_url 的 part 同理。 + 就地改它会一并污染调用方持有的 messages、后续重试尝试的请求体与缓存写入的 key, + 且全程无任何报错。 + """ + if cap is None: + return messages + capped: list[dict[str, Any]] = [] + for msg in messages: + content = msg.get("content") + if isinstance(content, str): + capped.append({**msg, "content": _cap_text(content, cap)}) + elif isinstance(content, list): + capped.append({**msg, "content": [_cap_part(part, cap) for part in content]}) + else: + capped.append(msg) + return capped + + @dataclass(frozen=True) class _AttemptUsage: """一次尝试的用量视图;默认值即"失败尝试"档(无用量可言,记 0 并标 unavailable)。 @@ -96,9 +134,20 @@ class _AttemptUsage: class TelemetryEmitter: """从请求与结果组装 24 字段并写入 recorder;一切写失败降级 warning。""" - def __init__(self, recorder: TelemetryRecorder, *, pricing: PricingTable | None = None) -> None: + def __init__( + self, + recorder: TelemetryRecorder, + *, + pricing: PricingTable | None = None, + text_cap: int | None, + ) -> None: + """`text_cap` 无默认值是有意的: 它是关键行为参数,漏传即静默改变落库正文。 + + 本类是库内部类,唯一构造者是三个公共 Client,必填能保证没有一处漏传。 + """ self._recorder = recorder self._pricing = pricing + self._text_cap = text_cap async def emit_attempt( self, @@ -241,8 +290,12 @@ class TelemetryEmitter: ) else: cost = None - # messages 落库前多模态摘要,与缓存 key 共用同一函数(VT R12) - messages_json = json.dumps(digest_messages(request.messages), ensure_ascii=False) + # messages 落库前多模态摘要,与缓存 key 共用同一函数(VT R12); + # 截断只发生在摘要之后、序列化之前的遥测分支,缓存路径不经过它(issue #12) + messages_json = json.dumps( + _cap_messages(digest_messages(request.messages), self._text_cap), + ensure_ascii=False, + ) await self._recorder.record_llm_call( call_id=call_id, parent_call_id=request.parent_call_id, @@ -251,8 +304,8 @@ class TelemetryEmitter: provider=provider, source_name=source_name, messages=messages_json, - response=response_text, - thinking=thinking, + response=_cap_text(response_text, self._text_cap), + thinking=_cap_text(thinking, self._text_cap), prompt_tokens=prompt_tokens, completion_tokens=completion_tokens, usage_source=usage_source, diff --git a/src/polygateway/ocr.py b/src/polygateway/ocr.py index 5a1b1d1..0c5ebe0 100644 --- a/src/polygateway/ocr.py +++ b/src/polygateway/ocr.py @@ -109,6 +109,7 @@ class OcrClient: backpressure: BackpressurePolicy, quota_full: str = "wait", telemetry: TelemetryRecorder | None = None, + text_cap: int | None = None, now: Callable[[], float] = time.monotonic, sleep: Callable[[float], Awaitable[None]] = asyncio.sleep, rng: Callable[[], float] = random.random, @@ -127,7 +128,7 @@ class OcrClient: self._retry = retry self._bp = backpressure self._quota_full = quota_full - self._emitter = TelemetryEmitter(telemetry) if telemetry else None + self._emitter = TelemetryEmitter(telemetry, text_cap=text_cap) if telemetry else None self._telemetry = telemetry self._memo = SourceCooldownMemo(now=now) self._now = now diff --git a/tests/unit/test_cache.py b/tests/unit/test_cache.py index c23340c..1b28405 100644 --- a/tests/unit/test_cache.py +++ b/tests/unit/test_cache.py @@ -9,7 +9,8 @@ import pytest from polygateway.backends.memory.cache import InMemoryCache from polygateway.errors import ResultInvalidError, TransientError from polygateway.middleware.cache import CacheMW, build_cache_key, digest_messages -from polygateway.types import ChatRequest, LLMResponse +from polygateway.middleware.telemetry import TelemetryEmitter +from polygateway.types import ChatRequest, LLMResponse, SourceConfig _MSGS = [{"role": "user", "content": "hi"}] @@ -327,3 +328,51 @@ class TestStructuredRehydration: key = build_cache_key("m", _MSGS, "proj", None) raw = await backend.get(key) assert raw is not None and "structured_data" not in json.loads(raw) + + +class TestTelemetryCapDoesNotPoisonTheCacheKey: + """红线之一(issue #12): 遥测截断绝不能改到缓存 key。 + + `digest_messages` 对 content 非 list 的消息**原样透传同一个 dict 对象** + (本文件上方公式测试依赖的也是这份对象),遥测拿到的与算 key 用的是同一份。 + 就地截断会让同一组 messages 在遥测前后算出两个不同的 key——全量 miss、 + 且没有任何报错。故这里测的是"截断没有就地改掉调用方的对象",不只是 + "截断函数是纯的"。 + """ + + class _Rows: + def __init__(self): + self.rows = [] + + async def record_llm_call(self, **fields): + self.rows.append(fields) + + async def test_key_is_byte_identical_across_a_capped_emit(self): + messages = [ + {"role": "user", "content": "合同正文" * 31}, + {"role": "user", "content": [{"type": "text", "text": "标书正文" * 30}]}, + ] + before = build_cache_key("m", messages, "proj", None) + + rec = self._Rows() + await TelemetryEmitter(rec, text_cap=8).emit_attempt( + request=ChatRequest(messages=messages), + source=SourceConfig( + name="s1", + provider="p", + base_url="https://gw.example/v1", + api_key="sk", + model="m", + timeout_s=10.0, + ), + call_id="c", + latency_ms=1, + response=_resp(), + error=None, + ) + # 截断确实发生了(否则本用例恒真) + logged = json.loads(rec.rows[0]["messages"]) + assert "(略 116 字)" in logged[0]["content"] + assert "(略 112 字)" in logged[1]["content"][0]["text"] + + assert build_cache_key("m", messages, "proj", None) == before diff --git a/tests/unit/test_openai_compat.py b/tests/unit/test_openai_compat.py index 2e2067c..01f2845 100644 --- a/tests/unit/test_openai_compat.py +++ b/tests/unit/test_openai_compat.py @@ -113,7 +113,7 @@ async def _recorded_cost(result, source): source_name=source.name, usage_source=result.usage_source, ) - await TelemetryEmitter(recorder, pricing=_PRICING).emit_attempt( + await TelemetryEmitter(recorder, pricing=_PRICING, text_cap=None).emit_attempt( request=ChatRequest(messages=[{"role": "user", "content": "hi"}]), source=source, call_id="cid-1", diff --git a/tests/unit/test_pricing.py b/tests/unit/test_pricing.py index b4bb24d..cdd76b4 100644 --- a/tests/unit/test_pricing.py +++ b/tests/unit/test_pricing.py @@ -169,7 +169,7 @@ def _source(model="qwen-max"): class TestEmitterCost: async def test_success_row_costed(self): rec = _MemoryRecorder() - emitter = TelemetryEmitter(rec, pricing=_TABLE) + emitter = TelemetryEmitter(rec, pricing=_TABLE, text_cap=None) await emitter.emit_attempt( request=_REQ, source=_source(), @@ -182,13 +182,13 @@ class TestEmitterCost: async def test_cache_hit_row_costs_zero(self): rec = _MemoryRecorder() - emitter = TelemetryEmitter(rec, pricing=_TABLE) + emitter = TelemetryEmitter(rec, pricing=_TABLE, text_cap=None) await emitter.emit_cache_hit(request=_REQ, response=_resp(cache_hit=True)) assert rec.rows[0]["cost"] == 0.0 async def test_failure_row_cost_none(self): rec = _MemoryRecorder() - emitter = TelemetryEmitter(rec, pricing=_TABLE) + emitter = TelemetryEmitter(rec, pricing=_TABLE, text_cap=None) await emitter.emit_attempt( request=_REQ, source=_source(), @@ -201,7 +201,7 @@ class TestEmitterCost: async def test_unknown_model_none_without_blocking(self): rec = _MemoryRecorder() - emitter = TelemetryEmitter(rec, pricing=_TABLE) + emitter = TelemetryEmitter(rec, pricing=_TABLE, text_cap=None) await emitter.emit_attempt( request=_REQ, source=_source(model="mystery"), @@ -215,7 +215,7 @@ class TestEmitterCost: async def test_no_pricing_keeps_none(self): """未注入价格表 = M1 现状: cost 恒 None(回归)。""" rec = _MemoryRecorder() - emitter = TelemetryEmitter(rec) + emitter = TelemetryEmitter(rec, text_cap=None) await emitter.emit_attempt( request=_REQ, source=_source(), diff --git a/tests/unit/test_telemetry.py b/tests/unit/test_telemetry.py index a04fb0d..6c1d03b 100644 --- a/tests/unit/test_telemetry.py +++ b/tests/unit/test_telemetry.py @@ -1,6 +1,7 @@ """遥测子系统测试: SQLiteRecorder(22 列)+ TelemetryEmitter 单一 helper + TelemetryMW。""" import asyncio +import copy import json import os import sqlite3 @@ -9,11 +10,27 @@ from pathlib import Path import pytest +from polygateway.backends.memory.breaker import InMemoryGate +from polygateway.backends.memory.limiter import InMemoryLimiter +from polygateway.embedding import EmbeddingClient from polygateway.errors import CircuitOpenError, RequestRejectedError +from polygateway.middleware.cache import digest_messages from polygateway.middleware.telemetry import TelemetryEmitter, TelemetryMW +from polygateway.ocr import OcrClient from polygateway.pricing import ModelPrice, PricingTable +from polygateway.sources import RoundRobinSelector from polygateway.telemetry.sqlite import SQLiteRecorder -from polygateway.types import ChatRequest, LLMResponse, SourceConfig +from polygateway.types import ( + BackpressurePolicy, + BreakerConfig, + ChatRequest, + EmbeddingTransportResult, + GlobalLimits, + LLMResponse, + OcrTextTransportResult, + RetryPolicy, + SourceConfig, +) _REQ = ChatRequest(messages=[{"role": "user", "content": "hi"}], session_id="sess-1") @@ -966,7 +983,7 @@ class TestEmitterRecorderContract: from polygateway.telemetry.schema import COLUMNS rec = _MemoryRecorder() - await TelemetryEmitter(rec).emit_attempt( + await TelemetryEmitter(rec, text_cap=None).emit_attempt( request=_REQ, source=_source(), call_id="cid-1", @@ -981,7 +998,7 @@ class TestEmitterRecorderContract: from polygateway.telemetry.schema import COLUMNS rec = _MemoryRecorder() - emitter = TelemetryEmitter(rec) + emitter = TelemetryEmitter(rec, text_cap=None) if emit == "attempt": await emitter.emit_attempt( request=_REQ, @@ -1005,7 +1022,7 @@ class TestEmitterObservabilityFields: async def test_attempt_carries_the_response_values(self): rec = _MemoryRecorder() - await TelemetryEmitter(rec).emit_attempt( + await TelemetryEmitter(rec, text_cap=None).emit_attempt( request=_REQ, source=_source(), call_id="cid-1", @@ -1019,7 +1036,7 @@ class TestEmitterObservabilityFields: async def test_failed_attempt_has_no_provider_facts(self): rec = _MemoryRecorder() - await TelemetryEmitter(rec).emit_attempt( + await TelemetryEmitter(rec, text_cap=None).emit_attempt( request=_REQ, source=_source(), call_id="cid-2", @@ -1034,7 +1051,7 @@ class TestEmitterObservabilityFields: async def test_cache_hit_replays_the_recorded_values(self): """决策 B1: 命中行原样回放,故命中率统计必须带 WHERE cache_hit = false。""" rec = _MemoryRecorder() - await TelemetryEmitter(rec).emit_cache_hit( + await TelemetryEmitter(rec, text_cap=None).emit_cache_hit( request=_REQ, response=_resp(cached_prompt_tokens=64, model_reported="m-real", reasoning_tokens=7), ) @@ -1045,7 +1062,7 @@ class TestEmitterObservabilityFields: async def test_terminal_failure_records_none(self): rec = _MemoryRecorder() - await TelemetryEmitter(rec).emit_terminal_failure( + await TelemetryEmitter(rec, text_cap=None).emit_terminal_failure( request=_REQ, call_id="c", latency_ms=1, error="dead" ) assert rec.rows[0]["cached_prompt_tokens"] is None @@ -1069,7 +1086,7 @@ class TestEmitterSamplingColumn: async def test_attempt_merges_source_extra_body(self): rec = _MemoryRecorder() - await TelemetryEmitter(rec).emit_attempt( + await TelemetryEmitter(rec, text_cap=None).emit_attempt( request=self._SAMPLED, source=_source(extra_body={"temperature": 0}), call_id="c", @@ -1082,7 +1099,7 @@ class TestEmitterSamplingColumn: async def test_response_format_never_leaks_into_the_column(self): """三行都不得出现 response_format——它不是采样参数。""" rec = _MemoryRecorder() - emitter = TelemetryEmitter(rec) + emitter = TelemetryEmitter(rec, text_cap=None) await emitter.emit_attempt( request=self._SAMPLED, source=_source(), @@ -1103,7 +1120,7 @@ class TestEmitterSamplingColumn: async def test_sourceless_entries_record_call_level_only(self, emit): """两个最外层入口没有"生效源"可言,与 model/source_name 置空同一先例。""" rec = _MemoryRecorder() - emitter = TelemetryEmitter(rec) + emitter = TelemetryEmitter(rec, text_cap=None) if emit == "cache_hit": await emitter.emit_cache_hit(request=self._SAMPLED, response=_resp()) else: @@ -1115,7 +1132,7 @@ class TestEmitterSamplingColumn: async def test_absent_sampling_is_null(self): """无采样参数时为 NULL,而非空字符串或 "{}"——便于 SQL 过滤。""" rec = _MemoryRecorder() - await TelemetryEmitter(rec).emit_attempt( + await TelemetryEmitter(rec, text_cap=None).emit_attempt( request=_REQ, source=_source(), call_id="c", @@ -1145,7 +1162,7 @@ class TestEmitterCallerDimensions: async def test_every_entry_point_carries_the_dimensions(self, emit): """三条路径写出的行都必须带维度: 漏掉任一条,该租户的账就永远对不上。""" rec = _MemoryRecorder() - emitter = TelemetryEmitter(rec) + emitter = TelemetryEmitter(rec, text_cap=None) if emit == "attempt": await emitter.emit_attempt( request=self._REQ_A, @@ -1178,7 +1195,7 @@ class TestEmitterCallerDimensions: meta={"batch": "old-batch"}, ) rec = _MemoryRecorder() - mw = TelemetryMW(TelemetryEmitter(rec)) + mw = TelemetryMW(TelemetryEmitter(rec, text_cap=None)) async def terminal(request): # 缓存层回放的是历史那次的响应对象(其 call_id 属于 historical 那次) @@ -1200,7 +1217,7 @@ class TestEmitterCallerDimensions: JSON 函数直接查询,NULL 则要每条查询都额外判空。 """ rec = _MemoryRecorder() - await TelemetryEmitter(rec).emit_attempt( + await TelemetryEmitter(rec, text_cap=None).emit_attempt( request=_REQ, # tenant_id=None, meta={} source=_source(), call_id="c", @@ -1215,7 +1232,7 @@ class TestEmitterCallerDimensions: async def test_meta_is_serialized_with_sorted_keys(self): """键序固定,同一份维度在任意两行里字节一致,可直接做等值比对与去重。""" rec = _MemoryRecorder() - await TelemetryEmitter(rec).emit_terminal_failure( + await TelemetryEmitter(rec, text_cap=None).emit_terminal_failure( request=self._REQ_A, call_id="c", latency_ms=1, error="dead" ) assert list(json.loads(rec.rows[0]["meta"])) == ["a_first", "m_mid", "z_last"] @@ -1224,7 +1241,7 @@ class TestEmitterCallerDimensions: """`ensure_ascii=False`: 中文维度按原文落库,而非 `\\uXXXX` 转义串。""" rec = _MemoryRecorder() req = ChatRequest(messages=[{"role": "user", "content": "hi"}], meta={"dept": "研发"}) - await TelemetryEmitter(rec).emit_terminal_failure( + await TelemetryEmitter(rec, text_cap=None).emit_terminal_failure( request=req, call_id="c", latency_ms=1, error="dead" ) assert "研发" in rec.rows[0]["meta"] @@ -1243,7 +1260,7 @@ class TestEmitterCallerDimensions: """ rec = _MemoryRecorder() req = ChatRequest(messages=[{"role": "user", "content": "hi"}], meta={"k": float("nan")}) - await TelemetryEmitter(rec).emit_terminal_failure( + await TelemetryEmitter(rec, text_cap=None).emit_terminal_failure( request=req, call_id="c", latency_ms=1, error="dead" ) assert rec.rows == [] @@ -1258,7 +1275,7 @@ class TestCostWithCachedTier: async def test_cached_hit_lowers_the_recorded_cost(self): rec = _MemoryRecorder() - emitter = TelemetryEmitter(rec, pricing=self._TABLE) + emitter = TelemetryEmitter(rec, pricing=self._TABLE, text_cap=None) full = _resp(prompt_tokens=1_000_000, completion_tokens=0) await emitter.emit_attempt( request=_REQ, @@ -1284,7 +1301,7 @@ class TestCostWithCachedTier: async def test_cache_hit_row_still_costs_zero(self): """缓存命中未产生新调用 → cost 恒 0.0,该短路必须排在任何换算之前。""" rec = _MemoryRecorder() - await TelemetryEmitter(rec, pricing=self._TABLE).emit_cache_hit( + await TelemetryEmitter(rec, pricing=self._TABLE, text_cap=None).emit_cache_hit( request=_REQ, response=_resp(prompt_tokens=1_000_000, cached_prompt_tokens=600_000), ) @@ -1292,7 +1309,7 @@ class TestCostWithCachedTier: async def test_unavailable_usage_still_costs_none(self): rec = _MemoryRecorder() - await TelemetryEmitter(rec, pricing=self._TABLE).emit_attempt( + await TelemetryEmitter(rec, pricing=self._TABLE, text_cap=None).emit_attempt( request=_REQ, source=_source(), call_id="c", @@ -1306,7 +1323,7 @@ class TestCostWithCachedTier: class TestEmitter: async def test_attempt_success_row(self): rec = _MemoryRecorder() - emitter = TelemetryEmitter(rec) + emitter = TelemetryEmitter(rec, text_cap=None) await emitter.emit_attempt( request=_REQ, source=_source(), @@ -1322,7 +1339,7 @@ class TestEmitter: async def test_attempt_failure_row(self): rec = _MemoryRecorder() - emitter = TelemetryEmitter(rec) + emitter = TelemetryEmitter(rec, text_cap=None) await emitter.emit_attempt( request=_REQ, source=_source(), @@ -1339,7 +1356,7 @@ class TestEmitter: async def test_terminal_failure_row_is_unavailable(self): rec = _MemoryRecorder() - await TelemetryEmitter(rec, pricing=_PRICING).emit_terminal_failure( + await TelemetryEmitter(rec, pricing=_PRICING, text_cap=None).emit_terminal_failure( request=_REQ, call_id="cid-t", latency_ms=5, error="cancelled" ) row = rec.rows[0] @@ -1352,7 +1369,7 @@ class TestEmitter: 参数第二组是改前兜底写出的 `0/4000` 形态: 那时换算出 0.032 的假金额。 """ rec = _MemoryRecorder() - await TelemetryEmitter(rec, pricing=_PRICING).emit_attempt( + await TelemetryEmitter(rec, pricing=_PRICING, text_cap=None).emit_attempt( request=_REQ, source=_source(), call_id="cid-u", @@ -1367,7 +1384,7 @@ class TestEmitter: async def test_measured_row_still_priced(self): """对照组: 同一价格表下 measured 行照常换算,证明 None 不是价格表没接上。""" rec = _MemoryRecorder() - await TelemetryEmitter(rec, pricing=_PRICING).emit_attempt( + await TelemetryEmitter(rec, pricing=_PRICING, text_cap=None).emit_attempt( request=_REQ, source=_source(), call_id="cid-m", @@ -1380,7 +1397,7 @@ class TestEmitter: async def test_cache_hit_keeps_zero_cost_even_when_unavailable(self): """缓存命中未产生新调用,0.0 是事实而非未知 → 短路必须排在 cache_hit 之后。""" rec = _MemoryRecorder() - await TelemetryEmitter(rec, pricing=_PRICING).emit_cache_hit( + await TelemetryEmitter(rec, pricing=_PRICING, text_cap=None).emit_cache_hit( request=_REQ, response=_resp(cache_hit=True, usage_source="unavailable", completion_tokens=4000), ) @@ -1388,7 +1405,7 @@ class TestEmitter: async def test_multimodal_messages_digested_before_storage(self): rec = _MemoryRecorder() - emitter = TelemetryEmitter(rec) + emitter = TelemetryEmitter(rec, text_cap=None) big = "data:image/png;base64," + "A" * 100_000 req = ChatRequest( messages=[ @@ -1415,7 +1432,7 @@ class TestEmitter: async def record_llm_call(self, **fields): raise OSError("disk full") - emitter = TelemetryEmitter(Broken()) + emitter = TelemetryEmitter(Broken(), text_cap=None) await emitter.emit_attempt( request=_REQ, source=_source(), @@ -1429,7 +1446,7 @@ class TestEmitter: class TestTelemetryMW: async def test_cache_hit_recorded(self): rec = _MemoryRecorder() - mw = TelemetryMW(TelemetryEmitter(rec)) + mw = TelemetryMW(TelemetryEmitter(rec, text_cap=None)) async def terminal(request): return _resp(cache_hit=True, latency_ms=0, call_id="cache-cid") @@ -1442,7 +1459,7 @@ class TestTelemetryMW: async def test_normal_success_not_double_recorded(self): """成功尝试由 RetryMW 逐次记录;最外层不得重复记。""" rec = _MemoryRecorder() - mw = TelemetryMW(TelemetryEmitter(rec)) + mw = TelemetryMW(TelemetryEmitter(rec, text_cap=None)) async def terminal(request): return _resp(cache_hit=False) @@ -1452,7 +1469,7 @@ class TestTelemetryMW: async def test_scope_level_failure_recorded(self): rec = _MemoryRecorder() - mw = TelemetryMW(TelemetryEmitter(rec)) + mw = TelemetryMW(TelemetryEmitter(rec, text_cap=None)) async def terminal(request): raise CircuitOpenError(scope="llm", retry_after_s=30.0) @@ -1464,7 +1481,7 @@ class TestTelemetryMW: async def test_attempt_level_failure_not_double_recorded(self): """RequestRejected 已被 RetryMW 逐次记录 → 最外层跳过。""" rec = _MemoryRecorder() - mw = TelemetryMW(TelemetryEmitter(rec)) + mw = TelemetryMW(TelemetryEmitter(rec, text_cap=None)) async def terminal(request): raise RequestRejectedError("400") @@ -1488,3 +1505,174 @@ def test_single_emitter_discipline(): if not p.endswith(("ports.py", "telemetry/sqlite.py", "telemetry/postgres.py")) ] assert callers == ["src/polygateway/middleware/telemetry.py"] + + +# —— issue #12 (a): 遥测正文可配置上限 —— + +_LONG = "甲乙丙丁戊己庚辛壬癸" * 5 # 50 字,cap=8 时省略 42 字 +_CAPPED = "甲乙丙丁戊己庚辛…(略 42 字)" + + +def _long_messages(): + """一条纯文本 + 一条多模态(text part + image_url part)。""" + return [ + {"role": "system", "content": _LONG}, + { + "role": "user", + "content": [ + {"type": "text", "text": _LONG}, + {"type": "image_url", "image_url": {"url": "https://gw.example/a.png"}}, + ], + }, + ] + + +async def _emit_with_cap(messages, *, cap, response=_LONG, thinking=_LONG): + rec = _MemoryRecorder() + await TelemetryEmitter(rec, text_cap=cap).emit_attempt( + request=ChatRequest(messages=messages, session_id="s"), + source=_source(), + call_id="c", + latency_ms=1, + response=_resp(content=response, thinking=thinking), + error=None, + ) + return rec.rows[0] + + +class TestTelemetryTextCap: + """截断发生在唯一遥测出口 `_record`(设计 §5.2);缺省 None = 不截断。""" + + async def test_cap_none_keeps_the_body_byte_for_byte(self): + """缺省不截断是人类决策(设计 §2 E-a): 落库正文与改前逐字节相同。""" + messages = _long_messages() + row = await _emit_with_cap(messages, cap=None) + assert row["messages"] == json.dumps(digest_messages(messages), ensure_ascii=False) + assert row["response"] == _LONG + assert row["thinking"] == _LONG + + async def test_cap_truncates_each_content_and_keeps_the_json_parsable(self): + """按每条文本切而非切整串 JSON: 否则该 TEXT 列此后无法按 JSON 解析。""" + row = await _emit_with_cap(_long_messages(), cap=8) + parsed = json.loads(row["messages"]) # 不抛 = 整串仍是合法 JSON + assert parsed[0]["content"] == _CAPPED + assert parsed[1]["content"][0]["text"] == _CAPPED + assert "(略 42 字)" in parsed[0]["content"] # 标记须含省略字数 + + async def test_image_digest_is_untouched_by_the_cap(self): + """多模态 image_url 的 sha256 摘要不是正文,不得被截断改形。""" + messages = _long_messages() + expected = digest_messages(messages)[1]["content"][1] + assert expected["type"] == "image_url" and len(expected["sha256"]) == 64 + row = await _emit_with_cap(messages, cap=8) + assert json.loads(row["messages"])[1]["content"][1] == expected + + async def test_non_string_content_passes_through_without_raising(self): + """外部输入形状不可控,遥测路径不得因此抛错(P5 + 降级方向)。 + + 同时钉住设计 §5.2 的覆盖面诚实声明: 只覆盖文本 content 与 text part, + 嵌套 dict 里的长文本**不在**覆盖范围内。 + """ + messages = [ + {"role": "user", "content": 123}, + {"role": "user", "content": None}, + {"role": "user", "content": {"nested": _LONG}}, + {"role": "user", "content": [{"type": "text", "text": 7}, "bare-part"]}, + ] + row = await _emit_with_cap(messages, cap=8) + assert json.loads(row["messages"]) == messages + + async def test_response_and_thinking_are_capped(self): + row = await _emit_with_cap([{"role": "user", "content": "hi"}], cap=8) + assert row["response"] == _CAPPED + assert row["thinking"] == _CAPPED + + async def test_cap_never_mutates_the_caller_messages(self): + """红线之二: 落库那份被截断,调用方持有的那份(含嵌套 part)一字未改。 + + `digest_messages` 对 content 非 list 的消息原样透传**同一个 dict 对象** + (`cache.py:43`),就地截断会连调用方的 messages、后续重试的请求体与缓存 + 写入的 key 一起改掉,且全程无任何报错。 + """ + messages = _long_messages() + snapshot = copy.deepcopy(messages) + row = await _emit_with_cap(messages, cap=8) + assert messages == snapshot + assert messages[0]["content"] == _LONG + assert messages[1]["content"][0]["text"] == _LONG + assert json.loads(row["messages"])[0]["content"] == _CAPPED # 落库那份确已截断 + + +class _StubEmbedTransport: + async def embed(self, *, texts, source, call_id): + return EmbeddingTransportResult( + vectors=[[1.0] for _ in texts], + dim=1, + prompt_tokens=1, + usage_source="measured", + raw={}, + ) + + +class _StubOcrTransport: + async def recognize_text(self, *, image, source, call_id): + return OcrTextTransportResult(text="识别结果" * 10, raw={"task_type": "text"}) + + async def parse_layout(self, *, image, source, call_id): + raise NotImplementedError + + +def _governance(scope, sources): + """embed/OCR 两条链路共用的最小治理装配(真实内存后端,不 mock)。""" + return { + "scope": scope, + "sources": sources, + "selector": RoundRobinSelector(), + "limiter": InMemoryLimiter( + scope=scope, + sources={s.name: s for s in sources}, + global_limits=GlobalLimits(max_concurrency=0, rpm=0, tpm=0), + lease_ttl_s=100.0, + ), + "breaker": InMemoryGate( + config=BreakerConfig(fail_threshold=3, cooldown_s=60.0, probe_ttl_s=120.0) + ), + "retry": RetryPolicy(max_attempts=3, backoff_base_s=0.001, backoff_max_s=0.01), + "backpressure": BackpressurePolicy(stall_window_s=300.0, poll_interval_s=0.001), + } + + +class TestTextCapCoversEmbedAndOcrChains: + """`_record` 是三条链路共同的出口,cap 自然覆盖全部三条(设计 §5.2)。 + + 同一张表不该一半受控一半不受控;embed/OCR 各自的 200 字上限保留不动, + 与新 cap 是"取更严者"的关系。 + """ + + async def test_embed_rows_are_capped(self): + rec = _MemoryRecorder() + client = EmbeddingClient( + **_governance("embed", [_source(name="e1", model="embed-1")]), + transport=_StubEmbedTransport(), + batch_size=2, + telemetry=rec, + text_cap=8, + ) + await client.embed([_LONG]) + row = rec.rows[0] + assert json.loads(row["messages"])[0]["content"] == _CAPPED + assert row["response"] == "` 共 19 字 + + async def test_ocr_rows_are_capped(self): + rec = _MemoryRecorder() + client = OcrClient( + **_governance("ocr", [_source(name="m1", model="monkey-ocr")]), + transport=_StubOcrTransport(), + telemetry=rec, + text_cap=8, + ) + await client.recognize_text(b"jpg") + row = rec.rows[0] + # 占位串 `` 共 24 字 + assert json.loads(row["messages"])[0]["content"] == "