diff --git a/litellm/_internal_context.py b/litellm/_internal_context.py index 389add8ed0f..845c2a6c84c 100644 --- a/litellm/_internal_context.py +++ b/litellm/_internal_context.py @@ -6,6 +6,8 @@ be settable from user input. Context variables are scoped to the current asyncio task and cannot be injected via HTTP request bodies. """ +import asyncio +import weakref from collections.abc import Generator from contextlib import contextmanager from contextvars import ContextVar @@ -23,6 +25,10 @@ _billing_time: Final[ContextVar[datetime | None]] = ContextVar("billing_time", d _post_response: Final[ContextVar[bool]] = ContextVar("post_response", default=False) +_provider_calls_in_flight: Final[ContextVar[frozenset[tuple[weakref.ref[asyncio.Task[object]], int]]]] = ContextVar( + "provider_calls_in_flight", default=frozenset() +) + @contextmanager def post_response_phase() -> Generator[None]: @@ -38,6 +44,22 @@ def in_post_response_phase() -> bool: return _post_response.get() +@contextmanager +def provider_call_scope(logging_obj: object) -> Generator[bool]: + """Yield whether an outer call in this same task, e.g. a Responses bridge, is already making this logging object's provider call.""" + task: Final = asyncio.current_task() + if task is None: + yield False + return + key: Final = (weakref.ref(task), id(logging_obj)) + in_flight: Final = _provider_calls_in_flight.get() + token: Final = _provider_calls_in_flight.set(in_flight.union((key,))) + try: + yield key in in_flight + finally: + _provider_calls_in_flight.reset(token) + + @contextmanager def pinned_billing_time(moment: datetime) -> Generator[None]: """Price every rate lookup inside this block at ``moment`` rather than at each one's own clock read.""" diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index e292ab7b2ec..20777c9cb42 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -75,6 +75,7 @@ from litellm.litellm_core_utils.core_helpers import ( reconstruct_model_name, set_response_cost_in_hidden_params, ) +from litellm.litellm_core_utils.coroutine_checker import coroutine_checker from litellm.litellm_core_utils.error_normalization import normalize_error from litellm.litellm_core_utils.get_litellm_params import get_litellm_params from litellm.litellm_core_utils.internal_call_metadata import ( @@ -1338,6 +1339,37 @@ class Logging(LiteLLMLoggingBaseClass): additional_args.get("api_base", "") ) + async def async_pre_call(self) -> None: + candidates: Final = ( + *(callback for callback in litellm._async_input_callback if callable(callback)), + *( + callback + for callback in litellm.input_callback + if callable(callback) and coroutine_checker.is_async_callable(callback) + ), + ) + callbacks: Final = tuple( + callback + for index, callback in enumerate(candidates) + if not any(earlier is callback for earlier in candidates[:index]) + ) + if not callbacks: + return + + request: Final = { # mutable-ok: callbacks receive the plain kwargs dict they always got + **self.model_call_details, + "model": self.model, + "messages": self.messages, + "log_event_type": "pre_api_call", + } + for callback in callbacks: + try: + await callback(request) + except Exception as e: + verbose_logger.exception( + "LiteLLM.LoggingError: [Non-Blocking] Exception occurred while async pre-call logging %s", e + ) + def pre_call(self, input, api_key, model=None, additional_args={}): # Log the exact input to the LLM API try: @@ -1446,7 +1478,11 @@ class Logging(LiteLLMLoggingBaseClass): messages=self.messages, kwargs=self.model_call_details, ) - elif callable(callback) and customLogger is not None: # custom logger functions + elif ( + callable(callback) + and customLogger is not None + and not coroutine_checker.is_async_callable(callback) + ): # custom logger functions customLogger.log_input_event( model=self.model, messages=self.messages, diff --git a/litellm/utils.py b/litellm/utils.py index 71186a28be4..b0419aca270 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -52,7 +52,7 @@ import litellm.litellm_core_utils # audio_utils.utils is lazy-loaded - only imported when needed for transcription calls import litellm.litellm_core_utils.json_validation_rule -from litellm._internal_context import is_internal_call +from litellm._internal_context import is_internal_call, provider_call_scope from litellm._lazy_imports import ( _get_default_encoding, _get_modified_max_tokens, @@ -2067,7 +2067,15 @@ def client(original_function): else kwargs ) try: - result = await original_function(*args, **call_kwargs) + if litellm._async_input_callback or ( + litellm.input_callback and any(check_coroutine(callback) for callback in litellm.input_callback) + ): + with provider_call_scope(logging_obj) as nested_provider_call: + if not nested_provider_call: + await logging_obj.async_pre_call() + result = await original_function(*args, **call_kwargs) + else: + result = await original_function(*args, **call_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: diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index 52b1714b4d9..4cd0afb0ec6 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -2,13 +2,17 @@ import asyncio import base64 import contextlib import contextvars +import functools +import gc import io import json import logging import os import queue import threading -from collections.abc import Callable, Iterator, Mapping +import warnings +import weakref +from collections.abc import Callable, Coroutine, Iterator, Mapping from concurrent.futures import Future, ThreadPoolExecutor from datetime import datetime, timedelta, timezone from pathlib import PurePath @@ -36,6 +40,7 @@ from litellm.constants import DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.get_litellm_params import get_litellm_params +from litellm.litellm_core_utils.litellm_logging import Logging from litellm.litellm_core_utils.thread_pool_executor import executor as logging_executor from litellm.llms.base_llm.base_model_iterator import MockResponseIterator from litellm.proxy.utils import is_valid_api_key @@ -5407,6 +5412,366 @@ async def test_wrapper_async_claims_the_budget_reservation_a_supplied_logging_ob assert reservation["callback_bound"] is True +@pytest.mark.asyncio +@pytest.mark.parametrize("registration", ["manager", "input_callback"]) +async def test_wrapper_async_awaits_async_input_callback_before_the_provider_call( + monkeypatch: pytest.MonkeyPatch, registration: str +) -> None: + events: Final[list[tuple[str, str]]] = [] + messages: Final = [{"role": "user", "content": "hello"}] + + async def async_callback(kwargs: dict) -> None: + assert kwargs["model"] == "gpt-4o-mini" + assert kwargs["messages"] == messages + events.append(("async", kwargs["log_event_type"])) + + def sync_callback(kwargs: dict) -> None: + events.append(("sync", kwargs["log_event_type"])) + + monkeypatch.setattr("litellm.litellm_core_utils.litellm_logging.customLogger", CustomLogger()) + + if registration == "manager": + litellm.logging_callback_manager.add_litellm_input_callback(async_callback) + else: + monkeypatch.setattr(litellm, "input_callback", [async_callback]) + litellm.logging_callback_manager.add_litellm_input_callback(sync_callback) + + response: Final = await litellm.acompletion(model="gpt-4o-mini", messages=messages, mock_response="ok") + + assert response.choices[0].message.content == "ok" + assert events == [("async", "pre_api_call"), ("sync", "pre_api_call")] + + +@pytest.mark.asyncio +async def test_wrapper_async_continues_when_an_async_input_callback_raises( + monkeypatch: pytest.MonkeyPatch, +) -> None: + events: Final[list[str]] = [] + + async def failing_callback(kwargs: dict) -> None: + events.append("failing") + raise RuntimeError("callback failed") + + async def following_callback(kwargs: dict) -> None: + events.append(kwargs["log_event_type"]) + + monkeypatch.setattr(litellm, "_async_input_callback", [failing_callback, following_callback]) + response: Final = await litellm.acompletion( + model="gpt-4o-mini", messages=[{"role": "user", "content": "hello"}], mock_response="ok" + ) + + assert events == ["failing", "pre_api_call"] + assert response.choices[0].message.content == "ok" + + +@pytest.mark.asyncio +async def test_wrapper_async_skips_async_input_callback_on_cache_hit( + monkeypatch: pytest.MonkeyPatch, +) -> None: + events: Final[list[str]] = [] + + async def async_callback(kwargs: dict) -> None: + events.append("async") + + def sync_callback(kwargs: dict) -> None: + events.append("sync") + + monkeypatch.setattr("litellm.litellm_core_utils.litellm_logging.customLogger", CustomLogger()) + + monkeypatch.setattr(litellm, "_async_input_callback", [async_callback]) + monkeypatch.setattr(litellm, "input_callback", [sync_callback]) + monkeypatch.setattr(litellm, "cache", Cache(type="local")) + request: Final = { + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hello"}], + "mock_response": "ok", + "caching": True, + } + + first: Final = await litellm.acompletion(**request) + await asyncio.gather(*tuple(_PENDING_CACHE_WRITES)) + second: Final = await litellm.acompletion(**request) + + assert first.choices[0].message.content == second.choices[0].message.content == "ok" + assert events == ["async", "sync"] + + +def _async_input_callback_logging_object( + original_function: str, litellm_call_id: str +) -> tuple[Logging, Mapping[str, object]]: + return litellm.utils.function_setup( + original_function=original_function, + rules_obj=litellm.utils.Rules(), + start_time=datetime.now(), + model="gpt-4o-mini", + messages=[{"role": "user", "content": "hello"}], + litellm_call_id=litellm_call_id, + ) + + +@pytest.mark.asyncio +async def test_wrapper_async_reruns_async_input_callback_when_a_failed_attempt_is_retried_on_the_same_logging_object( + monkeypatch: pytest.MonkeyPatch, +) -> None: + events: Final[list[object]] = [] + + async def callback(kwargs: Mapping[str, object]) -> None: + events.append(kwargs["log_event_type"]) + + monkeypatch.setattr(litellm, "_async_input_callback", [callback]) + logging_obj, kwargs = _async_input_callback_logging_object("acompletion", "retried-logging-object") + rate_limited: Final = litellm.RateLimitError(message="busy", llm_provider="openai", model="gpt-4o-mini") + + with pytest.raises(litellm.RateLimitError): + await litellm.acompletion(**kwargs, litellm_logging_obj=logging_obj, mock_response=rate_limited) + retried: Final = await litellm.acompletion(**kwargs, litellm_logging_obj=logging_obj, mock_response="ok") + + assert retried.choices[0].message.content == "ok" + assert events == ["pre_api_call", "pre_api_call"] + + +@pytest.mark.asyncio +async def test_wrapper_async_runs_async_input_callback_once_when_responses_bridges_to_acompletion( + monkeypatch: pytest.MonkeyPatch, +) -> None: + call_ids: Final[list[object]] = [] + + async def callback(kwargs: Mapping[str, object]) -> None: + call_ids.append(kwargs["litellm_call_id"]) + + monkeypatch.setattr(litellm, "_async_input_callback", [callback]) + + response: Final = await litellm.aresponses( + model="openai/gpt-4o-mini", + input="hello", + use_chat_completions_api=True, + mock_response=ModelResponse(choices=[{"message": {"role": "assistant", "content": "bridged"}}]), + ) + + assert response.output[0].content[0].text == "bridged" + assert len(call_ids) == 1 + + +@pytest.mark.asyncio +async def test_wrapper_async_awaits_async_input_callback_appended_directly_when_the_logging_object_is_supplied( + monkeypatch: pytest.MonkeyPatch, +) -> None: + events: Final[list[tuple[str, object]]] = [] + + async def async_callback(kwargs: Mapping[str, object]) -> None: + events.append(("async", kwargs["log_event_type"])) + + def sync_callback(kwargs: Mapping[str, object]) -> None: + events.append(("sync", kwargs["log_event_type"])) + + monkeypatch.setattr("litellm.litellm_core_utils.litellm_logging.customLogger", CustomLogger()) + logging_obj, kwargs = _async_input_callback_logging_object("acompletion", "supplied-logging-object") + monkeypatch.setattr(litellm, "input_callback", [async_callback, sync_callback]) + + response: Final = await litellm.acompletion(**kwargs, litellm_logging_obj=logging_obj, mock_response="ok") + + assert response.choices[0].message.content == "ok" + assert events == [("async", "pre_api_call"), ("sync", "pre_api_call")] + + +def test_completion_does_not_leave_an_async_input_callback_unawaited(monkeypatch: pytest.MonkeyPatch) -> None: + events: Final[list[object]] = [] + + async def async_input_callback(kwargs: Mapping[str, object]) -> None: + events.append("async") + + def sync_callback(kwargs: Mapping[str, object]) -> None: + events.append("sync") + + monkeypatch.setattr("litellm.litellm_core_utils.litellm_logging.customLogger", CustomLogger()) + logging_obj, kwargs = _async_input_callback_logging_object("completion", "sync-supplied-logging-object") + monkeypatch.setattr(litellm, "input_callback", [async_input_callback, sync_callback]) + + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + response: Final = litellm.completion(**kwargs, litellm_logging_obj=logging_obj, mock_response="ok") + gc.collect() + + assert response.choices[0].message.content == "ok" + assert events == ["sync"] + assert [str(w.message) for w in caught if async_input_callback.__name__ in str(w.message)] == [] + + +@pytest.mark.asyncio +async def test_async_pre_call_hands_callbacks_the_request_without_changing_the_logging_object( + monkeypatch: pytest.MonkeyPatch, +) -> None: + received: Final[list[tuple[object, object, object]]] = [] + + async def callback(kwargs: Mapping[str, object]) -> None: + received.append((kwargs["model"], kwargs["messages"], kwargs["log_event_type"])) + + monkeypatch.setattr(litellm, "_async_input_callback", [callback]) + logging_obj, _ = _async_input_callback_logging_object("acompletion", "untouched-logging-object") + before: Final = dict(logging_obj.model_call_details) + + await logging_obj.async_pre_call() + + assert received == [("gpt-4o-mini", [{"role": "user", "content": "hello"}], "pre_api_call")] + assert logging_obj.model_call_details == before + + +@pytest.mark.asyncio +async def test_wrapper_async_runs_async_input_callback_for_concurrent_and_later_calls_sharing_a_logging_object( + monkeypatch: pytest.MonkeyPatch, +) -> None: + events: Final[list[object]] = [] + + async def callback(kwargs: Mapping[str, object]) -> None: + events.append(kwargs["log_event_type"]) + await asyncio.sleep(0) + + monkeypatch.setattr(litellm, "_async_input_callback", [callback]) + logging_obj, kwargs = _async_input_callback_logging_object("acompletion", "shared-logging-object") + + concurrent: Final = await asyncio.gather( + litellm.acompletion(**kwargs, litellm_logging_obj=logging_obj, mock_response="first"), + litellm.acompletion(**kwargs, litellm_logging_obj=logging_obj, mock_response="second"), + ) + reused: Final = await litellm.acompletion(**kwargs, litellm_logging_obj=logging_obj, mock_response="third") + + assert [response.choices[0].message.content for response in (*concurrent, reused)] == ["first", "second", "third"] + assert events == ["pre_api_call", "pre_api_call", "pre_api_call"] + + +@pytest.mark.asyncio +async def test_wrapper_async_runs_async_input_callback_for_a_child_task_reusing_the_logging_object_after_the_parent_call( + monkeypatch: pytest.MonkeyPatch, +) -> None: + events: Final[list[object]] = [] + parent_done: Final = asyncio.Event() + children: Final[list[asyncio.Task[ModelResponse | CustomStreamWrapper]]] = [] + logging_obj, kwargs = _async_input_callback_logging_object("acompletion", "child-task-logging-object") + + async def child_call() -> ModelResponse | CustomStreamWrapper: + await parent_done.wait() + return await litellm.acompletion(**kwargs, litellm_logging_obj=logging_obj, mock_response="child") + + async def callback(kwargs: Mapping[str, object]) -> None: + events.append(kwargs["log_event_type"]) + if not children: + children.append(asyncio.create_task(child_call())) + + monkeypatch.setattr(litellm, "_async_input_callback", [callback]) + + parent: Final = await litellm.acompletion(**kwargs, litellm_logging_obj=logging_obj, mock_response="parent") + parent_done.set() + child: Final = await children[0] + + assert (parent.choices[0].message.content, child.choices[0].message.content) == ("parent", "child") + assert events == ["pre_api_call", "pre_api_call"] + + +@pytest.mark.asyncio +async def test_wrapper_async_does_not_keep_the_parent_task_alive_for_a_child_task_that_outlives_it( + monkeypatch: pytest.MonkeyPatch, +) -> None: + release_child: Final = asyncio.Event() + children: Final[list[asyncio.Task[ModelResponse | CustomStreamWrapper]]] = [] + logging_obj, kwargs = _async_input_callback_logging_object("acompletion", "outlived-parent-task") + + async def child_call() -> ModelResponse | CustomStreamWrapper: + await release_child.wait() + return await litellm.acompletion(**kwargs, litellm_logging_obj=logging_obj, mock_response="child") + + async def callback(request: Mapping[str, object]) -> None: + if not children: + children.append(asyncio.create_task(child_call())) + + async def run_parent() -> tuple[ + ModelResponse | CustomStreamWrapper, weakref.ref[asyncio.Task[ModelResponse | CustomStreamWrapper]] + ]: + parent_task: Final = asyncio.create_task( + litellm.acompletion(**kwargs, litellm_logging_obj=logging_obj, mock_response="parent") + ) + return await parent_task, weakref.ref(parent_task) + + monkeypatch.setattr(litellm, "_async_input_callback", [callback]) + + parent, parent_task_ref = await run_parent() + await asyncio.sleep(0) + gc.collect() + + assert parent_task_ref() is None + assert not children[0].done() + release_child.set() + child: Final = await children[0] + assert (parent.choices[0].message.content, child.choices[0].message.content) == ("parent", "child") + + +@pytest.mark.asyncio +async def test_wrapper_async_awaits_registered_async_input_callbacks_that_are_not_coroutine_functions( + monkeypatch: pytest.MonkeyPatch, +) -> None: + events: Final[list[object]] = [] + + async def record(label: str, kwargs: Mapping[str, object]) -> None: + events.append(f"{label}:{kwargs['log_event_type']}") + + def returns_coroutine(kwargs: Mapping[str, object]) -> Coroutine[object, object, None]: + return record("sync-returning-coroutine", kwargs) + + monkeypatch.setattr(litellm, "_async_input_callback", [functools.partial(record, "partial"), returns_coroutine]) + + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + response: Final = await litellm.acompletion( + model="gpt-4o-mini", messages=[{"role": "user", "content": "hello"}], mock_response="ok" + ) + gc.collect() + + assert response.choices[0].message.content == "ok" + assert events == ["partial:pre_api_call", "sync-returning-coroutine:pre_api_call"] + assert [str(w.message) for w in caught if "never awaited" in str(w.message)] == [] + + +@pytest.mark.asyncio +async def test_wrapper_async_runs_an_async_input_callback_registered_in_both_lists_once( + monkeypatch: pytest.MonkeyPatch, +) -> None: + events: Final[list[object]] = [] + + async def callback(kwargs: Mapping[str, object]) -> None: + events.append(kwargs["log_event_type"]) + + monkeypatch.setattr(litellm, "_async_input_callback", [callback]) + monkeypatch.setattr(litellm, "input_callback", [callback]) + + response: Final = await litellm.acompletion( + model="gpt-4o-mini", messages=[{"role": "user", "content": "hello"}], mock_response="ok" + ) + + assert response.choices[0].message.content == "ok" + assert events == ["pre_api_call"] + + +@pytest.mark.asyncio +async def test_wrapper_async_does_not_need_async_pre_call_on_a_supplied_logging_object_with_only_sync_input_callbacks( + monkeypatch: pytest.MonkeyPatch, +) -> None: + response: Final = object() + logger: Final = MagicMock() + del logger.async_pre_call + + async def acompletion(**kwargs: object) -> object: + return response + + def sync_callback(kwargs: Mapping[str, object]) -> None: + pass + + monkeypatch.setattr(litellm, "_async_input_callback", []) + monkeypatch.setattr(litellm, "input_callback", [sync_callback]) + + result: Final = await client(acompletion)(model="gpt-4o-mini", litellm_logging_obj=logger) + + assert result is response + + @pytest.mark.asyncio async def test_wrapper_async_hands_the_budget_reservation_back_when_the_call_fails() -> None: reservation = _budget_reservation()