mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge 70e5db056c into 3930c5bab6
This commit is contained in:
commit
6790292945
5 changed files with 181 additions and 18 deletions
|
|
@ -3897,7 +3897,9 @@ class ProxyBaseLLMRequestProcessing:
|
|||
# pure overhead on the streaming hot path (the default config).
|
||||
caps: Final = ProxyLogging._callback_capabilities()
|
||||
cost_injection_enabled: Final = bool(getattr(litellm, "include_cost_in_streaming_usage", False))
|
||||
fast_path = not caps.has_streaming_chunk_override and not caps.has_guardrail and not cost_injection_enabled
|
||||
fast_path = (
|
||||
not caps.has_streaming_chunk_override and not caps.has_streaming_guardrail and not cost_injection_enabled
|
||||
)
|
||||
debug_enabled: Final = verbose_proxy_logger.isEnabledFor(logging.DEBUG)
|
||||
stream_completed = False
|
||||
client_disconnected = False
|
||||
|
|
|
|||
|
|
@ -78,7 +78,7 @@ from litellm.proxy.common_utils.openai_error_payload import (
|
|||
with_litellm_call_id,
|
||||
)
|
||||
from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.guardrails import GuardrailEventHooks, Mode
|
||||
from litellm.types.proxy.model_listing import ModelInfoResponse
|
||||
from litellm.types.utils import MCP_GUARDRAIL_CALL_TYPES, CallTypes, CallTypesLiteral, ModelInfo, Usage
|
||||
|
||||
|
|
@ -1128,6 +1128,11 @@ class _CallbackCapabilities:
|
|||
has_iterator_override: bool = False
|
||||
has_streaming_chunk_override: bool = False
|
||||
has_guardrail: bool = False
|
||||
# True when any CustomGuardrail is configured for a response-side hook
|
||||
# (not exclusively pre_call / pre_mcp_call). Used by streaming fast paths.
|
||||
has_streaming_guardrail: bool = False
|
||||
# True when any CustomGuardrail can run on during_call.
|
||||
has_during_call_guardrail: bool = False
|
||||
has_pre_call_override: bool = False
|
||||
has_content_enforcer: bool = False
|
||||
has_moderation_override: bool = False
|
||||
|
|
@ -1151,6 +1156,76 @@ def _overrides_moderation_hook(callback: CustomLogger) -> bool:
|
|||
return _overrides_hook(callback, "async_moderation_hook")
|
||||
|
||||
|
||||
_PRE_CALL_ONLY_HOOKS: Final = frozenset(
|
||||
{
|
||||
GuardrailEventHooks.pre_call.value,
|
||||
GuardrailEventHooks.pre_mcp_call.value,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _normalize_guardrail_hook_name(hook: object) -> str:
|
||||
if isinstance(hook, GuardrailEventHooks):
|
||||
return hook.value
|
||||
return str(hook)
|
||||
|
||||
|
||||
def _guardrail_configured_hooks(guardrail: CustomGuardrail) -> frozenset[str] | None:
|
||||
"""Return the hook names this guardrail is configured for.
|
||||
|
||||
``None`` means unrestricted (``event_hook`` unset), which matches every
|
||||
lifecycle event — including response / streaming hooks.
|
||||
"""
|
||||
event_hook = guardrail.event_hook
|
||||
if event_hook is None:
|
||||
return None
|
||||
if isinstance(event_hook, Mode):
|
||||
hooks: set[str] = set()
|
||||
for tag_value in event_hook.tags.values():
|
||||
if isinstance(tag_value, list):
|
||||
hooks.update(_normalize_guardrail_hook_name(value) for value in tag_value)
|
||||
else:
|
||||
hooks.add(_normalize_guardrail_hook_name(tag_value))
|
||||
if event_hook.default:
|
||||
default_list = event_hook.default if isinstance(event_hook.default, list) else [event_hook.default]
|
||||
hooks.update(_normalize_guardrail_hook_name(value) for value in default_list)
|
||||
return frozenset(hooks)
|
||||
if isinstance(event_hook, list):
|
||||
return frozenset(_normalize_guardrail_hook_name(hook) for hook in event_hook)
|
||||
return frozenset({_normalize_guardrail_hook_name(event_hook)})
|
||||
|
||||
|
||||
def _guardrail_affects_streaming(guardrail: CustomGuardrail) -> bool:
|
||||
"""True unless the guardrail is request-only (``pre_call`` / ``pre_mcp_call``).
|
||||
|
||||
Streaming fast-path gates must not treat a purely pre-call guardrail as a
|
||||
reason to pay the per-chunk ``async_post_call_streaming_hook`` path.
|
||||
"""
|
||||
hooks = _guardrail_configured_hooks(guardrail)
|
||||
if hooks is None:
|
||||
return True
|
||||
return bool(hooks - _PRE_CALL_ONLY_HOOKS)
|
||||
|
||||
|
||||
def _guardrail_affects_during_call(guardrail: CustomGuardrail) -> bool:
|
||||
"""True if this guardrail can run on during_call / during_mcp_call events.
|
||||
|
||||
MCP tool calls reach ``during_call_hook``, which remaps to
|
||||
``during_mcp_call``. That capability must keep the hook live so
|
||||
MCP-only guardrails are not skipped by the early return.
|
||||
"""
|
||||
hooks = _guardrail_configured_hooks(guardrail)
|
||||
if hooks is None:
|
||||
return True
|
||||
return bool(
|
||||
hooks
|
||||
& {
|
||||
GuardrailEventHooks.during_call.value,
|
||||
GuardrailEventHooks.during_mcp_call.value,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
_LISTED_MODEL_NAMES: Final = TypeAdapter(tuple[str, ...])
|
||||
|
||||
|
||||
|
|
@ -2787,6 +2862,8 @@ class ProxyLogging:
|
|||
has_iterator_override = False
|
||||
has_streaming_chunk_override = False
|
||||
has_guardrail = False
|
||||
has_streaming_guardrail = False
|
||||
has_during_call_guardrail = False
|
||||
has_pre_call_override = False
|
||||
has_content_enforcer = False
|
||||
has_moderation_override = False
|
||||
|
|
@ -2808,6 +2885,10 @@ class ProxyLogging:
|
|||
continue
|
||||
if isinstance(resolved, CustomGuardrail):
|
||||
has_guardrail = True
|
||||
if _guardrail_affects_streaming(resolved):
|
||||
has_streaming_guardrail = True
|
||||
if _guardrail_affects_during_call(resolved):
|
||||
has_during_call_guardrail = True
|
||||
elif _overrides_moderation_hook(resolved):
|
||||
has_moderation_override = True
|
||||
# Use the same leaf-class ``__dict__`` check as the other hook
|
||||
|
|
@ -2822,7 +2903,12 @@ class ProxyLogging:
|
|||
if "async_post_call_streaming_iterator_hook" in cls_attrs:
|
||||
has_iterator_override = True
|
||||
iterator_overrides.append((resolved, "override"))
|
||||
elif "apply_guardrail" in cls_attrs and not getattr(resolved, "use_native_lifecycle_hooks", False):
|
||||
elif (
|
||||
"apply_guardrail" in cls_attrs
|
||||
and not getattr(resolved, "use_native_lifecycle_hooks", False)
|
||||
and (not isinstance(resolved, CustomGuardrail) or _guardrail_affects_streaming(resolved))
|
||||
):
|
||||
# pre_call-only guardrails never scan the response stream
|
||||
iterator_overrides.append((resolved, "apply_guardrail"))
|
||||
# Walk the MRO for ``async_post_call_streaming_hook`` rather than
|
||||
# using the leaf-class ``__dict__`` check used by the other flags:
|
||||
|
|
@ -2852,6 +2938,8 @@ class ProxyLogging:
|
|||
or any(kind == "apply_guardrail" for _, kind in iterator_overrides),
|
||||
has_streaming_chunk_override=has_streaming_chunk_override,
|
||||
has_guardrail=has_guardrail,
|
||||
has_streaming_guardrail=has_streaming_guardrail,
|
||||
has_during_call_guardrail=has_during_call_guardrail,
|
||||
has_pre_call_override=has_pre_call_override,
|
||||
has_content_enforcer=has_content_enforcer,
|
||||
has_moderation_override=has_moderation_override,
|
||||
|
|
@ -2885,14 +2973,14 @@ class ProxyLogging:
|
|||
@staticmethod
|
||||
def has_streaming_callbacks() -> bool:
|
||||
caps: Final = ProxyLogging._callback_capabilities()
|
||||
return caps.has_iterator_override or caps.has_streaming_chunk_override or caps.has_guardrail
|
||||
return caps.has_iterator_override or caps.has_streaming_chunk_override or caps.has_streaming_guardrail
|
||||
|
||||
@staticmethod
|
||||
def has_streaming_chunk_hook_overrides() -> bool:
|
||||
"""True iff any callback overrides ``async_post_call_streaming_hook``
|
||||
(the per-chunk hook, distinct from the iterator wrapper)."""
|
||||
caps: Final = ProxyLogging._callback_capabilities()
|
||||
return caps.has_streaming_chunk_override or caps.has_guardrail
|
||||
return caps.has_streaming_chunk_override or caps.has_streaming_guardrail
|
||||
|
||||
def needs_iterator_wrap(self) -> bool:
|
||||
"""Whether ``async_data_generator`` needs to wrap the upstream stream
|
||||
|
|
@ -2907,11 +2995,11 @@ class ProxyLogging:
|
|||
method for the same reason as :py:meth:`needs_iterator_wrap`.
|
||||
"""
|
||||
caps: Final = ProxyLogging._callback_capabilities()
|
||||
return caps.has_streaming_chunk_override or caps.has_guardrail
|
||||
return caps.has_streaming_chunk_override or caps.has_streaming_guardrail
|
||||
|
||||
@staticmethod
|
||||
def has_during_call_guardrails() -> bool:
|
||||
return ProxyLogging._callback_capabilities().has_guardrail
|
||||
return ProxyLogging._callback_capabilities().has_during_call_guardrail
|
||||
|
||||
async def during_call_hook(
|
||||
self,
|
||||
|
|
@ -2920,7 +3008,7 @@ class ProxyLogging:
|
|||
call_type: CallTypesLiteral,
|
||||
):
|
||||
caps: Final = ProxyLogging._callback_capabilities()
|
||||
if not caps.has_guardrail and not caps.has_moderation_override:
|
||||
if not caps.has_during_call_guardrail and not caps.has_moderation_override:
|
||||
return data
|
||||
# Step 1: Collect all guardrail tasks to run in parallel
|
||||
guardrail_tasks: Final = []
|
||||
|
|
@ -3826,7 +3914,7 @@ class ProxyLogging:
|
|||
# chunk so paying it per chunk for no-op callbacks dominated stream
|
||||
# CPU time even after the iterator-chain fix.
|
||||
caps: Final = ProxyLogging._callback_capabilities()
|
||||
if not caps.has_streaming_chunk_override and not caps.has_guardrail:
|
||||
if not caps.has_streaming_chunk_override and not caps.has_streaming_guardrail:
|
||||
return response
|
||||
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
|
|
@ -3844,7 +3932,7 @@ class ProxyLogging:
|
|||
_cached_guardrail_data: dict | None = None
|
||||
_guardrail_data_computed = False
|
||||
pipeline_gated: Final = (
|
||||
stream_gated_guardrail_names(data, user_api_key_dict) if caps.has_guardrail else frozenset()
|
||||
stream_gated_guardrail_names(data, user_api_key_dict) if caps.has_streaming_guardrail else frozenset()
|
||||
)
|
||||
|
||||
for callback in litellm.callbacks:
|
||||
|
|
|
|||
|
|
@ -52,10 +52,24 @@ def test_has_streaming_callbacks_uses_custom_logger_detection(monkeypatch):
|
|||
|
||||
|
||||
def test_has_streaming_callbacks_detects_guardrails(monkeypatch):
|
||||
# Unrestricted CustomGuardrail (event_hook=None) matches every lifecycle
|
||||
# event, including response / streaming hooks.
|
||||
monkeypatch.setattr(litellm, "callbacks", [CustomGuardrail()])
|
||||
assert ProxyLogging.has_streaming_callbacks() is True
|
||||
|
||||
|
||||
def test_has_streaming_callbacks_ignores_pre_call_only_guardrails(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"callbacks",
|
||||
[CustomGuardrail(event_hook=GuardrailEventHooks.pre_call, default_on=True)],
|
||||
)
|
||||
assert ProxyLogging.has_streaming_callbacks() is False
|
||||
caps = ProxyLogging._callback_capabilities()
|
||||
assert caps.has_guardrail is True
|
||||
assert caps.has_streaming_guardrail is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_response_headers_hook_returns_early_without_callbacks(
|
||||
monkeypatch,
|
||||
|
|
|
|||
|
|
@ -271,27 +271,84 @@ def test_needs_per_chunk_streaming_hook_error_raises(proxy_logging, monkeypatch)
|
|||
|
||||
def test_has_during_call_guardrails_truth_table(monkeypatch, mock_callbacks_disabled):
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
class _G(CustomGuardrail):
|
||||
class _PreCallOnly(CustomGuardrail):
|
||||
def __init__(self):
|
||||
super().__init__(guardrail_name="g", event_hook="pre_call")
|
||||
super().__init__(guardrail_name="g-pre", event_hook=GuardrailEventHooks.pre_call)
|
||||
|
||||
class _DuringCall(CustomGuardrail):
|
||||
def __init__(self):
|
||||
super().__init__(guardrail_name="g-during", event_hook=GuardrailEventHooks.during_call)
|
||||
|
||||
class _DuringMcpCall(CustomGuardrail):
|
||||
def __init__(self):
|
||||
super().__init__(
|
||||
guardrail_name="g-during-mcp",
|
||||
event_hook=GuardrailEventHooks.during_mcp_call,
|
||||
)
|
||||
|
||||
snapshot = {
|
||||
"empty_false": ProxyLogging.has_during_call_guardrails(),
|
||||
}
|
||||
monkeypatch.setattr(litellm, "callbacks", [_G()])
|
||||
monkeypatch.setattr(litellm, "callbacks", [_PreCallOnly()])
|
||||
ProxyLogging._callback_capabilities_cache.clear()
|
||||
snapshot["with_guardrail_true"] = ProxyLogging.has_during_call_guardrails()
|
||||
snapshot["pre_call_only_false"] = ProxyLogging.has_during_call_guardrails()
|
||||
monkeypatch.setattr(litellm, "callbacks", [_DuringCall()])
|
||||
ProxyLogging._callback_capabilities_cache.clear()
|
||||
snapshot["with_during_call_true"] = ProxyLogging.has_during_call_guardrails()
|
||||
monkeypatch.setattr(litellm, "callbacks", [_DuringMcpCall()])
|
||||
ProxyLogging._callback_capabilities_cache.clear()
|
||||
snapshot["with_during_mcp_call_true"] = ProxyLogging.has_during_call_guardrails()
|
||||
monkeypatch.setattr(litellm, "callbacks", [_PlainLogger()])
|
||||
ProxyLogging._callback_capabilities_cache.clear()
|
||||
snapshot["only_plain_logger_false"] = ProxyLogging.has_during_call_guardrails()
|
||||
assert snapshot == {
|
||||
"empty_false": False,
|
||||
"with_guardrail_true": True,
|
||||
"pre_call_only_false": False,
|
||||
"with_during_call_true": True,
|
||||
"with_during_mcp_call_true": True,
|
||||
"only_plain_logger_false": False,
|
||||
}
|
||||
|
||||
|
||||
def test_streaming_fast_path_ignores_pre_call_only_guardrails(monkeypatch, mock_callbacks_disabled):
|
||||
"""pre_call-only CustomGuardrails must not disable the streaming fast path."""
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
class _PreCallOnly(CustomGuardrail):
|
||||
def __init__(self):
|
||||
super().__init__(
|
||||
guardrail_name="pre-only",
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
default_on=True,
|
||||
)
|
||||
|
||||
class _PostCall(CustomGuardrail):
|
||||
def __init__(self):
|
||||
super().__init__(
|
||||
guardrail_name="post",
|
||||
event_hook=GuardrailEventHooks.post_call,
|
||||
default_on=True,
|
||||
)
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", [_PreCallOnly()])
|
||||
ProxyLogging._callback_capabilities_cache.clear()
|
||||
caps = ProxyLogging._callback_capabilities()
|
||||
assert caps.has_guardrail is True
|
||||
assert caps.has_streaming_guardrail is False
|
||||
assert ProxyLogging.has_streaming_callbacks() is False
|
||||
assert ProxyLogging.has_streaming_chunk_hook_overrides() is False
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", [_PreCallOnly(), _PostCall()])
|
||||
ProxyLogging._callback_capabilities_cache.clear()
|
||||
caps = ProxyLogging._callback_capabilities()
|
||||
assert caps.has_guardrail is True
|
||||
assert caps.has_streaming_guardrail is True
|
||||
assert ProxyLogging.has_streaming_callbacks() is True
|
||||
|
||||
|
||||
def test_has_during_call_guardrails_resolution_error_raises(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "callbacks", ["x"])
|
||||
monkeypatch.setattr(
|
||||
|
|
@ -319,9 +376,7 @@ def test_get_combined_callback_list_matrix(proxy_logging):
|
|||
"none_dynamic_returns_global_copy": proxy_logging.get_combined_callback_list(
|
||||
dynamic_success_callbacks=None, global_callbacks=["a", "b", "c"]
|
||||
),
|
||||
"empty_both": proxy_logging.get_combined_callback_list(
|
||||
dynamic_success_callbacks=[], global_callbacks=[]
|
||||
),
|
||||
"empty_both": proxy_logging.get_combined_callback_list(dynamic_success_callbacks=[], global_callbacks=[]),
|
||||
}
|
||||
assert snapshot == {
|
||||
"merge_dedupes_shared": ["dyn-1", "glob-1", "shared"],
|
||||
|
|
|
|||
|
|
@ -16,6 +16,8 @@ def test_callback_capabilities_default_values():
|
|||
"has_iterator_override": caps.has_iterator_override,
|
||||
"has_streaming_chunk_override": caps.has_streaming_chunk_override,
|
||||
"has_guardrail": caps.has_guardrail,
|
||||
"has_streaming_guardrail": caps.has_streaming_guardrail,
|
||||
"has_during_call_guardrail": caps.has_during_call_guardrail,
|
||||
"has_pre_call_override": caps.has_pre_call_override,
|
||||
"iterator_overrides": caps.iterator_overrides,
|
||||
"resolved_callbacks": caps.resolved_callbacks,
|
||||
|
|
@ -25,6 +27,8 @@ def test_callback_capabilities_default_values():
|
|||
"has_iterator_override": False,
|
||||
"has_streaming_chunk_override": False,
|
||||
"has_guardrail": False,
|
||||
"has_streaming_guardrail": False,
|
||||
"has_during_call_guardrail": False,
|
||||
"has_pre_call_override": False,
|
||||
"iterator_overrides": (),
|
||||
"resolved_callbacks": (),
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue