Files
PolyGateway/tests/unit/test_ports.py
T
iomgaa e06cd8e8b7 feat: record which tier a call actually ran at
Twenty-five columns and not one of them answered "which tier was this?",
so the question the whole issue exists to settle - does a higher tier buy
anything - had no way to group its data.

The three emit entry points deliberately disagree, the way sampling
already does. A successful attempt records what the transport actually
sent: with EFFORT_FALLBACK=nearest a request for medium goes out as low,
and recomputing here would file the row under a tier that never left the
process. A failed attempt has no response to read, so it falls back to
the requested tier - which is exactly right for the tier errors that are
rejected before any HTTP happens, because the rejected tier is the
signal. Cache hits and terminal failures have no chosen source at all,
so a source-level tier is not a thing they could report.

emit_attempt now demands to be told whether the path reasons at all.
Embedding and OCR share the emitter but never send reasoning parameters;
without the flag a source that mistakenly carries ENABLE_THINKING would
hang a tier on a call that could not possibly have run at one.

The value lands as a plain str. StrEnum is a str subclass and asyncpg
promises nothing about encoding subclasses, and a telemetry write that
fails is only a warning - Postgres would just quietly lose the column.
NULL means nobody declared a tier, which is not the same statement as
'none', and the two must never be folded together.
2026-09-05 05:57:29 -04:00

316 lines
9.7 KiB
Python

"""ports.py 端口冻结测试(M1 设计 §4): Protocol 结构性检查 + Gate 快照校验。"""
import inspect
from typing import Any
import pytest
from polygateway.ports import (
CacheBackend,
EmbeddingTransport,
GateDecision,
GateState,
GateUpdate,
Middleware,
OcrTransport,
Permit,
ProviderGate,
RateLimiter,
SourceSelector,
StructuredOutputStrategy,
TelemetryRecorder,
TelemetryStatusProvider,
Transport,
)
from polygateway.types import LLMResponse, SourceStats, TelemetryStatus
def _resp() -> LLMResponse:
return LLMResponse("c", "t", "m", "p", 1, 2, 3, None, None, False, "cid")
class _DummyPermit:
async def release(self) -> None: ...
async def settle(self, actual_tokens: int) -> None: ...
class _DummyLimiter:
async def try_acquire(self, source_key: str, est_tokens: int):
return _DummyPermit()
async def acquire(self, source_key: str, est_tokens: int):
return _DummyPermit()
async def source_stats(self, source_key: str):
return SourceStats(0, 0, 0)
async def mark_progress(self) -> None: ...
async def progress_age_s(self) -> float:
return 0.0
class _DummyGate:
async def try_enter(self, source_name: str, owner: str):
raise NotImplementedError
async def record_success(self, entry):
raise NotImplementedError
async def record_failure(self, entry, reason: str, force_open: bool):
raise NotImplementedError
async def release_probe(self, entry):
raise NotImplementedError
async def retry_after_s(self, sources):
return 0.0
class _DummyMw:
async def __call__(self, request, call_next):
return await call_next(request)
class _DummyTransport:
async def complete(self, *, messages, source, stream, overlay, call_id, reasoning_effort):
raise NotImplementedError
class _DummyCache:
async def get(self, key: str):
return None
async def set(self, key: str, value: str, ttl_s: int) -> None: ...
class _DummySelector:
def order(self, sources, stats):
return list(sources)
class _DummyStrategy:
def request_overlay(self, schema):
return {}
def parse(self, text: str) -> Any:
return {}
class _DummyRecorder:
async def record_llm_call(
self,
*,
call_id,
parent_call_id,
session_id,
model,
provider,
source_name,
messages,
response,
thinking,
prompt_tokens,
completion_tokens,
usage_source,
latency_ms,
ttft_ms,
max_inter_token_ms,
cache_hit,
error,
cost,
cached_prompt_tokens,
model_reported,
sampling,
reasoning_tokens,
tenant_id,
meta,
) -> None: ...
@pytest.mark.parametrize(
("impl", "protocol"),
[
(_DummyPermit(), Permit),
(_DummyLimiter(), RateLimiter),
(_DummyGate(), ProviderGate),
(_DummyMw(), Middleware),
(_DummyTransport(), Transport),
(_DummyCache(), CacheBackend),
(_DummySelector(), SourceSelector),
(_DummyStrategy(), StructuredOutputStrategy),
(_DummyRecorder(), TelemetryRecorder),
],
)
def test_protocols_are_runtime_checkable(impl, protocol):
assert isinstance(impl, protocol)
class TestReasoningTierIsOnlyOnTheChatPort:
"""档位属于 chat 端口,且**只属于**它(Task 5b)。
`@runtime_checkable` 只查方法名不查签名,故协议签名本身必须被显式断言——
否则实现漏改一个参数,要到运行期调用才会以 `TypeError` 现形,而那时的现场
离根因已经很远。
"""
def test_chat_transport_carries_the_per_call_tier(self):
params = inspect.signature(Transport.complete).parameters
assert "reasoning_effort" in params
# 不给默认值是有意的(与 TelemetryRecorder 同一既有约定): 库外无第三方
# 实现者,写全签名成本为零,而默认值会把"漏传"变成静默的"不表态"
assert params["reasoning_effort"].default is inspect.Parameter.empty
@pytest.mark.parametrize(
("protocol", "method"),
[
(EmbeddingTransport, "embed"),
(OcrTransport, "recognize_text"),
(OcrTransport, "parse_layout"),
],
)
def test_other_transports_have_no_reasoning_tier(self, protocol, method):
"""embedding 与 OCR 没有推理语义,给它们加档位只会静默无效(issue #4 同款决策)。"""
assert "reasoning_effort" not in inspect.signature(getattr(protocol, method)).parameters
class _DummyStatusProvider(_DummyRecorder):
@property
def telemetry_status(self) -> TelemetryStatus:
return TelemetryStatus(
degraded=False,
fatal=False,
reason=None,
degraded_for_s=None,
dropped_rows=0,
retry_after_s=None,
)
def test_status_provider_is_a_separate_optional_port():
"""状态**不得**并进 TelemetryRecorder: 那会让只实现 record_llm_call 的对象
当场不再满足 @runtime_checkable 的结构检查(设计 §3.3,Codex 审查)。"""
assert isinstance(_DummyStatusProvider(), TelemetryStatusProvider)
assert isinstance(_DummyStatusProvider(), TelemetryRecorder)
assert not isinstance(_DummyRecorder(), TelemetryStatusProvider)
assert isinstance(_DummyRecorder(), TelemetryRecorder) # 这条断言是那条决策的执法点
def _decision(**overrides) -> GateDecision:
base = {
"source_name": "qwen_1",
"allowed": True,
"state": GateState.CLOSED,
"epoch": 0,
"is_probe": False,
"probe_owner": None,
"retry_after_s": 0.0,
}
base.update(overrides)
return GateDecision(**base)
class TestGateDecisionInvariants:
"""校验逐条移植 CHS ports.py:405-440。"""
def test_valid_probe_decision(self):
d = _decision(state=GateState.HALF_OPEN, is_probe=True, probe_owner="w1")
assert d.is_probe and d.probe_owner == "w1"
def test_empty_source_rejected(self):
with pytest.raises(ValueError):
_decision(source_name=" ")
def test_negative_epoch_and_retry_after_rejected(self):
with pytest.raises(ValueError):
_decision(epoch=-1)
with pytest.raises(ValueError):
_decision(retry_after_s=-0.1)
def test_open_state_cannot_allow(self):
with pytest.raises(ValueError):
_decision(state=GateState.OPEN, allowed=True)
def test_half_open_admission_must_be_probe(self):
with pytest.raises(ValueError):
_decision(state=GateState.HALF_OPEN, is_probe=False)
def test_probe_requires_owner_and_half_open(self):
with pytest.raises(ValueError):
_decision(state=GateState.HALF_OPEN, is_probe=True, probe_owner=None)
with pytest.raises(ValueError):
_decision(state=GateState.CLOSED, is_probe=True, probe_owner="w1")
def test_non_probe_cannot_carry_owner(self):
with pytest.raises(ValueError):
_decision(probe_owner="w1")
class TestGateUpdate:
def test_bounds(self):
u = GateUpdate(
applied=True, state=GateState.CLOSED, epoch=0, failure_count=0, retry_after_s=0.0
)
assert u.applied
with pytest.raises(ValueError):
GateUpdate(
applied=True, state=GateState.CLOSED, epoch=-1, failure_count=0, retry_after_s=0.0
)
with pytest.raises(ValueError):
GateUpdate(
applied=True, state=GateState.CLOSED, epoch=0, failure_count=-1, retry_after_s=0.0
)
class TestTelemetryRecorderSignature:
"""`record_llm_call` 的冻结签名以 `inspect.signature` 实测,不凭记忆断言。
该 Protocol 的纪律是新增参数**不设默认值**(ports.py docstring):库外无第三方
实现者,而带默认值的参数会让 emitter 漏传时静默落默认值——遥测里的租户归属
一旦静默错位,事后无从分辨是"没传"还是"就是空的"。
"""
def test_caller_dimensions_are_declared(self):
import inspect
params = inspect.signature(TelemetryRecorder.record_llm_call).parameters
assert {"tenant_id", "meta"} <= set(params)
@pytest.mark.parametrize(
"name", ["tenant_id", "meta", "thinking_observation", "reasoning_effort"]
)
def test_caller_dimensions_have_no_default(self, name):
import inspect
param = inspect.signature(TelemetryRecorder.record_llm_call).parameters[name]
assert param.default is inspect.Parameter.empty
assert param.kind is inspect.Parameter.KEYWORD_ONLY
class TestOcrPorts:
"""M3 三个 OCR Protocol(设计 §3.2): runtime_checkable 结构判定。"""
def test_ocr_ports_runtime_checkable(self):
from polygateway.ports import OcrLayoutPort, OcrTextPort, OcrTransport
class GoodClient:
async def recognize_text(self, image): ...
async def parse_layout(self, image): ...
class GoodTransport:
async def recognize_text(self, *, image, source, call_id): ...
async def parse_layout(self, *, image, source, call_id): ...
async def check_health(self, *, source): ...
assert isinstance(GoodClient(), OcrTextPort)
assert isinstance(GoodClient(), OcrLayoutPort)
assert isinstance(GoodTransport(), OcrTransport)
def test_missing_method_rejected(self):
from polygateway.ports import OcrTransport
class NoHealth:
async def recognize_text(self, *, image, source, call_id): ...
async def parse_layout(self, *, image, source, call_id): ...
assert not isinstance(NoHealth(), OcrTransport)