Files
PolyGateway/src/polygateway/ocr.py
T
iomgaa 0477d9534b fix: apply the non-productive stall budget to the ocr loop
Same failure path as the embedding loop: one timed-out attempt drains the
wall-clock window, and the next round without a runnable source declares
the scope dead in _on_no_runnable. All three governance loops now meter
stall the same way.
2026-08-06 09:51:30 -04:00

524 lines
20 KiB
Python

"""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,
SourceNotConfiguredError,
TransientError,
)
from polygateway.middleware.breaker import BreakerGate
from polygateway.middleware.ratelimit import QuotaGate
from polygateway.middleware.retry import StallClock, _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,
strip_unsupported_extra_body,
)
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
# MonkeyOCR 只发 multipart 表单,带 extra_body 的源必须先剥离,否则
# 遥测会记录一个从未发出的采样参数(issue #4 决策 G)
self._sources = strip_unsupported_extra_body(list(sources), path="OCR")
self._selector = selector
self._feed_health = isinstance(selector, OutcomeAwareSelector)
self._quota = QuotaGate(limiter, scope=self._scope)
self._breaker = BreakerGate(breaker, scope=self._scope)
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] = {}
# 只计非生产性等待(issue #8): 真实尝试由重试预算治理,不重复烧 stall 预算
clock = StallClock(self._now)
while True:
picked, gate_rejections = await self._pick_runnable(reasons)
if picked is None:
await self._on_no_runnable(gate_rejections, reasons, clock)
continue
async with clock.attempting():
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], clock: StallClock
) -> 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 clock.stalled_s() > 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, SourceNotConfiguredError) 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",
)
# 错误带异常类名前缀(metric ocr-call-success 注册口径: 按类名归组)
if error is None or isinstance(error, str):
error_text = error
else:
error_text = f"{type(error).__name__}: {error}"
await self._emitter.emit_attempt(
request=request,
source=source,
call_id=call_id,
latency_ms=latency_ms,
response=response,
error=error_text,
)
@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,
)