fix(callbacks): scope per-call callbacks to requests

This commit is contained in:
Yujong Lee 2026-09-19 13:59:18 -07:00
parent a93bfdc749
commit 8d8094573e
5 changed files with 214 additions and 29 deletions

View file

@ -276,11 +276,9 @@ _UNTRUSTED_ROOT_CONTROL_FIELDS: Final = (
"pillar_response_headers",
"_guardrail_pipelines",
"_pipeline_managed_guardrails",
# Callback-registration fields. ``callbacks``, ``service_callback``,
# and ``logger_fn`` are read by ``litellm.utils.function_setup`` and
# appended to process-wide ``litellm.{input,success,failure,_async_*,
# service}_callback`` lists / ``litellm.user_logger_fn`` — one request
# poisons the worker for every subsequent caller.
# Callback-control fields are proxy-owned. Letting request bodies select
# callbacks, service callbacks, or logger functions would let callers
# choose server integrations and logging destinations.
# ``litellm_disabled_callbacks`` is the inverse primitive: the
# legitimate path reads it from key/team metadata, the request-body
# version silently turns off admin-configured audit/observability

View file

@ -571,15 +571,23 @@ def custom_llm_setup():
litellm._custom_providers.append(custom_llm["provider"])
def _add_custom_logger_callback_to_specific_event(callback: str, logging_event: Literal["success", "failure"]) -> None:
"""
Add a custom logger callback to the specific event
"""
def _initialize_custom_logger_callback(callback: str) -> CustomLogger | None:
from litellm import _custom_logger_compatible_callbacks_literal
from litellm.litellm_core_utils.litellm_logging import (
_init_custom_logger_compatible_class,
)
return _init_custom_logger_compatible_class(
cast(_custom_logger_compatible_callbacks_literal, callback),
internal_usage_cache=None,
llm_router=None,
)
def _add_custom_logger_callback_to_specific_event(callback: str, logging_event: Literal["success", "failure"]) -> None:
"""
Add a custom logger callback to the specific event
"""
if callback not in litellm._known_custom_logger_compatible_callbacks:
verbose_logger.debug(
"Callback %s is not a valid custom logger compatible callback. Known list - %s",
@ -588,11 +596,7 @@ def _add_custom_logger_callback_to_specific_event(callback: str, logging_event:
)
return
callback_class: Final = _init_custom_logger_compatible_class(
cast(_custom_logger_compatible_callbacks_literal, callback),
internal_usage_cache=None,
llm_router=None,
)
callback_class: Final = _initialize_custom_logger_callback(callback)
if callback_class:
if logging_event == "success" and _custom_logger_class_exists_in_success_callbacks(callback_class) is False:
@ -694,11 +698,26 @@ def load_credentials_from_list(kwargs: dict):
def get_dynamic_callbacks(
dynamic_callbacks: list[str | Callable | CustomLogger] | None,
) -> list:
returned_callbacks: Final = litellm.callbacks.copy()
if dynamic_callbacks:
returned_callbacks.extend(dynamic_callbacks)
return returned_callbacks
) -> list[Callable | CustomLogger] | None:
if not dynamic_callbacks:
return None
initialized_callbacks: Final = tuple(
(_initialize_custom_logger_callback(callback) if isinstance(callback, str) else callback)
for callback in dynamic_callbacks
)
return [
callback
for index, callback in enumerate(initialized_callbacks)
if callback is not None and callback not in initialized_callbacks[:index]
]
def _merge_dynamic_callbacks(
callbacks: Sequence[str | Callable | CustomLogger] | None,
additional_callbacks: Sequence[str | Callable | CustomLogger] | None,
) -> list[str | Callable | CustomLogger] | None:
merged_callbacks: Final = (*(callbacks or ()), *(additional_callbacks or ()))
return list(merged_callbacks) if merged_callbacks else None
def _is_gemini_model(model: str | None, custom_llm_provider: str | None) -> bool:
@ -912,7 +931,13 @@ def function_setup(
## DYNAMIC CALLBACKS ##
dynamic_callbacks: Final[list[str | Callable | CustomLogger] | None] = kwargs.pop("callbacks", None)
all_callbacks: Final = get_dynamic_callbacks(dynamic_callbacks=dynamic_callbacks)
request_callbacks: Final = get_dynamic_callbacks(dynamic_callbacks=dynamic_callbacks)
request_sync_callbacks: Final = (
[callback for callback in request_callbacks if not coroutine_checker.is_async_callable(callback)]
if request_callbacks
else None
)
all_callbacks: Final = litellm.callbacks.copy()
if len(all_callbacks) > 0:
for callback in all_callbacks:
@ -990,7 +1015,7 @@ def function_setup(
dynamic_success_callbacks: list[str | Callable | CustomLogger] | None = None
dynamic_async_success_callbacks: list[str | Callable | CustomLogger] | None = None
dynamic_failure_callbacks: list[str | Callable | CustomLogger] | None = None
dynamic_async_failure_callbacks: Final[list[str | Callable | CustomLogger] | None] = None
dynamic_async_failure_callbacks: list[str | Callable | CustomLogger] | None = None
if kwargs.get("success_callback", None) is not None and isinstance(kwargs["success_callback"], list):
removed_async_items = []
for index, callback in enumerate(kwargs["success_callback"]):
@ -1009,6 +1034,16 @@ def function_setup(
if kwargs.get("failure_callback", None) is not None and isinstance(kwargs["failure_callback"], list):
dynamic_failure_callbacks = kwargs.pop("failure_callback")
dynamic_input_callbacks: Final = (
[callback for callback in request_sync_callbacks if callback not in litellm.input_callback]
if request_sync_callbacks
else None
)
dynamic_success_callbacks = _merge_dynamic_callbacks(request_sync_callbacks, dynamic_success_callbacks)
dynamic_async_success_callbacks = _merge_dynamic_callbacks(request_callbacks, dynamic_async_success_callbacks)
dynamic_failure_callbacks = _merge_dynamic_callbacks(request_sync_callbacks, dynamic_failure_callbacks)
dynamic_async_failure_callbacks = _merge_dynamic_callbacks(request_callbacks, dynamic_async_failure_callbacks)
if add_breadcrumb:
try:
from litellm.litellm_core_utils.core_helpers import safe_deep_copy
@ -1188,6 +1223,7 @@ def function_setup(
function_id=function_id or "",
call_type=call_type,
start_time=start_time,
dynamic_input_callbacks=dynamic_input_callbacks,
dynamic_success_callbacks=dynamic_success_callbacks,
dynamic_failure_callbacks=dynamic_failure_callbacks,
dynamic_async_success_callbacks=dynamic_async_success_callbacks,

View file

@ -1107,14 +1107,10 @@ async def test_key_metadata_enable_prompt_caching_promoted_to_request_root(key_v
async def test_add_litellm_data_to_request_strips_callback_control_fields(
control_field,
):
"""``callbacks`` / ``service_callback`` / ``logger_fn`` get appended to
the worker-wide ``litellm.{input,success,failure,_async_*,service}_callback``
lists and ``litellm.user_logger_fn`` from inside ``function_setup`` —
one request poisons every subsequent caller in that worker.
``litellm_disabled_callbacks`` is the inverse: a request-body value
silently disables admin-configured audit/observability for the call.
None has a documented per-request use, so all four are stripped at
the proxy boundary alongside the existing internal-only fields."""
"""Callback controls are proxy-owned, including the request-scoped
``callbacks`` SDK parameter. A request-body ``litellm_disabled_callbacks``
value could also disable admin-configured audit/observability for the call,
so the proxy strips every callback control field."""
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()

View file

@ -3922,6 +3922,130 @@ def test_custom_logger_guards_ignore_subclass_instances(monkeypatch: pytest.Monk
assert _custom_logger_class_exists_in_failure_callbacks(builtin_instance) is True
@pytest.mark.asyncio
async def test_completion_callbacks_are_request_scoped_while_registered_callbacks_remain_global(
monkeypatch: pytest.MonkeyPatch,
) -> None:
request_ids: Final = tuple(f"request-{index}" for index in range(8))
global_calls: Final[list[str]] = []
observed: Final = {request_id: [] for request_id in request_ids}
class Observe(CustomLogger):
def __init__(self, calls: list[str]) -> None:
super().__init__()
self.calls: Final = calls
def log_pre_api_call(self, model, messages, kwargs):
self.calls.append(kwargs["litellm_call_id"])
global_callback: Final = Observe(global_calls)
request_callbacks: Final = {request_id: Observe(observed[request_id]) for request_id in request_ids}
monkeypatch.setattr(litellm, "callbacks", [global_callback])
for attribute in (
"input_callback",
"success_callback",
"failure_callback",
"_async_input_callback",
"_async_success_callback",
"_async_failure_callback",
):
monkeypatch.setattr(litellm, attribute, [])
monkeypatch.setattr(litellm.utils, "callback_list", [])
async def call(request_id: str) -> None:
await litellm.acompletion(
model="gpt-5.6",
messages=[{"role": "user", "content": request_id}],
mock_response="ok",
mock_delay=0.01,
litellm_call_id=request_id,
callbacks=[request_callbacks[request_id]],
)
await asyncio.gather(*(call(request_id) for request_id in request_ids))
await litellm.acompletion(
model="gpt-5.6",
messages=[{"role": "user", "content": "without callback"}],
mock_response="ok",
litellm_call_id="without-callback",
)
await litellm.acompletion(
model="gpt-5.6",
messages=[{"role": "user", "content": "registered and requested"}],
mock_response="ok",
litellm_call_id="registered-and-requested",
callbacks=[global_callback],
)
assert len(global_calls) == len(request_ids) + 2
assert set(global_calls) == {*request_ids, "without-callback", "registered-and-requested"}
assert observed == {request_id: [request_id] for request_id in request_ids}
assert all(
callback not in getattr(litellm, attribute)
for callback in request_callbacks.values()
for attribute in (
"callbacks",
"input_callback",
"success_callback",
"failure_callback",
"_async_input_callback",
"_async_success_callback",
"_async_failure_callback",
)
)
def test_completion_callback_does_not_leak_to_the_next_sync_request(
monkeypatch: pytest.MonkeyPatch,
) -> None:
calls: Final[list[str]] = []
class Observe(CustomLogger):
def log_pre_api_call(self, model, messages, kwargs):
calls.append(kwargs["litellm_call_id"])
callback: Final = Observe()
monkeypatch.setattr(litellm, "callbacks", [])
for attribute in (
"input_callback",
"success_callback",
"failure_callback",
"_async_input_callback",
"_async_success_callback",
"_async_failure_callback",
):
monkeypatch.setattr(litellm, attribute, [])
monkeypatch.setattr(litellm.utils, "callback_list", [])
litellm.completion(
model="gpt-5.6",
messages=[{"role": "user", "content": "with callback"}],
mock_response="ok",
litellm_call_id="with-callback",
callbacks=[callback],
)
litellm.completion(
model="gpt-5.6",
messages=[{"role": "user", "content": "without callback"}],
mock_response="ok",
litellm_call_id="without-callback",
)
assert calls == ["with-callback"]
assert all(
callback not in getattr(litellm, attribute)
for attribute in (
"callbacks",
"input_callback",
"success_callback",
"failure_callback",
"_async_input_callback",
"_async_success_callback",
"_async_failure_callback",
)
)
@pytest.mark.asyncio
async def test_s3_v2_success_callback_registers_alongside_user_subclass(
monkeypatch: pytest.MonkeyPatch,

View file

@ -44,6 +44,37 @@ async def call_native_aocr_with_callbacks(server: RecordingServer, callbacks: li
return await call_native_aocr(server, callbacks=callbacks, **kwargs)
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"])
async def test_native_ocr_callbacks_are_request_scoped(ocr_server: RecordingServer, asynchronous: bool) -> None:
request_ids: Final = tuple(f"request-{index}" for index in range(8))
observed: Final = {request_id: [] for request_id in request_ids}
ocr_server.expected_requests = len(request_ids) + 1
class Observe(CustomLogger):
def __init__(self, request_id: str) -> None:
super().__init__()
self.request_id: Final = request_id
def log_pre_api_call(self, model, messages, kwargs):
observed[self.request_id].append(kwargs["litellm_call_id"])
await asyncio.gather(
*(
call_native(
ocr_server,
asynchronous,
callbacks=[Observe(request_id)],
litellm_call_id=request_id,
)
for request_id in request_ids
)
)
await call_native(ocr_server, asynchronous, litellm_call_id="without-callback")
assert observed == {request_id: [request_id] for request_id in request_ids}
def test_native_ocr_pre_call_callback_receives_transformed_provider_request(ocr_server: RecordingServer) -> None:
observations: Final = []