mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(guardrails): read pipeline state from the bucket the policy engine wrote
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
34b513074e
commit
012e4f98e4
2 changed files with 40 additions and 14 deletions
|
|
@ -88,10 +88,7 @@ from litellm.integrations.custom_logger import CustomLogger
|
|||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
|
||||
from litellm.integrations.SlackAlerting.utils import _add_langfuse_trace_id_to_alert
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
coerce_token_limit,
|
||||
get_metadata_variable_name_from_kwargs,
|
||||
)
|
||||
from litellm.litellm_core_utils.core_helpers import coerce_token_limit
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
|
|
@ -379,6 +376,28 @@ def _exception_changes_request_flow(exc: BaseException) -> bool:
|
|||
return isinstance(exc, (SensitiveDataRouteException, ModifyResponseException))
|
||||
|
||||
|
||||
def _policy_state_metadata(data: Mapping[str, object]) -> Mapping[str, object]:
|
||||
"""
|
||||
Return the metadata bucket the policy engine wrote its pipeline state into.
|
||||
|
||||
The route decides the bucket (``litellm_metadata`` for ``/v1/messages``,
|
||||
responses, batches, files and bedrock, ``metadata`` everywhere else), and both
|
||||
buckets can be present at once because callers send their own provider-facing
|
||||
``metadata`` (Claude Code sends ``metadata.user_id``) or their own
|
||||
``litellm_metadata``. Pipeline slots are stripped from caller input before the
|
||||
policy engine runs, so whichever bucket carries them is the proxy's own write.
|
||||
"""
|
||||
return next(
|
||||
(
|
||||
bucket
|
||||
for bucket in (data.get("metadata"), data.get("litellm_metadata"))
|
||||
if isinstance(bucket, dict)
|
||||
and ("_guardrail_pipelines" in bucket or "_pipeline_managed_guardrails" in bucket)
|
||||
),
|
||||
{},
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _CallbackCapabilities:
|
||||
"""Cached per-hook capability flags derived from ``litellm.callbacks``.
|
||||
|
|
@ -1251,7 +1270,7 @@ class ProxyLogging:
|
|||
|
||||
Returns the (possibly modified) data dict.
|
||||
"""
|
||||
metadata: Final = data.get(get_metadata_variable_name_from_kwargs(data)) or {}
|
||||
metadata: Final = _policy_state_metadata(data)
|
||||
pipelines: Final = metadata.get("_guardrail_pipelines")
|
||||
if not pipelines:
|
||||
return data
|
||||
|
|
@ -1403,7 +1422,7 @@ class ProxyLogging:
|
|||
)
|
||||
|
||||
# Get pipeline-managed guardrails to skip in normal loop
|
||||
metadata: Final = data.get(get_metadata_variable_name_from_kwargs(data)) or {}
|
||||
metadata: Final = _policy_state_metadata(data)
|
||||
pipeline_managed: Final[set] = metadata.get("_pipeline_managed_guardrails", set())
|
||||
|
||||
caps: Final = ProxyLogging._callback_capabilities()
|
||||
|
|
|
|||
|
|
@ -354,13 +354,20 @@ async def test_maybe_execute_pipelines_skips_pipelines_with_other_mode(proxy_log
|
|||
assert out is data
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("policy_state_key", "caller_metadata_key", "call_type"),
|
||||
[
|
||||
("litellm_metadata", "metadata", "anthropic_messages"),
|
||||
("metadata", "litellm_metadata", "completion"),
|
||||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_maybe_execute_pipelines_reads_litellm_metadata_when_caller_sends_own_metadata(
|
||||
proxy_logging, make_user_api_key_auth, monkeypatch
|
||||
async def test_maybe_execute_pipelines_finds_policy_state_when_caller_sends_own_metadata(
|
||||
proxy_logging, make_user_api_key_auth, monkeypatch, policy_state_key, caller_metadata_key, call_type
|
||||
):
|
||||
"""On /v1/messages the proxy stores policy state in ``litellm_metadata``, while the
|
||||
caller's provider-facing ``metadata`` (Claude Code sends ``metadata.user_id``) stays
|
||||
untouched. The pipeline must still run and block."""
|
||||
"""The route picks the bucket the policy engine writes to (``litellm_metadata`` on
|
||||
/v1/messages, ``metadata`` on chat completions), and the caller can populate the other
|
||||
one, e.g. Claude Code sending ``metadata.user_id``. The pipeline must still run and block."""
|
||||
|
||||
class BlockingGuardrail(CustomGuardrail):
|
||||
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
|
||||
|
|
@ -369,8 +376,8 @@ async def test_maybe_execute_pipelines_reads_litellm_metadata_when_caller_sends_
|
|||
monkeypatch.setattr(litellm, "callbacks", [BlockingGuardrail(guardrail_name="gr-1")])
|
||||
pipeline = GuardrailPipeline(mode="pre_call", steps=[PipelineStep(guardrail="gr-1", on_fail="block")])
|
||||
data = {
|
||||
"metadata": {"user_id": "user_abc"},
|
||||
"litellm_metadata": {"_guardrail_pipelines": [("policy-1", pipeline)]},
|
||||
caller_metadata_key: {"user_id": "user_abc"},
|
||||
policy_state_key: {"_guardrail_pipelines": [("policy-1", pipeline)]},
|
||||
"messages": [],
|
||||
"model": "m",
|
||||
}
|
||||
|
|
@ -379,7 +386,7 @@ async def test_maybe_execute_pipelines_reads_litellm_metadata_when_caller_sends_
|
|||
await proxy_logging._maybe_execute_pipelines(
|
||||
data=data,
|
||||
user_api_key_dict=make_user_api_key_auth(),
|
||||
call_type="anthropic_messages",
|
||||
call_type=call_type,
|
||||
event_hook="pre_call",
|
||||
)
|
||||
assert exc_info.value.detail["error"] == "blocked by pipeline"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue