diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 59c44ab8101..0a71e8dfd40 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -3894,7 +3894,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_streaming_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 2f2e59f6ee4..c31f304237f 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -1156,7 +1156,6 @@ def _overrides_moderation_hook(callback: CustomLogger) -> bool: return _overrides_hook(callback, "async_moderation_hook") - _PRE_CALL_ONLY_HOOKS: Final = frozenset( { GuardrailEventHooks.pre_call.value, @@ -1188,9 +1187,7 @@ def _guardrail_configured_hooks(guardrail: CustomGuardrail) -> frozenset[str] | 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] - ) + 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): @@ -1211,11 +1208,22 @@ def _guardrail_affects_streaming(guardrail: CustomGuardrail) -> bool: def _guardrail_affects_during_call(guardrail: CustomGuardrail) -> bool: - """True if this guardrail can run on the ``during_call`` event.""" + """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 GuardrailEventHooks.during_call.value in hooks + return bool( + hooks + & { + GuardrailEventHooks.during_call.value, + GuardrailEventHooks.during_mcp_call.value, + } + ) _LISTED_MODEL_NAMES: Final = TypeAdapter(tuple[str, ...]) @@ -2898,10 +2906,7 @@ class ProxyLogging: 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) - ) + 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")) 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 97a027dbb86..36751315cc2 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 @@ -281,6 +281,13 @@ def test_has_during_call_guardrails_truth_table(monkeypatch, mock_callbacks_disa 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(), } @@ -290,6 +297,9 @@ def test_has_during_call_guardrails_truth_table(monkeypatch, mock_callbacks_disa 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() @@ -297,6 +307,7 @@ def test_has_during_call_guardrails_truth_table(monkeypatch, mock_callbacks_disa "empty_false": False, "pre_call_only_false": False, "with_during_call_true": True, + "with_during_mcp_call_true": True, "only_plain_logger_false": False, } @@ -365,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"],