mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
4f7d86b1b3
commit
70e5db056c
3 changed files with 30 additions and 14 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue