"""`usage_source` 值域封闭: 库内所有生产点的产出恒落在 `USAGE_SOURCES` 内。 设计 §3.1 裁定值域**只约束生产侧**——公共 frozen dataclass 不加运行时校验 (裸 `ValueError` 不属四分类,会逃出 `chat()`;该裁决的锁定断言在 `test_types.py::TestUsageSourceDomain`)。因此封闭性只能由"逐个驱动生产点、 断言其产出在三态内"来保证,本文件即该断言的载体。 独立成文件而非并入 `test_types.py`: 断言横跨 transports / embedding / telemetry 三层,放进最内层内核的类型测试会让它反向依赖具体实现。 覆盖的生产点(设计 §3.2 逐处改动表的字面量产出方): `_resolve_usage`、`_resolve_embedding_usage`、`_resolve_stream_usage`(打捞覆盖)、 `EmbeddingClient._merge`、`EmbeddingClient.embed` 空输入短路、 `OcrClient._emit`、`TelemetryEmitter.emit_attempt/emit_cache_hit/emit_terminal_failure`。 """ import itertools import json import httpx import pytest from polygateway.backends.memory.breaker import InMemoryGate from polygateway.backends.memory.limiter import InMemoryLimiter from polygateway.embedding import EmbeddingClient, _BatchOutcome from polygateway.middleware.telemetry import TelemetryEmitter from polygateway.ocr import OcrClient from polygateway.sources import RoundRobinSelector from polygateway.transports.openai_compat import ( OpenAICompatTransport, _resolve_embedding_usage, _resolve_usage, ) from polygateway.types import ( USAGE_SOURCES, BackpressurePolicy, BreakerConfig, ChatRequest, EmbeddingTransportResult, GlobalLimits, LLMResponse, OcrTextTransportResult, RetryPolicy, SourceConfig, ) _REQ = ChatRequest(messages=[{"role": "user", "content": "hi"}]) _DOMAIN = sorted(USAGE_SOURCES) def _src(): return SourceConfig( name="s1", provider="p", base_url="https://gw.example/v1", api_key="sk", model="m", timeout_s=10.0, est_tokens=4000, # 兜底口径的历史来源: 生产点不得因它落到三态之外 ) class _MemoryRecorder: def __init__(self): self.rows = [] async def record_llm_call(self, **fields): self.rows.append(fields) @pytest.mark.parametrize( "usage", [ {"prompt_tokens": 12, "completion_tokens": 34}, # 完整可信 {}, # 整帧缺失 {"prompt_tokens": 0, "completion_tokens": 0}, # 全 0(和不为正) {"prompt_tokens": "12", "completion_tokens": 34}, # 类型非法 {"prompt_tokens": None, "completion_tokens": None}, {"prompt_tokens": 12}, # 半帧 ], ) def test_resolve_usage_stays_in_domain(usage): assert _resolve_usage(usage)[2] in USAGE_SOURCES @pytest.mark.parametrize( "data", [ {"usage": {"prompt_tokens": 12}}, {}, {"usage": None}, {"usage": {}}, {"usage": {"prompt_tokens": 0}}, {"usage": {"prompt_tokens": "12"}}, ], ) def test_resolve_embedding_usage_stays_in_domain(data): assert _resolve_embedding_usage(data)[1] in USAGE_SOURCES def _sse(*frames, done): """构造 SSE 响应;done=False 触发打捞路径(`_complete_stream` 的覆盖分支)。""" text = "".join(f"data: {json.dumps(f)}\n\n" for f in frames) + ( "data: [DONE]\n\n" if done else "" ) return httpx.Response(200, content=text.encode(), headers={"content-type": "text/event-stream"}) @pytest.mark.parametrize("usage", [{"prompt_tokens": 11, "completion_tokens": 7}, None]) async def test_salvage_override_stays_in_domain(usage): """打捞覆盖(`openai_compat._complete_stream`)是第三个字面量产出方。""" frames = [{"choices": [{"delta": {"content": "partial"}}]}] if usage is not None: frames.append({"choices": [], "usage": usage}) transport = OpenAICompatTransport( client_factory=lambda src: httpx.AsyncClient( base_url=src.base_url, transport=httpx.MockTransport(lambda request: _sse(*frames, done=False)), ) ) source = SourceConfig( name="s1", provider="qwen", base_url="https://gw.example/v1", api_key="sk", model="m", timeout_s=10.0, est_tokens=4000, missing_done="salvage", ) result = await transport.complete( messages=[{"role": "user", "content": "hi"}], source=source, stream=True, overlay={}, call_id="cid", ) assert result.usage_source in USAGE_SOURCES class _ScriptedOcrTransport: async def recognize_text(self, *, image, source, call_id): return OcrTextTransportResult(text="LINE-1", raw={"task_type": "text"}) async def parse_layout(self, *, image, source, call_id): raise NotImplementedError async def check_health(self, *, source): raise NotImplementedError async def test_ocr_emit_stays_in_domain(): """`OcrClient._emit` 的字面量(ocr.py:411)同样纳入封闭性断言。 值取 `measured` 是设计 §3.3 的裁决(OCR 的 0 token 属事实);此处只断言 落在三态内,精确取值的防回归钉在 `test_ocr_client.py`。 """ source = SourceConfig( name="m1", provider="monkey", base_url="http://gw.example", api_key="none", model="monkey-ocr", timeout_s=10.0, ) recorder = _MemoryRecorder() client = OcrClient( scope="ocr", sources=[source], selector=RoundRobinSelector(), limiter=InMemoryLimiter( scope="ocr", sources={source.name: source}, 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) ), transport=_ScriptedOcrTransport(), retry=RetryPolicy(max_attempts=1, backoff_base_s=0.001, backoff_max_s=0.01), backpressure=BackpressurePolicy(stall_window_s=300.0, poll_interval_s=0.001), telemetry=recorder, ) await client.recognize_text(b"jpg") assert recorder.rows[0]["usage_source"] in USAGE_SOURCES def _merge_client(): """构造仅用于调用 `_merge` 的最小 EmbeddingClient(不发起任何调用)。""" source = _src() return EmbeddingClient( scope="embed", sources=[source], selector=RoundRobinSelector(), limiter=InMemoryLimiter( scope="embed", sources={source.name: source}, 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) ), transport=object(), retry=RetryPolicy(max_attempts=1, backoff_base_s=0.001, backoff_max_s=0.01), backpressure=BackpressurePolicy(stall_window_s=300.0, poll_interval_s=0.001), batch_size=2, ) @pytest.mark.parametrize(("first", "second"), list(itertools.product(_DOMAIN, repeat=2))) def test_merge_stays_in_domain(first, second): """任意两批 usage_source 组合(含尚无生产者的 unavailable)合并后仍在三态内。""" source = _src() outcomes = [ _BatchOutcome( result=EmbeddingTransportResult( vectors=[[1.0]], dim=1, prompt_tokens=1, usage_source=value, raw={} ), source=source, call_id="c", latency_ms=1, ) for value in (first, second) ] assert _merge_client()._merge(outcomes).usage_source in USAGE_SOURCES async def test_empty_input_short_circuit_stays_in_domain(): """空输入短路自造响应(embedding.py:151),不经 transport 也须落在三态内。""" resp = await _merge_client().embed([]) assert resp.usage_source in USAGE_SOURCES def _resp(usage_source): return LLMResponse( content="ok", thinking="", model="m", provider="p", prompt_tokens=1, completion_tokens=2, latency_ms=30, ttft_ms=None, max_inter_token_ms=None, cache_hit=False, call_id="cid", source_name="s1", usage_source=usage_source, ) @pytest.mark.parametrize("emitted", _DOMAIN) async def test_emit_attempt_success_stays_in_domain(emitted): recorder = _MemoryRecorder() await TelemetryEmitter(recorder).emit_attempt( request=_REQ, source=_src(), call_id="cid", latency_ms=10, response=_resp(emitted), error=None, ) assert recorder.rows[0]["usage_source"] in USAGE_SOURCES async def test_emit_attempt_failed_attempt_stays_in_domain(): """失败尝试无 response,`usage_source` 取 emitter 自己的字面量。""" recorder = _MemoryRecorder() await TelemetryEmitter(recorder).emit_attempt( request=_REQ, source=_src(), call_id="cid", latency_ms=10, response=None, error="boom", ) assert recorder.rows[0]["usage_source"] in USAGE_SOURCES @pytest.mark.parametrize("emitted", _DOMAIN) async def test_emit_cache_hit_stays_in_domain(emitted): recorder = _MemoryRecorder() await TelemetryEmitter(recorder).emit_cache_hit(request=_REQ, response=_resp(emitted)) assert recorder.rows[0]["usage_source"] in USAGE_SOURCES async def test_emit_terminal_failure_stays_in_domain(): """终态失败无具体源,`usage_source` 同样取 emitter 字面量。""" recorder = _MemoryRecorder() await TelemetryEmitter(recorder).emit_terminal_failure( request=_REQ, call_id="cid", latency_ms=10, error="cancelled" ) assert recorder.rows[0]["usage_source"] in USAGE_SOURCES