diff --git a/src/polygateway/ocr.py b/src/polygateway/ocr.py index 72d3d2a..5a1b1d1 100644 --- a/src/polygateway/ocr.py +++ b/src/polygateway/ocr.py @@ -18,7 +18,7 @@ import random import time import uuid from dataclasses import dataclass -from typing import TYPE_CHECKING, Literal +from typing import TYPE_CHECKING, Any, Literal from loguru import logger @@ -46,6 +46,7 @@ from polygateway.types import ( OcrTextResult, Usage, strip_unsupported_extra_body, + validate_caller_dimensions, ) if TYPE_CHECKING: @@ -142,9 +143,21 @@ class OcrClient: *, session_id: str | None = None, parent_call_id: str | None = None, + tenant_id: str | None = None, + meta: Mapping[str, Any] | None = None, ) -> OcrTextResult: - """一次治理文本转录(/ocr/text);text 空串 = 合法"无文字"。""" - outcome = await self._call("text", image, session_id, parent_call_id) + """一次治理文本转录(/ocr/text);text 空串 = 合法"无文字"。 + + `tenant_id` 与 `meta` 是调用方自定义维度,只进遥测(issue #11)。 + """ + # 必须在进链路之前校验: 链路内的一切失败都被遥测层降级成 warning + # (库铁律「遥测写失败降级不冒泡」),校验放下游等于没有校验 + dimension_tenant_id, dimensions = validate_caller_dimensions( + tenant_id, meta, origin="recognize_text(tenant_id=..., meta=...)" + ) + outcome = await self._call( + "text", image, session_id, parent_call_id, dimension_tenant_id, dimensions + ) result = outcome.result return OcrTextResult( text=result.text, @@ -161,9 +174,20 @@ class OcrClient: *, session_id: str | None = None, parent_call_id: str | None = None, + tenant_id: str | None = None, + meta: Mapping[str, Any] | None = None, ) -> OcrLayoutResult: - """一次治理版面解析(/parse → ZIP);elements 空 = 合法"无元素"。""" - outcome = await self._call("layout", image, session_id, parent_call_id) + """一次治理版面解析(/parse → ZIP);elements 空 = 合法"无元素"。 + + `tenant_id` 与 `meta` 是调用方自定义维度,只进遥测(issue #11)。 + """ + # 校验早于链路,理由同 recognize_text;origin 标明方法名以便定位入口 + dimension_tenant_id, dimensions = validate_caller_dimensions( + tenant_id, meta, origin="parse_layout(tenant_id=..., meta=...)" + ) + outcome = await self._call( + "layout", image, session_id, parent_call_id, dimension_tenant_id, dimensions + ) result = outcome.result return OcrLayoutResult( elements=result.elements, @@ -195,6 +219,8 @@ class OcrClient: image: bytes, session_id: str | None, parent_call_id: str | None, + tenant_id: str | None, + meta: dict[str, Any], ) -> _AttemptOutcome: if not isinstance(image, bytes): raise TypeError("image 必须是 bytes(路径读取/批量拼帧留业务侧,D9)") @@ -213,7 +239,7 @@ class OcrClient: continue async with clock.attempting(): outcome = await self._attempt( - kind, image, *picked, reasons, session_id, parent_call_id + kind, image, *picked, reasons, session_id, parent_call_id, tenant_id, meta ) if isinstance(outcome, _AttemptOutcome): return outcome @@ -294,9 +320,13 @@ class OcrClient: reasons: dict[str, str], session_id: str | None, parent_call_id: str | None, + tenant_id: str | None, + meta: dict[str, Any], ) -> _AttemptOutcome | _FailedAttempt: call_id = str(uuid.uuid4()) started = self._now() + # 四个 emit 分支(成功/终态拒绝/取消/可重试失败)都必须带调用方维度: + # 失败行与取消行同样需要租户归属,漏掉任一分支就会写出无归属的行 try: result = await self._invoke(kind, image, source, call_id) await self._record_quietly(self._breaker.record_success(entry)) @@ -304,20 +334,47 @@ class OcrClient: 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 + kind, + image, + source, + call_id, + started, + session_id, + parent_call_id, + tenant_id, + meta, + 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 + kind, + image, + source, + call_id, + started, + session_id, + parent_call_id, + tenant_id, + meta, + 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" + kind, + image, + source, + call_id, + started, + session_id, + parent_call_id, + tenant_id, + meta, + error="cancelled", ) raise except (SourceDeadError, TransientError) as exc: @@ -327,7 +384,16 @@ class OcrClient: 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 + kind, + image, + source, + call_id, + started, + session_id, + parent_call_id, + tenant_id, + meta, + error=exc, ) return _FailedAttempt(exc, immediate=dead) finally: @@ -389,16 +455,22 @@ class OcrClient: started: float, session_id: str | None, parent_call_id: str | None, + tenant_id: str | None, + meta: dict[str, Any], result: OcrTextTransportResult | OcrLayoutTransportResult | None = None, error: object | None = None, ) -> None: """逐尝试遥测(单一 Emitter): messages 占位摘要,图像 bytes 绝不入库。""" if self._emitter is None: return + # 这个 ChatRequest 只为复用同一个 Emitter 而现场构造(OCR 不走 chat 洋葱), + # 故调用方维度必须在这里显式填回,否则 OCR 行的维度恒为空 request = ChatRequest( messages=[{"role": "user", "content": f""}], session_id=session_id, parent_call_id=parent_call_id, + tenant_id=tenant_id, + meta=meta, ) latency_ms = int((self._now() - started) * 1000) response = None diff --git a/tests/unit/test_ocr_client.py b/tests/unit/test_ocr_client.py index 2c0cecb..85e875b 100644 --- a/tests/unit/test_ocr_client.py +++ b/tests/unit/test_ocr_client.py @@ -480,6 +480,59 @@ class TestTelemetry: assert (await limiter.source_stats("m1")).tpm_used == 0 # settle(0) 全额退回预扣 +class TestOcrCallerDimensions: + """issue #11: 调用方自定义维度必须沿 OCR 链四层透传到每一行遥测。 + + OCR 行与 chat 行落在同一张 `llm_calls` 表: 不覆盖这条链会让同一张表里 + 一部分行有租户归属、一部分永远空白,而"先启用后加列则归属无法还原"。 + """ + + async def test_recognize_text_row_carries_dimensions(self): + recorder = _MemoryRecorder() + client, _, _ = _client([_src()], ["text"], telemetry=recorder) + await client.recognize_text(b"jpg", tenant_id="t1", meta={"batch": "b-42"}) + assert recorder.rows[0]["tenant_id"] == "t1" + assert recorder.rows[0]["meta"] == '{"batch": "b-42"}' + + async def test_parse_layout_row_carries_dimensions(self): + """两个公共方法都是入口: 只测一个会漏掉另一个的透传缺口。""" + recorder = _MemoryRecorder() + client, _, _ = _client([_src()], ["layout"], telemetry=recorder) + await client.parse_layout(b"jpg", tenant_id="t2", meta={"batch": "b-43"}) + assert recorder.rows[0]["tenant_id"] == "t2" + assert recorder.rows[0]["meta"] == '{"batch": "b-43"}' + + async def test_failed_attempt_row_also_carries_dimensions(self): + """失败行同样需要归属: 某租户的请求没被服务,正是审计最需要的一行。""" + recorder = _MemoryRecorder() + client, _, _ = _client( + [_src()], + [TransientError("boom", status_code=500), "text"], + telemetry=recorder, + ) + await client.recognize_text(b"jpg", tenant_id="t1", meta={"batch": "b-42"}) + assert len(recorder.rows) == 2 # 失败尝试 + 成功尝试 + assert [r["tenant_id"] for r in recorder.rows] == ["t1", "t1"] + assert [r["meta"] for r in recorder.rows] == ['{"batch": "b-42"}'] * 2 + + @pytest.mark.parametrize("method", ["recognize_text", "parse_layout"]) + async def test_invalid_meta_rejected_before_any_telemetry(self, method): + """校验必须早于遥测: 链路内的失败都被降级成 warning,放下游等于没有校验。""" + recorder = _MemoryRecorder() + client, _, _ = _client([_src()], ["text"], telemetry=recorder) + with pytest.raises(ValueError, match="meta"): + await getattr(client, method)(b"jpg", meta={"Bad Key": 1}) + assert recorder.rows == [] + assert client._transport.calls == [] # 连调用都没发出 + + async def test_defaults_land_as_sentinels(self): + recorder = _MemoryRecorder() + client, _, _ = _client([_src()], ["text"], telemetry=recorder) + await client.recognize_text(b"jpg") + assert recorder.rows[0]["tenant_id"] == "" # 空串哨兵,不是 None + assert recorder.rows[0]["meta"] == "{}" + + class TestAssembly: _ENV = { "OCR__MONKEY__1__BASE_URL": "http://10.77.0.20:7866",