From 8d8094573ecd18012b7318b67cf73f2de67a4091 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sat, 19 Sep 2026 13:59:18 -0700 Subject: [PATCH] fix(callbacks): scope per-call callbacks to requests --- litellm/proxy/litellm_pre_call_utils.py | 8 +- litellm/utils.py | 68 +++++++--- .../proxy/test_litellm_pre_call_utils.py | 12 +- tests/test_litellm/test_utils.py | 124 ++++++++++++++++++ tests/test_litellm_rust/ocr/test_callbacks.py | 31 +++++ 5 files changed, 214 insertions(+), 29 deletions(-) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 9a973755894..8121a65aff9 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -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 diff --git a/litellm/utils.py b/litellm/utils.py index 48d13bc16af..0b08224715c 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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, diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 88d38d74f49..ed9d154869d 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -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() diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index b40c10de428..64fb8678d18 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -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, diff --git a/tests/test_litellm_rust/ocr/test_callbacks.py b/tests/test_litellm_rust/ocr/test_callbacks.py index ac4a1a11a80..e1befab9b17 100644 --- a/tests/test_litellm_rust/ocr/test_callbacks.py +++ b/tests/test_litellm_rust/ocr/test_callbacks.py @@ -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 = []