refactor(ocr): extract call completion boundary

This commit is contained in:
Yujong Lee 2026-09-11 08:06:37 -07:00 • committed by yujonglee
parent 83ab0113f0
commit 5ac9eec465
5 changed files with 464 additions and 155 deletions

View file

@ -0,0 +1,229 @@
from __future__ import annotations
import asyncio
import contextvars
import datetime
from collections.abc import Awaitable, Callable, Coroutine
from concurrent.futures import Future
from typing import Protocol
class CompletionLogging(Protocol):
def success_handler(
self,
result: object,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
def failure_handler(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
def async_success_handler(
self,
result: object,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> Coroutine[object, object, None]: ...
def async_failure_handler(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> Awaitable[None]: ...
def handle_sync_success_callbacks_for_async_calls(
self,
result: object,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
class CompletionExecutor(Protocol):
def submit(
self,
function: Callable[..., object],
*args: object,
) -> Future[object]: ...
class Completion(Protocol):
def success(
self,
result: object,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
def failure(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
async def async_failure(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
class PythonCompletion:
def __init__(
self,
logging_obj: CompletionLogging,
executor: CompletionExecutor | None,
*,
async_call: bool,
internal_call: bool,
completion_with_fallbacks: bool,
) -> None:
self._logging_obj = logging_obj
self._executor = executor
self._async_call = async_call
self._internal_call = internal_call
self._completion_with_fallbacks = completion_with_fallbacks
def success(
self,
result: object,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None:
if not self._async_call:
assert self._executor is not None
context = contextvars.copy_context()
self._executor.submit(
context.run,
self._logging_obj.success_handler,
result,
start_time,
end_time,
)
return
if not self._internal_call:
if getattr(self._logging_obj, "_defer_async_logging", False):
def enqueue_deferred_logging() -> None:
asyncio.create_task(self._dispatch_async_success(result, start_time, end_time))
setattr( # noqa: B010 # optional legacy logger field is absent from narrow test doubles
self._logging_obj,
"_enqueue_deferred_logging",
enqueue_deferred_logging,
)
else:
asyncio.create_task(self._dispatch_async_success(result, start_time, end_time))
self._logging_obj.handle_sync_success_callbacks_for_async_calls(
result=result,
start_time=start_time,
end_time=end_time,
)
def failure(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None:
if self._async_call and self._internal_call:
return
self._logging_obj.failure_handler(exception, traceback_exception, start_time, end_time)
async def async_failure(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None:
if not self._async_call or self._internal_call:
return
await self._logging_obj.async_failure_handler(exception, traceback_exception, start_time, end_time)
async def _dispatch_async_success(
self,
result: object,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None:
if self._completion_with_fallbacks:
return
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue( # pyright: ignore[reportUnknownMemberType] # legacy worker lacks generic coroutine annotations
async_coroutine=self._logging_obj.async_success_handler(
result=result,
start_time=start_time,
end_time=end_time,
)
)
self._logging_obj.handle_sync_success_callbacks_for_async_calls(
result=result,
start_time=start_time,
end_time=end_time,
)
class CallCompletion:
def __init__(self, implementation: Completion) -> None:
self._python_implementation = implementation
self._implementation = implementation
self._attached = False
@property
def python_implementation(self) -> Completion:
return self._python_implementation
def attach(self, implementation: Completion) -> bool:
if self._attached:
return False
self._implementation = implementation
self._attached = True
return True
def success(
self,
result: object,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None:
self._implementation.success(result, start_time, end_time)
def failure(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None:
self._implementation.failure(exception, traceback_exception, start_time, end_time)
async def async_failure(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None:
await self._implementation.async_failure(
exception,
traceback_exception,
start_time,
end_time,
)

View file

@ -522,6 +522,7 @@ async def aocr(
)
```
"""
call_completion: Final = kwargs.pop("_litellm_call_completion", None)
completion_kwargs: Final[dict[str, object]] = {
"model": model,
"document": document,
@ -541,6 +542,7 @@ async def aocr(
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
kwargs=kwargs,
call_completion=call_completion,
)
try:
if rust_enabled() and _rust_ocr_supported(request):
@ -804,6 +806,7 @@ 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,
@ -823,6 +826,7 @@ def ocr(
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
kwargs=kwargs,
call_completion=call_completion,
)
try:
_is_async: Final = kwargs.pop("aocr", False) is True

View file

@ -52,6 +52,7 @@ class LiteLLMOcrRequest:
custom_llm_provider: str | None
extra_headers: dict[str, object] | None
kwargs: Mapping[str, object]
call_completion: object = None
input_sources: Mapping[str, str] | None = None
@ -294,6 +295,7 @@ def _marshal(
custom_llm_provider=request.custom_llm_provider,
extra_headers=request.extra_headers,
kwargs=optional_params,
call_completion=request.call_completion,
input_sources=input_sources,
)

View file

@ -7,7 +7,6 @@ import ast
import asyncio
import base64
import binascii
import contextvars
import copy
import datetime
import hashlib
@ -66,7 +65,6 @@ from litellm.constants import (
DEFAULT_EMBEDDING_PARAM_VALUES,
DEFAULT_MAX_LRU_CACHE_SIZE,
DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT,
DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT,
DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET,
DEFAULT_TRIM_RATIO,
FUNCTION_DEFINITION_TOKEN_COUNT,
@ -81,6 +79,7 @@ from litellm.constants import (
PROVIDERS_THAT_AUTHENTICATE_ON_PROVIDER_INFO,
TOOL_CHOICE_OBJECT_TOKEN_COUNT,
)
from litellm.litellm_core_utils.call_completion import CallCompletion, PythonCompletion
from litellm.litellm_core_utils.core_helpers import normalize_drop_params
from litellm.litellm_core_utils.fallback_generalizations import (
match_capability_generalizations,
@ -279,7 +278,7 @@ except (ImportError, AttributeError, TypeError):
# Convert to str (if necessary)
claude_json_str = json.dumps(json_data)
import importlib.metadata
from collections.abc import AsyncIterator, Callable, Iterable, Iterator, Mapping, Sequence
from collections.abc import Callable, Iterable, Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast, get_args
from litellm import utils as litellm_utils
@ -1196,79 +1195,6 @@ def function_setup(
raise e
def _dispatch_success_logging(
logging_obj: LiteLLMLoggingObject,
result: object,
start_time: datetime.datetime,
end_time: datetime.datetime,
is_completion_with_fallbacks: bool,
is_litellm_internal_call: bool,
) -> None:
if not is_litellm_internal_call:
if getattr(logging_obj, "_defer_async_logging", False):
def _enqueue_deferred_logging() -> None:
asyncio.create_task(
_client_async_logging_helper(
logging_obj=logging_obj,
result=result,
start_time=start_time,
end_time=end_time,
is_completion_with_fallbacks=is_completion_with_fallbacks,
)
)
logging_obj._enqueue_deferred_logging = _enqueue_deferred_logging
else:
asyncio.create_task(
_client_async_logging_helper(
logging_obj=logging_obj,
result=result,
start_time=start_time,
end_time=end_time,
is_completion_with_fallbacks=is_completion_with_fallbacks,
)
)
logging_obj.handle_sync_success_callbacks_for_async_calls(
result=result,
start_time=start_time,
end_time=end_time,
)
async def _client_async_logging_helper(
logging_obj: LiteLLMLoggingObject,
result,
start_time,
end_time,
is_completion_with_fallbacks: bool,
):
if (
is_completion_with_fallbacks is False
): # don't log the parent event litellm.completion_with_fallbacks as a 'log_success_event', this will lead to double logging the same call - https://github.com/BerriAI/litellm/issues/7477
print_verbose(
f"Async Wrapper: Completed Call, calling async_success_handler: {logging_obj.async_success_handler}"
)
################################################
# Async Logging Worker
################################################
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
async_coroutine=logging_obj.async_success_handler(result=result, start_time=start_time, end_time=end_time)
)
################################################
# Sync Logging Worker
################################################
logging_obj.handle_sync_success_callbacks_for_async_calls(
result=result,
start_time=start_time,
end_time=end_time,
)
def _get_wrapper_num_retries(kwargs: dict[str, Any], exception: Exception) -> tuple[int | None, dict[str, Any]]:
"""
Get the number of retries from the kwargs and the retry policy.
@ -1537,6 +1463,7 @@ def client(original_function):
start_time: Final = datetime.datetime.now()
result = None
logging_obj: LiteLLMLoggingObject | None = kwargs.get("litellm_logging_obj", None)
completion: CallCompletion | None = None
# only set litellm_call_id if its not in kwargs
if "litellm_call_id" not in kwargs:
@ -1552,7 +1479,17 @@ def client(original_function):
# Type assertion: logging_obj is guaranteed to be non-None after function_setup
assert logging_obj is not None, "logging_obj should not be None after function_setup"
from litellm.litellm_core_utils.thread_pool_executor import executor as logging_executor
completion = CallCompletion(
PythonCompletion(
logging_obj,
logging_executor,
async_call=False,
internal_call=False,
completion_with_fallbacks=False,
)
)
## LOAD CREDENTIALS
load_credentials_from_list(kwargs)
kwargs["litellm_logging_obj"] = logging_obj
@ -1650,7 +1587,12 @@ def client(original_function):
except Exception as e:
print_verbose(f"Error while checking max token limit: {e}")
# MODEL CALL
result = original_function(*args, **kwargs)
invocation_kwargs: Final = (
{**kwargs, "_litellm_call_completion": completion}
if original_function.__name__ == CallTypes.ocr.value
else kwargs
)
result = original_function(*args, **invocation_kwargs)
end_time = datetime.datetime.now()
if _is_streaming_request(
kwargs=kwargs,
@ -1704,8 +1646,11 @@ def client(original_function):
kwargs=kwargs,
)
_update_response_metadata: Final = getattr(sys.modules[__name__], "update_response_metadata")
_update_response_metadata(
verbose_logger.info("Wrapper: Completed Call, calling success_handler")
completion.success(result, start_time, end_time)
# RETURN RESULT
update_response_metadata = getattr(sys.modules[__name__], "update_response_metadata")
update_response_metadata(
result=result,
logging_obj=logging_obj,
model=model,
@ -1713,21 +1658,6 @@ def client(original_function):
start_time=start_time,
end_time=end_time,
)
# LOG SUCCESS - handle streaming success logging in the _next_ object, remove `handle_success` once it's deprecated
verbose_logger.info("Wrapper: Completed Call, calling success_handler")
# Copy the current context to propagate it to the background thread
# This is essential for OpenTelemetry span context propagation
ctx: Final = contextvars.copy_context()
executor: Final = getattr(sys.modules[__name__], "executor")
executor.submit(
ctx.run,
logging_obj.success_handler,
result,
start_time,
end_time,
)
# RETURN RESULT
return result
except Exception as e:
call_type = original_function.__name__
@ -1801,11 +1731,14 @@ def client(original_function):
end_time = datetime.datetime.now()
# LOG FAILURE - handle streaming failure logging in the _next_ object, remove `handle_failure` once it's deprecated
if logging_obj:
logging_obj.failure_handler(
e, traceback_exception, start_time, end_time
) # DO NOT MAKE THREADED - router retry fallback relies on this!
if completion is not None:
completion.failure(e, traceback_exception, start_time, end_time)
elif logging_obj:
logging_obj.failure_handler(e, traceback_exception, start_time, end_time)
raise e
finally:
if completion is not None:
completion.release()
@wraps(original_function)
async def wrapper_async(*args, **kwargs):
@ -1814,6 +1747,7 @@ def client(original_function):
result = None
_update_response_metadata: Final = getattr(sys.modules[__name__], "update_response_metadata")
logging_obj: LiteLLMLoggingObject | None = kwargs.get("litellm_logging_obj", None)
completion: CallCompletion | None = None
LLMCachingHandler: Final = _get_cached_llm_caching_handler()
_llm_caching_handler: Final[LLMCachingHandler] = LLMCachingHandler(
original_function=original_function,
@ -1839,7 +1773,15 @@ def client(original_function):
# Type assertion: logging_obj is guaranteed to be non-None after function_setup
assert logging_obj is not None, "logging_obj should not be None after function_setup"
completion = CallCompletion(
PythonCompletion(
logging_obj,
None,
async_call=True,
internal_call=_is_litellm_internal_call,
completion_with_fallbacks=is_completion_with_fallbacks,
)
)
modified_kwargs: Final = await async_pre_call_deployment_hook(kwargs, call_type)
if modified_kwargs is not None:
kwargs = modified_kwargs
@ -1924,7 +1866,12 @@ def client(original_function):
# MODEL CALL
try:
result = await original_function(*args, **kwargs)
invocation_kwargs: Final = (
{**kwargs, "_litellm_call_completion": completion}
if original_function.__name__ == CallTypes.aocr.value
else kwargs
)
result = await original_function(*args, **invocation_kwargs)
except Exception as deployment_error:
_deployment_call_end_time = datetime.datetime.now() # noqa: DTZ005 # matches the naive datetimes this whole function already times start_time/end_time with
try:
@ -1987,20 +1934,13 @@ def client(original_function):
args=args,
)
completion.success(result, start_time, end_time)
# REBUILD EMBEDDING CACHING
if (
isinstance(result, EmbeddingResponse)
and _caching_handler_response is not None
and _caching_handler_response.final_embedding_cached_response is not None
):
_dispatch_success_logging(
logging_obj=logging_obj,
result=result,
start_time=start_time,
end_time=end_time,
is_completion_with_fallbacks=is_completion_with_fallbacks,
is_litellm_internal_call=_is_litellm_internal_call,
)
return _llm_caching_handler._combine_cached_embedding_response_with_api_result(
_caching_handler_response=_caching_handler_response,
embedding_response=result,
@ -2016,14 +1956,6 @@ def client(original_function):
start_time=start_time,
end_time=end_time,
)
_dispatch_success_logging(
logging_obj=logging_obj,
result=result,
start_time=start_time,
end_time=end_time,
is_completion_with_fallbacks=is_completion_with_fallbacks,
is_litellm_internal_call=_is_litellm_internal_call,
)
return result
except Exception as e:
@ -2031,17 +1963,12 @@ def client(original_function):
# Reuse the timestamp taken right when the deployment call itself failed, before
# the failure hook ran, so a slow callback doesn't inflate the reported duration.
end_time = _deployment_call_end_time if _deployment_call_end_time is not None else datetime.datetime.now() # noqa: DTZ005 # matches the naive datetimes this whole function already times start_time/end_time with
if logging_obj and not _is_litellm_internal_call:
try:
logging_obj.failure_handler(
e, traceback_exception, start_time, end_time
) # DO NOT MAKE THREADED - router retry fallback relies on this!
except Exception as e:
raise e
try:
await logging_obj.async_failure_handler(e, traceback_exception, start_time, end_time)
except Exception as e:
raise e
if completion is not None:
completion.failure(e, traceback_exception, start_time, end_time)
await completion.async_failure(e, traceback_exception, start_time, end_time)
elif logging_obj and not _is_litellm_internal_call:
logging_obj.failure_handler(e, traceback_exception, start_time, end_time)
await logging_obj.async_failure_handler(e, traceback_exception, start_time, end_time)
call_type = original_function.__name__
num_retries, kwargs = _get_wrapper_num_retries(kwargs=kwargs, exception=e)
@ -2111,6 +2038,8 @@ def client(original_function):
raise e
finally:
if completion is not None:
completion.release()
# Restore trace_id/session_id contextvars to their pre-call value once
# this call (in this asyncio Task) is fully done - see
# request_correlation_in_logs. Unlike wrapper()'s sync path, it's safe to
@ -7024,26 +6953,7 @@ class TextCompletionStreamWrapper:
raise StopAsyncIteration
def mock_stream_usage_chunk(model_response: ModelResponseStream, model: str, prompt_tokens: int) -> ModelResponseStream:
return ModelResponseStream(
id=model_response.id,
choices=[], # mutable-ok: ModelResponseStream only treats a list as explicit choices, a tuple gets a default choice
model=model,
usage=Usage(
prompt_tokens=prompt_tokens,
completion_tokens=DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT,
total_tokens=prompt_tokens + DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT,
),
)
def mock_completion_streaming_obj(
model_response: ModelResponseStream,
mock_response: str | MockException | ModelResponseStream,
model: str,
n: int | None = None,
prompt_tokens: int | None = None,
) -> Iterator[ModelResponseStream]:
def mock_completion_streaming_obj(model_response, mock_response, model, n: int | None = None):
if isinstance(mock_response, litellm.MockException):
raise mock_response
if isinstance(mock_response, ModelResponseStream):
@ -7063,17 +6973,14 @@ def mock_completion_streaming_obj(
_all_choices.append(_streaming_choice)
model_response.choices = _all_choices
yield model_response
if prompt_tokens is not None:
yield mock_stream_usage_chunk(model_response, model=model, prompt_tokens=prompt_tokens)
async def async_mock_completion_streaming_obj(
model_response: ModelResponseStream,
model_response,
mock_response: str | MockException | ModelResponseStream,
model: str,
model,
n: int | None = None,
prompt_tokens: int | None = None,
) -> AsyncIterator[ModelResponseStream]:
):
if isinstance(mock_response, litellm.MockException):
raise mock_response
if isinstance(mock_response, ModelResponseStream):
@ -7093,8 +7000,6 @@ async def async_mock_completion_streaming_obj(
_all_choices.append(_streaming_choice)
model_response.choices = _all_choices
yield model_response
if prompt_tokens is not None:
yield mock_stream_usage_chunk(model_response, model=model, prompt_tokens=prompt_tokens)
########## Reading Config File ############################

View file

@ -0,0 +1,169 @@
import contextvars
import datetime
from collections.abc import Callable, Coroutine
from concurrent.futures import Future
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import pytest
from litellm.litellm_core_utils.call_completion import CallCompletion, PythonCompletion
class RecordingExecutor:
def __init__(self) -> None:
self.submissions: list[tuple[Callable[..., object], tuple[object, ...]]] = []
def submit(self, function: Callable[..., object], *args: object) -> Future[object]:
self.submissions.append((function, args))
future: Final[Future[object]] = Future()
future.set_result(function(*args))
return future
class RecordingCompletion:
def __init__(self) -> None:
self.successes: list[object] = []
self.failures: list[Exception] = []
def success(
self,
result: object,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None:
self.successes.append(result)
def failure(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None:
self.failures.append(exception)
async def async_failure(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None:
self.failures.append(exception)
class RecordingLogging:
def __init__(self, marker: contextvars.ContextVar[str], observed: list[tuple[object, str]]) -> None:
self._marker = marker
self._observed = observed
self._defer_async_logging = False
self._enqueue_deferred_logging: Callable[[], None] | None = None
def success_handler(
self,
result: object,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None:
self._observed.append((result, self._marker.get()))
def failure_handler(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
async def async_success_handler(
self,
result: object,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
async def async_failure_handler(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
def handle_sync_success_callbacks_for_async_calls(
self,
result: object,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
def test_python_completion_preserves_sync_context_and_response_identity() -> None:
marker: Final[contextvars.ContextVar[str]] = contextvars.ContextVar("marker")
marker.set("request-context")
executor: Final = RecordingExecutor()
response: Final = object()
observed: Final[list[tuple[object, str]]] = []
logging_obj: Final = RecordingLogging(marker, observed)
completion: Final = PythonCompletion(
logging_obj,
executor,
async_call=False,
internal_call=False,
completion_with_fallbacks=False,
)
now: Final = datetime.datetime.now(datetime.timezone.utc)
completion.success(response, now, now)
assert observed == [(response, "request-context")]
assert len(executor.submissions) == 1
@pytest.mark.asyncio
async def test_call_completion_attaches_once_and_forwards_final_objects() -> None:
python_completion: Final = RecordingCompletion()
native_completion: Final = RecordingCompletion()
ignored_completion: Final = RecordingCompletion()
completion: Final = CallCompletion(python_completion)
response: Final = object()
error: Final = ValueError("mapped failure")
now: Final = datetime.datetime.now(datetime.timezone.utc)
assert completion.python_implementation is python_completion
assert completion.attach(native_completion)
assert not completion.attach(ignored_completion)
completion.success(response, now, now)
completion.failure(error, "traceback", now, now)
await completion.async_failure(error, "traceback", now, now)
assert native_completion.successes == [response]
assert native_completion.failures == [error, error]
assert python_completion.successes == []
assert ignored_completion.successes == []
@pytest.mark.asyncio
async def test_python_completion_retains_deferred_success_arguments(monkeypatch: pytest.MonkeyPatch) -> None:
response: Final = object()
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,
)
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)
logging_obj._enqueue_deferred_logging()
assert len(scheduled) == 1
scheduled[0].close()