diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 97de2488b8d..59c44ab8101 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -3894,7 +3894,7 @@ 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 diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index ea294b76e92..2f2e59f6ee4 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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,68 @@ 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 the ``during_call`` event.""" + hooks = _guardrail_configured_hooks(guardrail) + if hooks is None: + return True + return GuardrailEventHooks.during_call.value in hooks + + _LISTED_MODEL_NAMES: Final = TypeAdapter(tuple[str, ...]) @@ -2787,6 +2854,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 +2877,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 +2895,15 @@ 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 +2933,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 +2968,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 +2990,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 +3003,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 +3909,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 +3927,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: diff --git a/tests/test_litellm/proxy/test_proxy_logging_hook_detection.py b/tests/test_litellm/proxy/test_proxy_logging_hook_detection.py index 9eb89b2f301..a14d2c920f6 100644 --- a/tests/test_litellm/proxy/test_proxy_logging_hook_detection.py +++ b/tests/test_litellm/proxy/test_proxy_logging_hook_detection.py @@ -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, diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_callback_capabilities_class.py b/tests/test_litellm/proxy/utils/proxy_logging/test_callback_capabilities_class.py index 75a91177f00..97a027dbb86 100644 --- a/tests/test_litellm/proxy/utils/proxy_logging/test_callback_capabilities_class.py +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_callback_capabilities_class.py @@ -271,27 +271,73 @@ 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) 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", [_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, "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( diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_callback_capabilities_dataclass.py b/tests/test_litellm/proxy/utils/proxy_logging/test_callback_capabilities_dataclass.py index 931c832732e..68e47337fb9 100644 --- a/tests/test_litellm/proxy/utils/proxy_logging/test_callback_capabilities_dataclass.py +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_callback_capabilities_dataclass.py @@ -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": (),