From 89e86448b07fda1be1c760404e05aa6963755633 Mon Sep 17 00:00:00 2001 From: iomgaa Date: Tue, 21 Jul 2026 22:55:03 -0400 Subject: [PATCH] feat: add OcrClient governed loop with dual ports --- src/polygateway/ocr.py | 506 ++++++++++++++++++++++++++++++++++ tests/unit/test_ocr_client.py | 390 ++++++++++++++++++++++++++ 2 files changed, 896 insertions(+) create mode 100644 src/polygateway/ocr.py create mode 100644 tests/unit/test_ocr_client.py diff --git a/src/polygateway/ocr.py b/src/polygateway/ocr.py new file mode 100644 index 0000000..0a7ebd5 --- /dev/null +++ b/src/polygateway/ocr.py @@ -0,0 +1,506 @@ +"""OcrClient: 治理化 OCR 调用(M3 设计 §5,方案 A)。 + +独立精简治理循环,**复用**库的算法件: `RateLimiter`/`ProviderGate` 端口与 +两种后端、错误四分类、`backoff_delay` 退避公式、`SourceCooldownMemo`、 +`TelemetryEmitter`(遥测单一 helper 铁律)、选源器(含 OutcomeAwareSelector +喂数)。循环与 EmbeddingClient 同构——设计 §2.A 已声明的第三份有限重复 +(chat 循环的 AIMD/429 免预算/流式看门狗/缓存均不适用于 OCR)。 + +与 embedding 循环的有意差异(设计 §5): ① settle 恒为 0(OCR 无 token +计费,失败也不按 est 保守结算);② 无批处理外循环;③ 接 +OutcomeAwareSelector 喂数(多实例 LAN 服务单机可挂,健康选源正为此设计)。 +""" + +from __future__ import annotations + +import asyncio +import random +import time +import uuid +from dataclasses import dataclass +from typing import TYPE_CHECKING, Literal + +from loguru import logger + +from polygateway.errors import ( + AllSourcesExhausted, + CircuitOpenError, + GovernanceBackendError, + PolyGatewayError, + RequestRejectedError, + ResultInvalidError, + SourceDeadError, + TransientError, +) +from polygateway.middleware.breaker import BreakerGate +from polygateway.middleware.ratelimit import QuotaGate +from polygateway.middleware.retry import _failure_reason, backoff_delay +from polygateway.middleware.telemetry import TelemetryEmitter +from polygateway.ports import OutcomeAwareSelector +from polygateway.sources import SourceCooldownMemo +from polygateway.types import ( + ChatRequest, + LLMResponse, + OcrLayoutResult, + OcrTextResult, + Usage, +) + +if TYPE_CHECKING: + from collections.abc import Awaitable, Callable, Mapping + + from polygateway.config import OcrSettings + from polygateway.ports import ( + GateDecision, + OcrTransport, + Permit, + ProviderGate, + RateLimiter, + SourceSelector, + TelemetryRecorder, + ) + from polygateway.types import ( + BackpressurePolicy, + OcrLayoutTransportResult, + OcrTextTransportResult, + RetryPolicy, + SourceConfig, + ) + +_RESPONSE_TEXT_CAP = 200 # 遥测行 text 截断长度(与 embedding 口径一致) + +_OcrKind = Literal["text", "layout"] + + +@dataclass(frozen=True) +class _FailedAttempt: + exc: PolyGatewayError + immediate: bool + + +@dataclass(frozen=True) +class _AttemptOutcome: + result: OcrTextTransportResult | OcrLayoutTransportResult + source: SourceConfig + call_id: str + latency_ms: int + + +class OcrClient: + """治理化 OCR 入口: 同时实现 OcrTextPort 与 OcrLayoutPort(D9 端口族)。 + + 两端点打同一服务实例池,共享同一 scope 的限流/熔断/选源状态; + 与其他 scope 共享后端实例即共享全局闸(显式传入,禁止隐式全局)。 + """ + + def __init__( + self, + *, + scope: str, + sources: list[SourceConfig], + selector: SourceSelector, + limiter: RateLimiter, + breaker: ProviderGate, + transport: OcrTransport, + retry: RetryPolicy, + backpressure: BackpressurePolicy, + quota_full: str = "wait", + telemetry: TelemetryRecorder | None = None, + now: Callable[[], float] = time.monotonic, + sleep: Callable[[float], Awaitable[None]] = asyncio.sleep, + rng: Callable[[], float] = random.random, + ) -> None: + if quota_full not in ("wait", "fail_fast"): + raise ValueError(f"quota_full 必须是 wait|fail_fast: {quota_full!r}") + self._scope = scope + self._sources = list(sources) + self._selector = selector + self._feed_health = isinstance(selector, OutcomeAwareSelector) + self._quota = QuotaGate(limiter) + self._breaker = BreakerGate(breaker) + self._transport = transport + self._retry = retry + self._bp = backpressure + self._quota_full = quota_full + self._emitter = TelemetryEmitter(telemetry) if telemetry else None + self._telemetry = telemetry + self._memo = SourceCooldownMemo(now=now) + self._now = now + self._sleep = sleep + self._rng = rng + self._closed = False + + # —— 公共端口(OcrTextPort / OcrLayoutPort)—— + + async def recognize_text( + self, + image: bytes, + *, + session_id: str | None = None, + parent_call_id: str | None = None, + ) -> OcrTextResult: + """一次治理文本转录(/ocr/text);text 空串 = 合法"无文字"。""" + outcome = await self._call("text", image, session_id, parent_call_id) + result = outcome.result + return OcrTextResult( + text=result.text, + source_name=outcome.source.name, + usage=Usage(0, 0), + latency_ms=outcome.latency_ms, + call_id=outcome.call_id, + raw=result.raw, + ) + + async def parse_layout( + self, + image: bytes, + *, + session_id: str | None = None, + parent_call_id: str | None = None, + ) -> OcrLayoutResult: + """一次治理版面解析(/parse → ZIP);elements 空 = 合法"无元素"。""" + outcome = await self._call("layout", image, session_id, parent_call_id) + result = outcome.result + return OcrLayoutResult( + elements=result.elements, + page_sizes=result.page_sizes, + source_name=outcome.source.name, + usage=Usage(0, 0), + latency_ms=outcome.latency_ms, + call_id=outcome.call_id, + raw=result.raw, + ) + + async def check_health(self) -> dict[str, bool]: + """逐源并发健康预检(R10);探测失败=False 不上抛,取消穿透。""" + names = [s.name for s in self._sources] + results = await asyncio.gather( + *(self._transport.check_health(source=s) for s in self._sources) + ) + health = dict(zip(names, results, strict=True)) + for name, ok in health.items(): + if not ok: + logger.warning("OCR 源 {} 健康预检未通过", name) + return health + + # —— 治理循环(与 EmbeddingClient._embed_batch 同构;设计 §2.A)—— + + async def _call( + self, + kind: _OcrKind, + image: bytes, + session_id: str | None, + parent_call_id: str | None, + ) -> _AttemptOutcome: + if not isinstance(image, bytes): + raise TypeError("image 必须是 bytes(路径读取/批量拼帧留业务侧,D9)") + if not image: + raise ValueError("image 不能为空") + if not self._sources: + raise AllSourcesExhausted(scope=self._scope, reason="no_sources", retry_after_s=0.0) + fails = 0 + reasons: dict[str, str] = {} + entered_at = self._now() + while True: + picked, gate_rejections = await self._pick_runnable(reasons) + if picked is None: + await self._on_no_runnable(gate_rejections, reasons, entered_at) + continue + outcome = await self._attempt(kind, image, *picked, reasons, session_id, parent_call_id) + if isinstance(outcome, _AttemptOutcome): + return outcome + fails += 1 + if fails >= self._retry.max_attempts: + raise AllSourcesExhausted( + scope=self._scope, + reason="retry_exhausted", + retry_after_s=self._retry.backoff_base_s, + per_source_reasons=reasons, + ) from outcome.exc + if not outcome.immediate: + await self._sleep(backoff_delay(self._retry, fails, outcome.exc, self._rng)) + + async def _pick_runnable( + self, reasons: dict[str, str] + ) -> tuple[tuple[SourceConfig, Permit, GateDecision] | None, int]: + stats = {s.name: await self._quota.stats(s) for s in self._sources} + gate_rejections = 0 + for cand in self._selector.order(self._sources, stats): + if self._memo.active(cand.name): + gate_rejections += 1 + reasons[cand.name] = "cooldown" + continue + permit = await self._quota.try_acquire(cand) + if permit is None: + reasons.setdefault(cand.name, "rate_limited") + continue + entry = None + try: + entry = await self._breaker.try_enter(cand, uuid.uuid4().hex) + finally: + if entry is None: + await self._settle_and_release(permit) + if entry.allowed: + return (cand, permit, entry), gate_rejections + gate_rejections += 1 + reasons[cand.name] = "circuit_open" + self._memo.set_until(cand.name, self._now() + entry.retry_after_s) + await self._settle_and_release(permit) + return None, gate_rejections + + async def _on_no_runnable( + self, gate_rejections: int, reasons: dict[str, str], entered_at: float + ) -> None: + if gate_rejections == len(self._sources): + names = tuple(s.name for s in self._sources) + raise CircuitOpenError( + scope=self._scope, + retry_after_s=await self._breaker.retry_after_s(names), + per_source_reasons=reasons, + ) + if self._quota_full == "fail_fast": + raise AllSourcesExhausted( + scope=self._scope, + reason="quota_exhausted", + retry_after_s=self._bp.poll_interval_s, + per_source_reasons=reasons, + ) + stall = self._bp.stall_window_s + if self._now() - entered_at > stall and await self._quota.progress_age_s() > stall: + names = tuple(s.name for s in self._sources) + raise AllSourcesExhausted( + scope=self._scope, + reason="stalled", + retry_after_s=await self._breaker.retry_after_s(names), + per_source_reasons=reasons, + ) + await self._sleep(self._bp.poll_interval_s * (0.5 + 0.5 * self._rng())) + + async def _attempt( + self, + kind: _OcrKind, + image: bytes, + source: SourceConfig, + permit: Permit, + entry: GateDecision, + reasons: dict[str, str], + session_id: str | None, + parent_call_id: str | None, + ) -> _AttemptOutcome | _FailedAttempt: + call_id = str(uuid.uuid4()) + started = self._now() + try: + result = await self._invoke(kind, image, source, call_id) + await self._record_quietly(self._breaker.record_success(entry)) + await self._record_quietly(self._quota.mark_progress()) + self._feed_outcome(source.name, ok=True) + latency_ms = int((self._now() - started) * 1000) + await self._emit(kind, image, source, call_id, started, session_id, parent_call_id, result) + return _AttemptOutcome(result, source, call_id, latency_ms) + except (RequestRejectedError, ResultInvalidError) as exc: + await self._gate_on_terminal(exc, entry) + await self._emit(kind, image, source, call_id, started, session_id, parent_call_id, error=exc) + raise + except asyncio.CancelledError: + if entry.is_probe: + await self._record_quietly(self._breaker.release_probe(entry)) + await self._emit( + kind, image, source, call_id, started, session_id, parent_call_id, error="cancelled" + ) + raise + except (SourceDeadError, TransientError) as exc: + dead = isinstance(exc, SourceDeadError) + reason = _failure_reason(exc) + reasons[source.name] = reason + await self._record_quietly(self._breaker.record_failure(entry, reason, dead)) + self._feed_outcome(source.name, ok=False) + await self._emit(kind, image, source, call_id, started, session_id, parent_call_id, error=exc) + return _FailedAttempt(exc, immediate=dead) + finally: + await self._settle_and_release(permit) + + async def _invoke( + self, kind: _OcrKind, image: bytes, source: SourceConfig, call_id: str + ) -> OcrTextTransportResult | OcrLayoutTransportResult: + if kind == "text": + return await self._transport.recognize_text(image=image, source=source, call_id=call_id) + return await self._transport.parse_layout(image=image, source=source, call_id=call_id) + + # —— 辅助(与 embedding 同口径)—— + + async def _gate_on_terminal(self, exc: PolyGatewayError, entry: GateDecision) -> None: + """终态异常门控写回: 坏结果/服务的业务拒绝(有 HTTP 响应)≠ 坏服务 + → 记成功不计窗口样本;服务没响应的本地拒绝若持探针则归还。""" + if isinstance(exc, ResultInvalidError) or exc.status_code is not None: + await self._record_quietly(self._breaker.record_success(entry, count_attempt=False)) + elif entry.is_probe: + await self._record_quietly(self._breaker.release_probe(entry)) + + def _feed_outcome(self, source_name: str, *, ok: bool) -> None: + """健康喂数(M2.5 口径): 真实成败喂,ResultInvalid/健康拒绝不喂; + 选源器异常吞并降级 warning(喂数失败不得影响调用)。""" + if not self._feed_health: + return + try: + self._selector.record_outcome(source_name, ok) + except Exception as exc: # noqa: BLE001 — 喂数侧故障降级,不冒泡 + logger.warning("OCR 健康喂数失败(不冒泡): {}", exc) + + async def _record_quietly(self, write_back: Awaitable[object]) -> None: + try: + await write_back + except asyncio.CancelledError: + raise + except GovernanceBackendError as exc: + logger.warning("OCR 治理记账写回降级(不冒泡): {}", exc) + + async def _settle_and_release(self, permit: Permit) -> None: + """settle 恒 0: OCR 无 token 计费(设计 §5 差异①)。""" + try: + try: + await permit.settle(0) + finally: + await permit.release() + except asyncio.CancelledError: + raise + except Exception as exc: + logger.warning("OCR permit 结算/释放失败(不掩盖主异常): {}", exc) + + async def _emit( + self, + kind: _OcrKind, + image: bytes, + source: SourceConfig, + call_id: str, + started: float, + session_id: str | None, + parent_call_id: str | None, + result: OcrTextTransportResult | OcrLayoutTransportResult | None = None, + error: object | None = None, + ) -> None: + """逐尝试遥测(单一 Emitter): messages 占位摘要,图像 bytes 绝不入库。""" + if self._emitter is None: + return + request = ChatRequest( + messages=[{"role": "user", "content": f""}], + session_id=session_id, + parent_call_id=parent_call_id, + ) + latency_ms = int((self._now() - started) * 1000) + response = None + if result is not None: + response = LLMResponse( + content=self._summarize(result), + thinking="", + model=source.model, + provider=source.provider, + prompt_tokens=0, + completion_tokens=0, + latency_ms=latency_ms, + ttft_ms=None, + max_inter_token_ms=None, + cache_hit=False, + call_id=call_id, + source_name=source.name, + usage_source="measured", + ) + await self._emitter.emit_attempt( + request=request, + source=source, + call_id=call_id, + latency_ms=latency_ms, + response=response, + error=None if error is None else str(error), + ) + + @staticmethod + def _summarize(result: OcrTextTransportResult | OcrLayoutTransportResult) -> str: + if hasattr(result, "text"): + return result.text[:_RESPONSE_TEXT_CAP] + return f"" + + # —— 生命周期 —— + + async def aclose(self) -> None: + """幂等释放 transport 连接池与遥测连接(与 EmbeddingClient 对称)。""" + if self._closed: + return + self._closed = True + transport_aclose = getattr(self._transport, "aclose", None) + if transport_aclose is not None: + await transport_aclose() + telemetry_aclose = getattr(self._telemetry, "aclose", None) + if telemetry_aclose is not None: + await telemetry_aclose() + else: + telemetry_close = getattr(self._telemetry, "close", None) + if telemetry_close is not None: + telemetry_close() + + async def __aenter__(self) -> OcrClient: + return self + + async def __aexit__(self, *exc_info: object) -> None: + await self.aclose() + + # —— 工厂(与 EmbeddingClient 对称)—— + + @classmethod + def from_settings( + cls, + settings: OcrSettings, + *, + limiter: RateLimiter | None = None, + breaker: ProviderGate | None = None, + telemetry: TelemetryRecorder | None = None, + ) -> OcrClient: + """按配置装配;显式传入的后端实例即共享(与其他 scope 共享全局闸)。""" + from polygateway.client import ( + _build_breaker, + _build_limiter, + _build_selector, + _build_telemetry, + ) + from polygateway.transports.monkey_ocr import MonkeyOcrTransport + + gw = settings.gateway + sources = list(gw.sources) + # 装配防御(D9 GLM 预留档): 配了非 monkey 源必须失败, + # 严禁静默用 MonkeyOcrTransport 打别家端点(默认值掩盖错误) + alien = sorted({s.provider for s in sources if s.provider != "monkey"}) + if alien: + raise ValueError( + f"OCR 装配仅支持 provider=monkey(D9 其余后端预留未实现): 发现 {alien}" + ) + return cls( + scope=gw.scope, + sources=sources, + selector=_build_selector(gw.selector), + limiter=limiter or _build_limiter(gw, sources), + breaker=breaker or _build_breaker(gw), + transport=MonkeyOcrTransport(), + retry=gw.retry, + backpressure=gw.backpressure, + quota_full=gw.quota_full, + telemetry=telemetry if telemetry is not None else _build_telemetry(gw), + ) + + @classmethod + def from_env( + cls, + scope: str = "OCR", + *, + limiter: RateLimiter | None = None, + breaker: ProviderGate | None = None, + telemetry: TelemetryRecorder | None = None, + env: Mapping[str, str] | None = None, + ) -> OcrClient: + """从 .env/环境变量装配一个 OCR scope 的 client。""" + from polygateway.config import OcrSettings + + return cls.from_settings( + OcrSettings.from_env(scope, env=env), + limiter=limiter, + breaker=breaker, + telemetry=telemetry, + ) diff --git a/tests/unit/test_ocr_client.py b/tests/unit/test_ocr_client.py new file mode 100644 index 0000000..a81caa8 --- /dev/null +++ b/tests/unit/test_ocr_client.py @@ -0,0 +1,390 @@ +"""OcrClient 治理循环测试(M3 设计 §5): 换源/熔断口径/stall/取消/G1 契约。 + +循环与 EmbeddingClient 同构(设计 §2.A 有限重复裁决);本文件的 +retry_exhausted/circuit_open/stalled 三组断言即设计 §6 ③ 的 G1 契约钉。 +""" + +import asyncio + +import pytest + +from polygateway.backends.memory.breaker import InMemoryGate +from polygateway.backends.memory.limiter import InMemoryLimiter +from polygateway.errors import ( + AllSourcesExhausted, + CircuitOpenError, + RequestRejectedError, + ResultInvalidError, + SourceDeadError, + TransientError, +) +from polygateway.ocr import OcrClient +from polygateway.types import ( + BackpressurePolicy, + BreakerConfig, + GlobalLimits, + OcrLayoutElement, + OcrLayoutTransportResult, + OcrTextTransportResult, + RetryPolicy, + SourceConfig, +) + +_BREAKER = BreakerConfig(fail_threshold=3, cooldown_s=60.0, probe_ttl_s=120.0) +_NO_GLOBAL = GlobalLimits(max_concurrency=0, rpm=0, tpm=0) +_TEXT_OK = OcrTextTransportResult(text="LINE-1", raw={"task_type": "text"}) +_LAYOUT_OK = OcrLayoutTransportResult( + elements=[OcrLayoutElement(type="table", bbox=(41.0, 48.0, 218.0, 282.0), page_index=0)], + page_sizes=[(759.0, 540.0)], + raw={"success": True}, +) + + +def _src(**overrides): + base = { + "name": "m1", + "provider": "monkey", + "base_url": "http://10.77.0.20:7866", + "api_key": "none", + "model": "monkey-ocr", + "timeout_s": 120.0, + } + base.update(overrides) + return SourceConfig(**base) + + +class ScriptedOcrTransport: + """按脚本响应: Exception / "text" / "layout" / "hang";两方法共用一份脚本。""" + + def __init__(self, script): + self.script = list(script) + self.calls = [] + + async def _next(self, method, source, call_id): + self.calls.append((method, source.name, call_id)) + action = self.script.pop(0) + if isinstance(action, Exception): + raise action + if action == "hang": + await asyncio.Event().wait() + return _TEXT_OK if action == "text" else _LAYOUT_OK + + async def recognize_text(self, *, image, source, call_id): + return await self._next("text", source, call_id) + + async def parse_layout(self, *, image, source, call_id): + return await self._next("layout", source, call_id) + + async def check_health(self, *, source): + raise NotImplementedError + + +class StaticSelector: + def order(self, sources, stats): + return list(sources) + + +class RecordingSelector(StaticSelector): + def __init__(self): + self.outcomes = [] + + def record_outcome(self, source_name, ok): + self.outcomes.append((source_name, ok)) + + def health(self, source_name): + return 1.0 + + +class RecordingGate: + """InMemoryGate 包装: 录 record_success 的 count_attempt 与 release_probe。""" + + def __init__(self, inner): + self._inner = inner + self.successes = [] + self.probe_releases = 0 + + async def try_enter(self, source_name, owner): + return await self._inner.try_enter(source_name, owner) + + async def record_success(self, entry, *, count_attempt=True): + self.successes.append((entry.source_name, count_attempt)) + return await self._inner.record_success(entry, count_attempt=count_attempt) + + async def record_failure(self, entry, reason, force_open): + return await self._inner.record_failure(entry, reason, force_open) + + async def release_probe(self, entry): + self.probe_releases += 1 + return await self._inner.release_probe(entry) + + async def retry_after_s(self, sources): + return await self._inner.retry_after_s(sources) + + +class _MemoryRecorder: + def __init__(self): + self.rows = [] + + async def record_llm_call(self, **fields): + self.rows.append(fields) + + +def _client(sources, script, **overrides): + limiter = InMemoryLimiter( + scope="ocr", + sources={s.name: s for s in sources}, + global_limits=_NO_GLOBAL, + lease_ttl_s=100.0, + ) + gate = RecordingGate(InMemoryGate(config=_BREAKER)) + kwargs = { + "scope": "ocr", + "sources": sources, + "selector": RecordingSelector(), + "limiter": limiter, + "breaker": gate, + "transport": ScriptedOcrTransport(script), + "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), + } + kwargs.update(overrides) + return OcrClient(**kwargs), limiter, gate + + +class TestSuccessPaths: + async def test_text_result_fields(self): + client, _, gate = _client([_src()], ["text"]) + r = await client.recognize_text(b"jpg") + assert r.text == "LINE-1" + assert r.source_name == "m1" + assert r.usage.prompt_tokens == 0 and r.usage.completion_tokens == 0 + assert r.latency_ms >= 0 and r.call_id and r.raw["task_type"] == "text" + assert gate.successes == [("m1", True)] + assert client._selector.outcomes == [("m1", True)] + + async def test_layout_result_fields(self): + client, _, _ = _client([_src()], ["layout"]) + r = await client.parse_layout(b"jpg") + assert r.elements[0].type == "table" and r.page_sizes == [(759.0, 540.0)] + + async def test_input_validation(self): + client, _, _ = _client([_src()], []) + with pytest.raises(TypeError): + await client.recognize_text("not-bytes") + with pytest.raises(ValueError): + await client.parse_layout(b"") + + +class TestFailover: + async def test_transient_retries_with_backoff(self): + sleeps = [] + + async def fake_sleep(delay): + sleeps.append(delay) + + client, _, _ = _client( + [_src()], [TransientError("boom", status_code=500), "text"], sleep=fake_sleep + ) + r = await client.recognize_text(b"jpg") + assert r.text == "LINE-1" + assert any(s > 0 for s in sleeps) # Transient 退避后重试 + assert client._selector.outcomes == [("m1", False), ("m1", True)] + + async def test_source_dead_switches_immediately(self): + sleeps = [] + + async def fake_sleep(delay): + sleeps.append(delay) + + s1, s2 = _src(name="m1"), _src(name="m2", base_url="http://10.77.0.20:7867") + client, _, _ = _client( + [s1, s2], [SourceDeadError("401", status_code=401), "text"], sleep=fake_sleep + ) + r = await client.recognize_text(b"jpg") + assert r.source_name == "m2" + assert not [s for s in sleeps if s > 0.0005] # dead 立即换源不退避 + + async def test_retry_exhausted_carries_g1_fields(self): + client, _, _ = _client( + [_src()], + [TransientError("1", status_code=500)] * 3, + ) + with pytest.raises(AllSourcesExhausted) as ei: + await client.recognize_text(b"jpg") + exc = ei.value + assert exc.reason == "retry_exhausted" + assert exc.per_source_reasons == {"m1": "network_error"} # G1: 逐源原因 + assert exc.retry_after_s > 0 # G1: 可延期重投 + + async def test_circuit_open_carries_g1_fields(self): + client, _, _ = _client([_src()], [SourceDeadError("401", status_code=401)]) + with pytest.raises(CircuitOpenError) as ei: + await client.recognize_text(b"jpg") + exc = ei.value + assert exc.per_source_reasons["m1"] == "circuit_open" + assert exc.retry_after_s > 0 + + +class TestTerminalOutcomes: + async def test_result_invalid_counts_no_attempt_and_no_feed(self): + client, _, gate = _client([_src()], [ResultInvalidError("bad zip")]) + with pytest.raises(ResultInvalidError): + await client.parse_layout(b"jpg") + assert gate.successes == [("m1", False)] # 记成功但不入失败率窗 + assert client._selector.outcomes == [] # 坏结果 ≠ 坏服务,不喂健康 + + async def test_rejected_with_status_counts_no_attempt(self): + client, _, gate = _client( + [_src()], [RequestRejectedError("parse failed", status_code=200)] + ) + with pytest.raises(RequestRejectedError): + await client.parse_layout(b"jpg") + assert gate.successes == [("m1", False)] + + async def test_rejected_without_status_no_gate_success(self): + client, _, gate = _client([_src()], [RequestRejectedError("local refuse")]) + with pytest.raises(RequestRejectedError): + await client.recognize_text(b"jpg") + assert gate.successes == [] # 服务未响应: 不记成功(非探针也不归还) + + +class TestBackpressure: + async def test_fail_fast_when_quota_full(self): + src = _src(max_concurrency=1) + client, limiter, _ = _client([src], ["text"], quota_full="fail_fast") + permit = await limiter.acquire("m1", 0) # 占满唯一并发位 + try: + with pytest.raises(AllSourcesExhausted) as ei: + await client.recognize_text(b"jpg") + assert ei.value.reason == "quota_exhausted" + finally: + await permit.settle(0) + await permit.release() + + async def test_wait_until_permit_freed(self): + src = _src(max_concurrency=1) + client, limiter, _ = _client([src], ["text"]) + permit = await limiter.acquire("m1", 0) + + async def free_later(): + await asyncio.sleep(0.01) + await permit.settle(0) + await permit.release() + + release_task = asyncio.create_task(free_later()) + r = await client.recognize_text(b"jpg") + await release_task + assert r.text == "LINE-1" # quota_full=wait: 等到许可释放而非报错 + + async def test_stalled_when_no_progress(self): + src = _src(max_concurrency=1) + client, limiter, _ = _client( + [src], + ["text"], + backpressure=BackpressurePolicy(stall_window_s=0.01, poll_interval_s=0.001), + ) + permit = await limiter.acquire("m1", 0) # 永不释放且全局无进展 + try: + with pytest.raises(AllSourcesExhausted) as ei: + await client.recognize_text(b"jpg") + assert ei.value.reason == "stalled" + assert ei.value.per_source_reasons["m1"] == "rate_limited" + finally: + await permit.settle(0) + await permit.release() + + +class TestCancellation: + async def test_cancel_during_transport_releases_permit(self): + src = _src(max_concurrency=1) + client, limiter, _ = _client([src], ["hang"]) + task = asyncio.create_task(client.recognize_text(b"jpg")) + await asyncio.sleep(0.01) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + stats = await limiter.source_stats("m1") + assert stats.inflight == 0 # permit 在 finally 释放 + + +class TestCheckHealth: + class _HealthTransport(ScriptedOcrTransport): + def __init__(self, mapping): + super().__init__([]) + self.mapping = mapping + + async def check_health(self, *, source): + result = self.mapping[source.name] + if result == "hang": + await asyncio.Event().wait() + if isinstance(result, Exception): + raise result + return result + + async def test_per_source_dict(self): + s1, s2 = _src(name="m1"), _src(name="m2", base_url="http://10.77.0.20:7867") + client, _, _ = _client([s1, s2], []) + client._transport = self._HealthTransport({"m1": True, "m2": False}) + assert await client.check_health() == {"m1": True, "m2": False} + + async def test_cancellation_passes_through(self): + client, _, _ = _client([_src()], []) + client._transport = self._HealthTransport({"m1": "hang"}) + task = asyncio.create_task(client.check_health()) + await asyncio.sleep(0.01) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + +class TestTelemetry: + async def test_success_and_failure_recorded_without_image_bytes(self): + recorder = _MemoryRecorder() + client, _, _ = _client( + [_src()], + [TransientError("boom", status_code=500), "text"], + telemetry=recorder, + ) + await client.recognize_text(b"RAW-IMAGE-BYTES") + assert len(recorder.rows) == 2 # 失败尝试与成功尝试均必录 + for row in recorder.rows: + assert "" in row["messages"] + assert "RAW-IMAGE-BYTES" not in row["messages"] # 图像字节绝不入库 + assert recorder.rows[0]["error"] is not None + assert recorder.rows[1]["error"] is None + assert recorder.rows[1]["prompt_tokens"] == 0 + + +class TestAssembly: + _ENV = { + "OCR__MONKEY__1__BASE_URL": "http://10.77.0.20:7866", + "OCR__MONKEY__1__API_KEY": "none", + "OCR__MONKEY__1__MODEL": "monkey-ocr", + "OCR__MONKEY__1__TIMEOUT_S": "120", + "LLM_MAX_RETRIES": "3", + "LLM_RETRY_BASE_DELAY": "2.0", + "LLM_RETRY_MAX_DELAY": "30.0", + "LLM_CIRCUIT_BREAKER_THRESHOLD": "5", + "LLM_CIRCUIT_BREAKER_COOLDOWN": "60", + "PGW_CACHE_BACKEND": "none", + "PGW_TELEMETRY_BACKEND": "none", + } + + async def test_from_env_assembles(self): + client = OcrClient.from_env("OCR", env=dict(self._ENV)) + try: + assert client._scope == "ocr" + finally: + await client.aclose() + + async def test_non_monkey_provider_rejected(self): + env = { + k.replace("MONKEY", "GLM"): v for k, v in self._ENV.items() + } + with pytest.raises(ValueError, match="monkey"): + OcrClient.from_env("OCR", env=env) + + async def test_aclose_idempotent(self): + client = OcrClient.from_env("OCR", env=dict(self._ENV)) + await client.aclose() + await client.aclose()