mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(logging): run async input callbacks once per provider attempt
This commit is contained in:
parent
6eeb1b1c77
commit
ef80f9da19
4 changed files with 332 additions and 28 deletions
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -74,6 +74,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,15 +1339,31 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
)
|
||||
|
||||
async def async_pre_call(self) -> None:
|
||||
if self.model_call_details.get("has_logged_async_pre_call"):
|
||||
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
|
||||
|
||||
self.model_call_details["has_logged_async_pre_call"] = True
|
||||
self.model_call_details.update(model=self.model, messages=self.messages, log_event_type="pre_api_call")
|
||||
for callback in litellm._async_input_callback:
|
||||
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:
|
||||
if callable(callback):
|
||||
await callback(self.model_call_details)
|
||||
await callback(request)
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while async pre-call logging %s", e
|
||||
|
|
@ -1460,7 +1477,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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
@ -2079,9 +2079,15 @@ def client(original_function):
|
|||
else kwargs
|
||||
)
|
||||
try:
|
||||
if litellm._async_input_callback:
|
||||
await logging_obj.async_pre_call()
|
||||
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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -35,6 +39,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
|
||||
|
|
@ -5366,32 +5371,282 @@ async def test_wrapper_async_skips_async_input_callback_on_cache_hit(
|
|||
assert events == ["async", "sync"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_wrapper_async_runs_async_input_callback_once_per_logging_object(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
events: Final[list[str]] = []
|
||||
|
||||
async def callback(kwargs: dict) -> None:
|
||||
events.append(kwargs["log_event_type"])
|
||||
|
||||
monkeypatch.setattr(litellm, "_async_input_callback", [callback])
|
||||
logging_obj, kwargs = litellm.utils.function_setup(
|
||||
original_function="acompletion",
|
||||
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="shared-logging-object",
|
||||
litellm_call_id=litellm_call_id,
|
||||
)
|
||||
|
||||
first: Final = await litellm.acompletion(**kwargs, litellm_logging_obj=logging_obj, mock_response="ok")
|
||||
second: Final = await litellm.acompletion(**kwargs, litellm_logging_obj=logging_obj, mock_response="ok")
|
||||
|
||||
assert first.choices[0].message.content == second.choices[0].message.content == "ok"
|
||||
@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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue