fix(proxy): count during_mcp_call in during-call capability gate

MCP tool calls reach during_call_hook, which remaps to during_mcp_call.
Include that event in _guardrail_affects_during_call so MCP-only
guardrails are not skipped by the early return. Also ruff-format the
touched proxy files.

Signed-off-by: GokayAI <60583610+gokay-ai@users.noreply.github.com>
This commit is contained in:
GokayAI 2026-09-30 09:22:44 +03:00
parent 4f7d86b1b3
commit 70e5db056c
No known key found for this signature in database
3 changed files with 30 additions and 14 deletions

View file

@ -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

View file

@ -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"))

View file

@ -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"],