This commit is contained in:
GokayAI 2026-10-01 02:24:12 +08:00 • committed by GitHub
commit 6790292945
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 181 additions and 18 deletions

View file

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

View file

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

View file

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

View file

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

View file

@ -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": (),