mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-25 01:02:15 +00:00
fix(callbacks): scope per-call callbacks to requests
This commit is contained in:
parent
a93bfdc749
commit
8d8094573e
5 changed files with 214 additions and 29 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 = []
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue