diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 71a598a6fe7..365bae93efd 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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, diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index 1d583c16ad7..263e5d6b2cc 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -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 diff --git a/tests/test_litellm/ocr/callback_support.py b/tests/test_litellm/ocr/callback_support.py new file mode 100644 index 00000000000..068d896bb52 --- /dev/null +++ b/tests/test_litellm/ocr/callback_support.py @@ -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}