feat: add OcrClient governed loop with dual ports
This commit is contained in:
@@ -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"<ocr:{kind} image_bytes={len(image)}>"}],
|
||||
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"<elements n={len(result.elements)} pages={len(result.page_sizes)}>"
|
||||
|
||||
# —— 生命周期 ——
|
||||
|
||||
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,
|
||||
)
|
||||
@@ -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 "<ocr:text image_bytes=15>" 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()
|
||||
Reference in New Issue
Block a user