mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
refactor(ocr): extract call completion boundary
This commit is contained in:
parent
83ab0113f0
commit
5ac9eec465
5 changed files with 464 additions and 155 deletions
229
litellm/litellm_core_utils/call_completion.py
Normal file
229
litellm/litellm_core_utils/call_completion.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
215
litellm/utils.py
215
litellm/utils.py
|
|
@ -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 ############################
|
||||
|
|
|
|||
169
tests/test_litellm/litellm_core_utils/test_call_completion.py
Normal file
169
tests/test_litellm/litellm_core_utils/test_call_completion.py
Normal 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()
|
||||
Loading…
Add table
Reference in a new issue