89ff916bc8
Reasoning tokens are already counted inside completion_tokens, so the cost total was never wrong -- what was missing is the attribution: how much of a call was spent thinking rather than answering. LLMResponse and TransportResult each gain a trailing reasoning_tokens field, and the telemetry port grows from 21 to 22 columns with the new column appended in both backends so fresh and migrated schemas keep the same physical order. None means this particular call did not report the field, not that the source never reports it: a relay that falls back to a local tokenizer replaces the whole usage object and drops completion_tokens_details. Downstream checks must therefore read "in (None, 0)"; no provider was observed reporting a literal zero.
235 lines
6.4 KiB
Python
235 lines
6.4 KiB
Python
"""ports.py 端口冻结测试(M1 设计 §4): Protocol 结构性检查 + Gate 快照校验。"""
|
|
|
|
from typing import Any
|
|
|
|
import pytest
|
|
|
|
from polygateway.ports import (
|
|
CacheBackend,
|
|
GateDecision,
|
|
GateState,
|
|
GateUpdate,
|
|
Middleware,
|
|
Permit,
|
|
ProviderGate,
|
|
RateLimiter,
|
|
SourceSelector,
|
|
StructuredOutputStrategy,
|
|
TelemetryRecorder,
|
|
Transport,
|
|
)
|
|
from polygateway.types import LLMResponse, SourceStats
|
|
|
|
|
|
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):
|
|
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,
|
|
) -> 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)
|
|
|
|
|
|
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 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)
|