feat(proxy): enforce tag budgets for tags added by guardrails

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-09-12 08:04:24 +00:00 committed by yassin
parent 252c71c0b2
commit e79cf49722
2 changed files with 149 additions and 2 deletions

View file

@ -33,7 +33,11 @@ from litellm.constants import (
UNSAFE_PROXY_RESPONSE_HEADERS,
)
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket, is_expected_client_error
from litellm.litellm_core_utils.core_helpers import (
get_metadata_variable_name_from_kwargs,
get_or_create_metadata_bucket,
is_expected_client_error,
)
from litellm.litellm_core_utils.dd_tracing import NullTracer, tracer
from litellm.litellm_core_utils.get_supported_openai_params import (
get_supported_openai_params,
@ -49,12 +53,16 @@ from litellm.litellm_core_utils.streaming_handler import (
backfill_missing_cache_usage_fields,
)
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
from litellm.proxy.auth.auth_checks import can_key_call_resolved_model
from litellm.proxy.auth.auth_checks import (
_tag_max_budget_check, # pyright: ignore[reportPrivateUsage] # auth's tag budget gate, reused for tags guardrails add post-auth
can_key_call_resolved_model,
)
from litellm.proxy.auth.auth_utils import check_response_size_is_safe
from litellm.proxy.common_utils.callback_utils import (
get_logging_caching_headers,
get_remaining_tokens_and_requests_from_request_data,
)
from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body
from litellm.proxy.common_utils.openai_error_payload import (
attribute_of,
error_status_code,
@ -640,6 +648,30 @@ async def _resolve_per_request_model_group_alias(
return target
async def _enforce_guardrail_added_tag_budgets(
data: Mapping[str, object],
tags_before_guardrails: frozenset[str],
user_api_key_dict: UserAPIKeyAuth,
proxy_logging_obj: ProxyLogging,
) -> None:
added_tags: Final = tuple(
tag for tag in get_tags_from_request_body(request_body=data) if tag not in tags_before_guardrails
)
if not added_tags:
return
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
metadata_key: Final = get_metadata_variable_name_from_kwargs(data)
request_body: Final = {metadata_key: {"tags": list(added_tags)}} # mutable-ok: budget check takes a dict
await _tag_max_budget_check(
request_body=request_body,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
valid_token=user_api_key_dict,
)
async def _parse_event_data_for_error(event_line: str | bytes) -> int | None:
"""Parses an event line and returns an error code if present, else None."""
event_line = event_line.decode("utf-8") if isinstance(event_line, bytes) else event_line
@ -1999,11 +2031,18 @@ class ProxyBaseLLMRequestProcessing:
# to run below.
await _arm_auto_router_compression(data=self.data, llm_router=llm_router)
tags_before_guardrails: Final = frozenset(get_tags_from_request_body(request_body=self.data))
self.data = await proxy_logging_obj.pre_call_hook(
user_api_key_dict=user_api_key_dict,
data=self.data,
call_type=route_type,
)
await _enforce_guardrail_added_tag_budgets(
data=self.data,
tags_before_guardrails=tags_before_guardrails,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
)
if route_type == "aget_responses":
attach_post_call_pipelines_to_retrieval(
data=self.data,

View file

@ -377,6 +377,114 @@ class TestProxyBaseLLMRequestProcessing:
assert "litellm_logging_obj" not in persisted_body
json.dumps(persisted_body)
@pytest.mark.asyncio
async def test_common_processing_pre_call_logic_enforces_tag_budget_for_guardrail_added_tags(
self, monkeypatch
):
"""
Per-tag budgets are enforced at auth time against tags already in the request
body. A custom guardrail may add tags to data["metadata"]["tags"] inside
pre_call_hook, which runs after auth; those added tags must still be
budget-checked before the model call proceeds.
"""
processing_obj = ProxyBaseLLMRequestProcessing(data={})
mock_request = MagicMock(spec=Request)
mock_request.headers = {}
async def mock_add_litellm_data_to_request(*args, **kwargs):
return {
"messages": [{"role": "user", "content": "hello"}],
"metadata": {"tags": ["existing-tag"]},
}
async def mock_pre_call_hook(user_api_key_dict, data, call_type):
data["metadata"]["tags"].append("guardrail-tag")
return data
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
mock_proxy_logging_obj.pre_call_hook = AsyncMock(side_effect=mock_pre_call_hook)
monkeypatch.setattr(
litellm.proxy.common_request_processing,
"add_litellm_data_to_request",
mock_add_litellm_data_to_request,
)
tag_budget_check = AsyncMock(
side_effect=litellm.BudgetExceededError(current_cost=10.0, max_budget=5.0)
)
monkeypatch.setattr(
"litellm.proxy.common_request_processing._tag_max_budget_check",
tag_budget_check,
)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MagicMock())
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock())
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
with pytest.raises(litellm.BudgetExceededError):
await processing_obj.common_processing_pre_call_logic(
request=mock_request,
general_settings={},
user_api_key_dict=mock_user_api_key_dict,
proxy_logging_obj=mock_proxy_logging_obj,
proxy_config=MagicMock(spec=ProxyConfig),
route_type="acompletion",
)
tag_budget_check.assert_awaited_once()
_, call_kwargs = tag_budget_check.call_args
assert call_kwargs["request_body"] == {"metadata": {"tags": ["guardrail-tag"]}}
assert call_kwargs["valid_token"] is mock_user_api_key_dict
@pytest.mark.asyncio
async def test_common_processing_pre_call_logic_skips_tag_budget_check_when_guardrails_add_no_tags(
self, monkeypatch
):
"""
When pre_call_hook does not add tags, the post-hook tag budget check must
not run at all, so requests without guardrail-added tags pay no extra lookups.
"""
processing_obj = ProxyBaseLLMRequestProcessing(data={})
mock_request = MagicMock(spec=Request)
mock_request.headers = {}
async def mock_add_litellm_data_to_request(*args, **kwargs):
return {
"messages": [{"role": "user", "content": "hello"}],
"metadata": {"tags": ["existing-tag"]},
}
async def mock_pre_call_hook(user_api_key_dict, data, call_type):
return data
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
mock_proxy_logging_obj.pre_call_hook = AsyncMock(side_effect=mock_pre_call_hook)
monkeypatch.setattr(
litellm.proxy.common_request_processing,
"add_litellm_data_to_request",
mock_add_litellm_data_to_request,
)
tag_budget_check = AsyncMock()
monkeypatch.setattr(
"litellm.proxy.common_request_processing._tag_max_budget_check",
tag_budget_check,
)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MagicMock())
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock())
returned_data, _ = await processing_obj.common_processing_pre_call_logic(
request=mock_request,
general_settings={},
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth),
proxy_logging_obj=mock_proxy_logging_obj,
proxy_config=MagicMock(spec=ProxyConfig),
route_type="acompletion",
)
assert returned_data["metadata"]["tags"] == ["existing-tag"]
tag_budget_check.assert_not_awaited()
@pytest.mark.asyncio
async def test_common_processing_pre_call_logic_arms_auto_router_compression_before_guardrails(
self, monkeypatch