diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index f4f968c8d34..627049b9423 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -18,6 +18,7 @@ import httpx import litellm from litellm._logging import verbose_logger from litellm.constants import request_timeout +from litellm.litellm_core_utils.call_completion import CallCompletion from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.azure_ai.ocr.common_utils import ( is_azure_cohere_parse_model, @@ -462,6 +463,8 @@ async def aocr( timeout: float | httpx.Timeout | None = None, custom_llm_provider: str | None = None, extra_headers: dict[str, object] | None = None, + *, + _litellm_call_completion: CallCompletion | None = None, **kwargs: object, ) -> OCRResponse: """ @@ -522,7 +525,6 @@ async def aocr( ) ``` """ - call_completion: Final = kwargs.pop("_litellm_call_completion", None) completion_kwargs: Final[dict[str, object]] = { "model": model, "document": document, @@ -542,7 +544,7 @@ async def aocr( custom_llm_provider=custom_llm_provider, extra_headers=extra_headers, kwargs=kwargs, - call_completion=call_completion, + call_completion=_litellm_call_completion, ) try: if rust_enabled() and _rust_ocr_supported(request): @@ -742,6 +744,8 @@ def ocr( timeout: float | httpx.Timeout | None = None, custom_llm_provider: str | None = None, extra_headers: dict[str, object] | None = None, + *, + _litellm_call_completion: CallCompletion | None = None, **kwargs: object, ) -> OCRResponse | Coroutine[object, object, OCRResponse]: """ @@ -806,7 +810,6 @@ def ocr( print(f"Page {page.index}: {page.markdown}") ``` """ - call_completion: Final = kwargs.pop("_litellm_call_completion", None) completion_kwargs: Final[dict[str, object]] = { "model": model, "document": document, @@ -826,7 +829,7 @@ def ocr( custom_llm_provider=custom_llm_provider, extra_headers=extra_headers, kwargs=kwargs, - call_completion=call_completion, + call_completion=_litellm_call_completion, ) try: _is_async: Final = kwargs.pop("aocr", False) is True diff --git a/litellm/utils.py b/litellm/utils.py index 85b20bdd602..5970ea8876a 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -1658,7 +1658,6 @@ def client(original_function): start_time=start_time, end_time=end_time, ) - completion.release() return result except Exception as e: call_type = original_function.__name__ @@ -1942,14 +1941,12 @@ def client(original_function): and _caching_handler_response is not None and _caching_handler_response.final_embedding_cached_response is not None ): - combined_response: Final = _llm_caching_handler._combine_cached_embedding_response_with_api_result( + return _llm_caching_handler._combine_cached_embedding_response_with_api_result( _caching_handler_response=_caching_handler_response, embedding_response=result, start_time=start_time, end_time=end_time, ) - completion.release() - return combined_response _update_response_metadata( result=result, @@ -1960,7 +1957,6 @@ def client(original_function): end_time=end_time, ) - completion.release() return result except Exception as e: traceback_exception: Final = traceback.format_exc() diff --git a/tests/test_litellm/litellm_core_utils/test_call_completion.py b/tests/test_litellm/litellm_core_utils/test_call_completion.py index 140e482a3c1..41d7a1d07f0 100644 --- a/tests/test_litellm/litellm_core_utils/test_call_completion.py +++ b/tests/test_litellm/litellm_core_utils/test_call_completion.py @@ -1,13 +1,19 @@ +import asyncio import contextvars import datetime +import weakref from collections.abc import Callable, Coroutine from concurrent.futures import Future +from importlib import import_module from typing import Final from unittest.mock import AsyncMock, MagicMock import pytest +import litellm from litellm.litellm_core_utils.call_completion import CallCompletion, PythonCompletion +from litellm.llms.base_llm.ocr.transformation import OCRResponse +from litellm.rust_bridge import ocr as rust_ocr_bridge from litellm.utils import client @@ -152,22 +158,29 @@ async def test_python_completion_retains_deferred_success_arguments(monkeypatch: logging_obj: Final = MagicMock() logging_obj._defer_async_logging = True logging_obj.async_success_handler = AsyncMock() - completion: Final = PythonCompletion( - logging_obj, - RecordingExecutor(), - async_call=True, - internal_call=False, - completion_with_fallbacks=False, + completion: Final = CallCompletion( + PythonCompletion( + logging_obj, + RecordingExecutor(), + async_call=True, + internal_call=False, + completion_with_fallbacks=False, + ) ) + worker: Final = MagicMock() + monkeypatch.setattr("litellm.litellm_core_utils.logging_worker.GLOBAL_LOGGING_WORKER", worker) scheduled: Final[list[Coroutine[object, object, None]]] = [] monkeypatch.setattr("asyncio.create_task", scheduled.append) now: Final = datetime.datetime.now(datetime.timezone.utc) completion.success(response, now, now) + completion.release() logging_obj._enqueue_deferred_logging() assert len(scheduled) == 1 - scheduled[0].close() + await scheduled[0] + await worker.ensure_initialized_and_enqueue.call_args.kwargs["async_coroutine"] + logging_obj.async_success_handler.assert_awaited_once_with(result=response, start_time=now, end_time=now) @pytest.mark.asyncio @@ -265,3 +278,105 @@ async def test_async_ocr_wrapper_retains_completion_until_metadata_finishes( assert caught.value is metadata_error assert native_completion.successes == [response] assert native_completion.failures == [metadata_error, metadata_error] + + +@pytest.mark.parametrize("asynchronous", [False, True]) +@pytest.mark.asyncio +async def test_ocr_completion_stays_separate_from_marshaled_provider_options( + monkeypatch: pytest.MonkeyPatch, asynchronous: bool +) -> None: + native_completion: Final = RecordingCompletion() + response: Final = OCRResponse(model="mistral-ocr-latest", pages=[]) + metadata: Final = {"request": "shared"} + pages: Final = [0, 2] + + def run( + request: rust_ocr_bridge.LiteLLMOcrRequest, + resolve_secret: Callable[[str], str | None], + convert_file_document: Callable[[dict[str, object]], dict[str, str]], + ) -> OCRResponse: + assert request.kwargs["metadata"] is metadata + assert "_litellm_call_completion" not in request.kwargs + marshalled: Final = rust_ocr_bridge._marshal(request, resolve_secret, convert_file_document) + assert "_litellm_call_completion" not in marshalled.kwargs + assert marshalled.kwargs["pages"] is pages + assert marshalled.call_completion is request.call_completion + assert marshalled.call_completion is not None + assert marshalled.call_completion.attach(native_completion) + return response + + async def arun( + request: rust_ocr_bridge.LiteLLMOcrRequest, + resolve_secret: Callable[[str], str | None], + convert_file_document: Callable[[dict[str, object]], dict[str, str]], + ) -> OCRResponse: + return run(request, resolve_secret, convert_file_document) + + monkeypatch.setattr(import_module("litellm.ocr.main"), "rust_enabled", lambda: True) + monkeypatch.setattr(rust_ocr_bridge, "run", run) + monkeypatch.setattr(rust_ocr_bridge, "arun", arun) + arguments: Final = { + "model": "mistral/mistral-ocr-latest", + "document": {"type": "document_url", "document_url": "https://example.com/doc.pdf"}, + "api_key": "test-key", + "metadata": metadata, + "pages": pages, + } + + result: Final = await litellm.aocr(**arguments) if asynchronous else litellm.ocr(**arguments) + + assert result is response + assert native_completion.successes == [response] + assert native_completion.failures == [] + assert arguments["metadata"] is metadata + assert "_litellm_call_completion" not in arguments + + +@pytest.mark.parametrize( + ("asynchronous", "exit_path"), + [(False, "success"), (True, "success"), (False, "callback_error"), (True, "callback_error"), (True, "cancelled")], +) +@pytest.mark.asyncio +async def test_wrapper_releases_completion_resources_on_every_exit( + monkeypatch: pytest.MonkeyPatch, asynchronous: bool, exit_path: str +) -> None: + retained: Final[list[tuple[CallCompletion, weakref.ReferenceType[object]]]] = [] + callback_error: Final = RuntimeError("failure callback failed") + implementation: Final = MagicMock(spec=RecordingCompletion) + implementation.failure.side_effect = callback_error + implementation.async_failure = AsyncMock() + response: Final = object() + + def ocr(*, _litellm_call_completion: CallCompletion, **kwargs: object) -> object: + retained.append((_litellm_call_completion, weakref.ref(_litellm_call_completion.python_implementation))) + assert _litellm_call_completion.attach(implementation) + if exit_path == "cancelled": + raise asyncio.CancelledError + if exit_path == "callback_error": + raise ValueError("provider failed") + return response + + async def aocr(*, _litellm_call_completion: CallCompletion, **kwargs: object) -> object: + return ocr(_litellm_call_completion=_litellm_call_completion, **kwargs) + + monkeypatch.setattr("litellm.utils.function_setup", MagicMock(return_value=(MagicMock(), {}))) + monkeypatch.setattr("litellm.utils.load_credentials_from_list", MagicMock()) + wrapped: Final = client(aocr if asynchronous else ocr) + + if exit_path == "success": + result: Final = await wrapped() if asynchronous else wrapped() + assert result is response + assert implementation.success.call_args.args[0] is response + elif exit_path == "cancelled": + with pytest.raises(asyncio.CancelledError): + await wrapped() + implementation.success.assert_not_called() + implementation.failure.assert_not_called() + else: + with pytest.raises(RuntimeError, match="failure callback failed") as caught: + await wrapped() if asynchronous else wrapped() + assert caught.value is callback_error + implementation.async_failure.assert_not_called() + + assert len(retained) == 1 + assert retained[0][1]() is None