fix(logging): run async input callbacks once per provider attempt

This commit is contained in:
sharosoo 2026-09-28 22:22:54 +09:00
parent 6eeb1b1c77
commit ef80f9da19
4 changed files with 332 additions and 28 deletions

View file

@ -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."""

View file

@ -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,

View file

@ -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:

View file

@ -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()