fix(ocr): deliver Python post-call logging before normalization

This commit is contained in:
Yujong Lee 2026-09-02 22:48:13 -07:00
parent 62ed7e1942
commit 807d99f5a5
3 changed files with 227 additions and 3 deletions

View file

@ -1712,6 +1712,11 @@ class BaseLLMHTTPHandler:
except Exception as e:
raise self._handle_error(e=e, provider_config=provider_config)
logging_obj.post_call(
original_response=response.text,
input="OCR document processing",
api_key=api_key,
)
return self._transform_ocr_response(
provider_config=provider_config,
model=model,
@ -1775,7 +1780,11 @@ class BaseLLMHTTPHandler:
except Exception as e:
raise self._handle_error(e=e, provider_config=provider_config)
# Use async response transform for async operations
logging_obj.post_call(
original_response=response.text,
input="OCR document processing",
api_key=api_key,
)
return await provider_config.async_transform_ocr_response(
model=model,
raw_response=response,

View file

@ -2,6 +2,7 @@ import asyncio
import json
import logging
import time
from typing import Final
from unittest.mock import AsyncMock, Mock, patch
import httpx
@ -10,9 +11,10 @@ import pytest
import litellm
from litellm._logging import verbose_logger
from litellm.integrations.code_interpreter_interception.handler import (
CodeInterpreterInterceptionLogger,
LITELLM_CODE_EXECUTION_TOOL_NAME,
CodeInterpreterInterceptionLogger,
)
from litellm.llms.azure.videos.transformation import AzureVideoConfig
from litellm.llms.base_llm.audio_transcription.transformation import (
AudioTranscriptionRequestData,
BaseAudioTranscriptionConfig,
@ -26,7 +28,6 @@ from litellm.llms.custom_httpx.llm_http_handler import (
_has_pre_call_deployment_hook,
_rust_responses_websocket_enabled,
)
from litellm.llms.azure.videos.transformation import AzureVideoConfig
from litellm.llms.openai.videos.transformation import OpenAIVideoConfig
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.router import GenericLiteLLMParams
@ -36,6 +37,70 @@ _ACTIVE_KEY = "_code_interpreter_interception_active"
_SANDBOX_KEY = "_code_interpreter_interception_sandbox_key"
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.parametrize("outcome", ["success", "observer_error", "malformed", "401", "429", "500"])
async def test_ocr_sdk_logs_provider_response_before_normalization(asynchronous: bool, outcome: str) -> None:
from tests.test_litellm.ocr.callback_support import (
DOCUMENT,
MODEL,
RESPONSE,
CallbackRecorder,
assert_provider_request,
ocr_upstream,
)
success: Final = outcome in ("success", "observer_error")
recorder: Final = CallbackRecorder(asynchronous, raises="post" if outcome == "observer_error" else "")
follower: Final = CallbackRecorder(asynchronous, name="follower")
status: Final = int(outcome) if outcome.isdigit() else 200
body: Final = "not-json" if outcome == "malformed" else RESPONSE if status == 200 else '{"error":"rejected"}'
with ocr_upstream(status, body) as upstream:
params: Final = dict(
model=MODEL,
document=DOCUMENT,
api_base=upstream.api_base,
api_key="test-key",
rust=False,
callbacks=[recorder, follower],
num_retries=0,
)
try:
result: Final = (
await litellm.aocr(**params) if asynchronous else await asyncio.to_thread(litellm.ocr, **params)
)
except Exception as error:
assert not success
if status != 200:
assert getattr(error, "status_code", None) == status
else:
assert success
assert result.pages[0].markdown == "callback-test"
events: Final = await recorder.wait()
names: Final = tuple(event.name for event in events)
following_events: Final = await follower.wait()
assert sorted(event.name for event in following_events) == sorted(names)
prefix: Final = ("pre", "post") if status == 200 else ("pre",)
assert names[: len(prefix)] == prefix
terminal: Final = "success" if success else "failure"
expected: Final = (
("async_success",)
if asynchronous and success
else ((terminal, f"async_{terminal}") if asynchronous else (terminal,))
)
assert sorted(names[len(prefix) :]) == sorted(expected)
assert len({event.call_id for event in events}) == 1
if status == 200:
assert events[1].original_response == body
assert events[1].response_type == "NoneType"
assert events[1].start_time is not None
assert events[1].end_time is None
for event in events[len(prefix) :]:
assert event.start_time is not None and event.end_time is not None
assert event.end_time >= event.start_time
assert_provider_request(upstream)
def test_prepare_fake_stream_request():
# Initialize the BaseLLMHTTPHandler
handler = BaseLLMHTTPHandler()
@ -1467,6 +1532,7 @@ def _make_responses_handler_call(signed_body):
signing provider (e.g. Bedrock Mantle).
"""
from unittest.mock import MagicMock
from litellm.llms.custom_httpx.http_handler import HTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.types.router import GenericLiteLLMParams
@ -1522,6 +1588,7 @@ def test_responses_handler_signs_after_fake_stream_prep_strips_stream():
We snapshot request_data at sign time and assert "stream" is already gone.
"""
from unittest.mock import MagicMock
from litellm.llms.custom_httpx.http_handler import HTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.types.llms.openai import ResponsesAPIResponse
@ -1585,6 +1652,7 @@ def _make_compact_handler_call(signed_body, is_async):
signing provider (e.g. Bedrock Mantle SigV4 / bearer).
"""
from unittest.mock import MagicMock
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.types.router import GenericLiteLLMParams

View file

@ -0,0 +1,147 @@
import asyncio
import json
import threading
from collections.abc import Generator, Mapping
from contextlib import contextmanager
from dataclasses import dataclass
from datetime import datetime
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from typing import Final
from litellm.integrations.custom_logger import CustomLogger
MODEL: Final = "mistral/mistral-ocr-4-1"
DOCUMENT: Final = {"type": "document_url", "document_url": "https://example.com/document.pdf"}
RESPONSE: Final = (
'{"pages":[{"index":0,"markdown":"callback-test"}],"model":"mistral-ocr-4-1","usage_info":{"pages_processed":1}}'
)
@dataclass(frozen=True, slots=True)
class CallbackEvent:
name: str
model: object
call_id: object
original_response: object
response_type: str
start_time: datetime | None
end_time: datetime | None
class CallbackRecorder(CustomLogger):
def __init__(self, asynchronous: bool, raises: str = "", name: str = "recorder") -> None:
super().__init__() # pyright: ignore[reportUnknownMemberType] # CustomLogger exposes untyped keyword arguments
self.asynchronous = asynchronous
self.name = name
self.raises = raises
self.events: tuple[CallbackEvent, ...] = ()
self.done = threading.Event()
self.lock = threading.Lock()
def record(
self,
name: str,
kwargs: Mapping[str, object],
response: object = None,
start_time: datetime | None = None,
end_time: datetime | None = None,
) -> None:
event: Final = CallbackEvent(
name,
kwargs.get("model"),
kwargs.get("litellm_call_id"),
kwargs.get("original_response"),
type(response).__name__,
start_time,
end_time,
)
with self.lock:
self.events += (event,)
terminals: Final = tuple(
item.name for item in self.events if "success" in item.name or "failure" in item.name
)
if len(terminals) >= (2 if self.asynchronous and "failure" in name else 1):
self.done.set()
if name == self.raises:
raise ValueError("intentional observer failure")
def log_pre_api_call(self, model: str, messages: object, kwargs: Mapping[str, object]) -> None:
self.record("pre", kwargs)
def log_post_api_call(
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime | None
) -> None:
self.record("post", kwargs, response_obj, start_time, end_time)
def log_success_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime
) -> None:
self.record("success", kwargs, response_obj, start_time, end_time)
def log_failure_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime
) -> None:
self.record("failure", kwargs, response_obj, start_time, end_time)
async def async_log_success_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime
) -> None:
self.record("async_success", kwargs, response_obj, start_time, end_time)
async def async_log_failure_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime
) -> None:
self.record("async_failure", kwargs, response_obj, start_time, end_time)
async def wait(self) -> tuple[CallbackEvent, ...]:
assert await asyncio.to_thread(self.done.wait, 5), tuple(event.name for event in self.events)
return self.events
class OcrUpstream(ThreadingHTTPServer):
daemon_threads = True
def __init__(self, status: int, body: str) -> None:
super().__init__(("127.0.0.1", 0), OcrHandler)
self.status = status
self.body = body.encode()
self.requests: tuple[tuple[str, str], ...] = ()
@property
def api_base(self) -> str:
return f"http://127.0.0.1:{self.server_port}/v1"
class OcrHandler(BaseHTTPRequestHandler):
def do_POST(self) -> None:
assert isinstance(self.server, OcrUpstream)
body: Final = self.rfile.read(int(self.headers["content-length"])).decode()
self.server.requests += ((self.path, body),)
self.send_response(self.server.status)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(self.server.body)))
self.end_headers()
self.wfile.write(self.server.body)
def log_message(self, format: str, *args: object) -> None:
pass
@contextmanager
def ocr_upstream(status: int = 200, body: str = RESPONSE) -> Generator[OcrUpstream, None, None]:
with OcrUpstream(status, body) as server:
thread: Final = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
try:
yield server
finally:
server.shutdown()
thread.join(timeout=5)
assert not thread.is_alive()
def assert_provider_request(server: OcrUpstream) -> None:
assert len(server.requests) == 1
path, body = server.requests[0]
assert path == "/v1/ocr"
assert json.loads(body) == {"model": "mistral-ocr-4-1", "document": DOCUMENT}