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

Auth runs the tag budget check before pre_call_hook, so a tag that a custom guardrail adds is attributed spend but never budget checked. After the pre-call hook, budget check only the newly added tags with the same exemptions auth applied (budget-free routes, zero-cost models), keep the pre-guardrail tag baseline across fallback retries, and surface an over-budget tag as the same budget_exceeded 429 auth returns

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yassin 2026-09-16 23:23:00 +00:00 committed by Devin AI
parent 319f427c40
commit 488666ccae
3 changed files with 387 additions and 19 deletions

View file

@ -856,6 +856,16 @@ BUDGET_ENFORCED_SIDE_EFFECT_ROUTES: Final = frozenset(
)
def route_skips_budget_checks(route: str) -> bool:
return route not in BUDGET_ENFORCED_SIDE_EFFECT_ROUTES and (
route in MODEL_DISCOVERY_ROUTES or not RouteChecks.is_llm_api_route(route=route)
)
def request_skips_budget_checks(route: str, model: str | list[str] | None, llm_router: Router | None) -> bool:
return route_skips_budget_checks(route=route) or _is_model_cost_zero(model=model, llm_router=llm_router)
async def common_checks(
request_body: dict,
team_object: LiteLLM_TeamTable | None,
@ -903,10 +913,7 @@ async def common_checks(
team_id=valid_token.team_id if valid_token is not None else None,
)
skip_all_budget_checks: Final = skip_budget_checks or (
route not in BUDGET_ENFORCED_SIDE_EFFECT_ROUTES
and (route in MODEL_DISCOVERY_ROUTES or not RouteChecks.is_llm_api_route(route=route))
)
skip_all_budget_checks: Final = skip_budget_checks or route_skips_budget_checks(route=route)
membership_user_id: Final = (
valid_token.user_id if valid_token is not None and (bool(_model) or not skip_all_budget_checks) else None
@ -2104,7 +2111,7 @@ async def _fetch_uncached_tags(
@log_db_metrics
async def get_tag_objects_batch(
tag_names: list[str],
tag_names: Sequence[str],
prisma_client: PrismaClient | None,
user_api_key_cache: UserApiKeyCache,
parent_otel_span: Span | None = None,
@ -5863,15 +5870,25 @@ async def _tag_max_budget_check(
"""
from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body
if prisma_client is None:
await tag_max_budget_check_for_tags(
tags=get_tags_from_request_body(request_body=request_body),
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
valid_token=valid_token,
)
async def tag_max_budget_check_for_tags(
tags: Sequence[str],
prisma_client: PrismaClient | None,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
valid_token: UserAPIKeyAuth | None,
) -> None:
if prisma_client is None or not tags:
return
# Get tags from request metadata
tags: Final = get_tags_from_request_body(request_body=request_body)
if not tags:
return
# Batch fetch all tags in one go
tag_objects: Final = await get_tag_objects_batch(
tag_names=tags,
prisma_client=prisma_client,

View file

@ -25,7 +25,7 @@ import httpx
import orjson
from fastapi import HTTPException, Request, status
from fastapi.responses import JSONResponse, Response, StreamingResponse
from pydantic import ValidationError
from pydantic import TypeAdapter, ValidationError
from starlette.types import Receive, Scope, Send
import litellm
@ -64,14 +64,21 @@ from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
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_utils import check_response_size_is_safe
from litellm.proxy._types import ProxyErrorTypes, ProxyException, UserAPIKeyAuth
from litellm.proxy.auth.auth_checks import (
can_key_call_resolved_model,
request_skips_budget_checks,
tag_max_budget_check_for_tags,
)
from litellm.proxy.auth.auth_utils import check_response_size_is_safe, get_request_route
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_client_requested_model
from litellm.proxy.common_utils.http_parsing_utils import (
get_client_requested_model,
get_tags_from_request_body,
)
from litellm.proxy.common_utils.openai_error_payload import (
attribute_of,
error_status_code,
@ -658,6 +665,48 @@ async def _resolve_per_request_model_group_alias(
return target
_REQUEST_MODEL: Final[TypeAdapter[str | list[str] | None]] = TypeAdapter(str | list[str] | None)
def _request_model(data: Mapping[str, object]) -> str | list[str] | None:
try:
return _REQUEST_MODEL.validate_python(data.get("model"), strict=True)
except ValidationError:
return None
async def _enforce_guardrail_added_tag_budgets(
data: Mapping[str, object],
tags_before_guardrails: frozenset[str],
route: str,
llm_router: Router | None,
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 or request_skips_budget_checks(route=route, model=_request_model(data), llm_router=llm_router):
return
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
try:
await tag_max_budget_check_for_tags(
tags=added_tags,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
valid_token=user_api_key_dict,
)
except litellm.BudgetExceededError as e:
raise ProxyException(
message=e.message,
type=ProxyErrorTypes.budget_exceeded,
param=None,
code=e.status_code,
) from e
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
@ -1558,6 +1607,7 @@ def _timing_values(
class ProxyBaseLLMRequestProcessing:
def __init__(self, data: dict):
self.data = data
self._tags_before_guardrails: frozenset[str] | None = None
@property
def litellm_call_id(self) -> str | None:
@ -2051,11 +2101,21 @@ class ProxyBaseLLMRequestProcessing:
# to run below.
await _arm_auto_router_compression(data=self.data, llm_router=llm_router)
if self._tags_before_guardrails is None:
self._tags_before_guardrails = 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=self._tags_before_guardrails,
route=get_request_route(request=request),
llm_router=llm_router,
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

@ -3,7 +3,7 @@ import copy
import datetime
import json
from types import MappingProxyType, SimpleNamespace
from typing import AsyncGenerator, Callable, Final, Iterator, Optional
from typing import AsyncGenerator, Callable, Final, Iterator, Optional, Sequence
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
@ -45,9 +45,10 @@ from litellm.proxy.common_request_processing import (
)
from litellm.proxy.dd_span_tagger import DDSpanTagger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._types import ProxyException
from litellm.proxy._types import ProxyErrorTypes, ProxyException
from litellm.proxy._types import UserAPIKeyAuth as ProxyUserAPIKeyAuth
from litellm.proxy.utils import ProxyLogging
from litellm.router import Router
class TestProxyBaseLLMRequestProcessing:
@ -382,6 +383,296 @@ class TestProxyBaseLLMRequestProcessing:
assert "litellm_logging_obj" not in persisted_body
json.dumps(persisted_body)
@staticmethod
def _guardrail_tag_budget_harness(
monkeypatch,
request_body: dict,
guardrail_tags: Sequence[str],
route: str = "/v1/chat/completions",
) -> tuple[ProxyBaseLLMRequestProcessing, MagicMock, MagicMock, MagicMock, AsyncMock]:
processing_obj = ProxyBaseLLMRequestProcessing(data={})
mock_request = MagicMock(spec=Request)
mock_request.headers = {}
mock_request.scope = {"path": route}
async def mock_add_litellm_data_to_request(*args, **kwargs):
return copy.deepcopy(request_body)
async def mock_pre_call_hook(user_api_key_dict, data, call_type):
data.setdefault("metadata", {}).setdefault("tags", []).extend(guardrail_tags)
return data
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
mock_proxy_logging_obj.pre_call_hook = AsyncMock(side_effect=mock_pre_call_hook)
mock_proxy_config = MagicMock(spec=ProxyConfig)
mock_proxy_config._get_hierarchical_router_settings = AsyncMock(return_value=None)
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_for_tags", tag_budget_check)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MagicMock())
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock())
return processing_obj, mock_request, mock_proxy_logging_obj, mock_proxy_config, tag_budget_check
@staticmethod
def _router_with_free_and_paid_models() -> Router:
return Router(
model_list=[
{
"model_name": "free-model",
"litellm_params": {"model": "openai/gpt-4.1-mini", "api_key": "sk-test"},
"model_info": {"input_cost_per_token": 0, "output_cost_per_token": 0},
},
{
"model_name": "paid-model",
"litellm_params": {"model": "openai/gpt-4.1-mini", "api_key": "sk-test"},
},
]
)
@pytest.mark.asyncio
async def test_common_processing_pre_call_logic_enforces_tag_budget_for_guardrail_added_tags(
self, monkeypatch
):
processing_obj, mock_request, mock_proxy_logging_obj, mock_proxy_config, tag_budget_check = (
self._guardrail_tag_budget_harness(
monkeypatch,
request_body={
"model": "paid-model",
"messages": [{"role": "user", "content": "hello"}],
"metadata": {"tags": ["existing-tag"]},
},
guardrail_tags=["guardrail-tag"],
)
)
tag_budget_check.side_effect = litellm.BudgetExceededError(current_cost=10.0, max_budget=5.0)
user_api_key_dict = ProxyUserAPIKeyAuth(api_key="sk-test")
with pytest.raises(ProxyException) as exc_info:
await processing_obj.common_processing_pre_call_logic(
request=mock_request,
general_settings={},
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=mock_proxy_logging_obj,
proxy_config=mock_proxy_config,
route_type="acompletion",
llm_router=self._router_with_free_and_paid_models(),
)
assert exc_info.value.type == ProxyErrorTypes.budget_exceeded
assert exc_info.value.code == "429"
tag_budget_check.assert_awaited_once()
_, call_kwargs = tag_budget_check.call_args
assert call_kwargs["tags"] == ("guardrail-tag",)
assert call_kwargs["valid_token"] is 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
):
processing_obj, mock_request, mock_proxy_logging_obj, mock_proxy_config, tag_budget_check = (
self._guardrail_tag_budget_harness(
monkeypatch,
request_body={
"model": "paid-model",
"messages": [{"role": "user", "content": "hello"}],
"metadata": {"tags": ["existing-tag"]},
},
guardrail_tags=[],
)
)
returned_data, _ = await processing_obj.common_processing_pre_call_logic(
request=mock_request,
general_settings={},
user_api_key_dict=ProxyUserAPIKeyAuth(api_key="sk-test"),
proxy_logging_obj=mock_proxy_logging_obj,
proxy_config=mock_proxy_config,
route_type="acompletion",
llm_router=self._router_with_free_and_paid_models(),
)
assert returned_data["metadata"]["tags"] == ["existing-tag"]
tag_budget_check.assert_not_awaited()
@pytest.mark.asyncio
async def test_common_processing_pre_call_logic_skips_guardrail_tag_budget_check_for_zero_cost_model(
self, monkeypatch
):
processing_obj, mock_request, mock_proxy_logging_obj, mock_proxy_config, tag_budget_check = (
self._guardrail_tag_budget_harness(
monkeypatch,
request_body={"model": "free-model", "messages": [{"role": "user", "content": "hello"}]},
guardrail_tags=["guardrail-tag"],
)
)
returned_data, _ = await processing_obj.common_processing_pre_call_logic(
request=mock_request,
general_settings={},
user_api_key_dict=ProxyUserAPIKeyAuth(api_key="sk-test"),
proxy_logging_obj=mock_proxy_logging_obj,
proxy_config=mock_proxy_config,
route_type="acompletion",
llm_router=self._router_with_free_and_paid_models(),
)
assert returned_data["metadata"]["tags"] == ["guardrail-tag"]
tag_budget_check.assert_not_awaited()
@pytest.mark.asyncio
async def test_common_processing_pre_call_logic_skips_guardrail_tag_budget_check_on_budget_exempt_route(
self, monkeypatch
):
processing_obj, mock_request, mock_proxy_logging_obj, mock_proxy_config, tag_budget_check = (
self._guardrail_tag_budget_harness(
monkeypatch,
request_body={"model": "paid-model", "text": "hello"},
guardrail_tags=["guardrail-tag"],
route="/guardrails/apply_guardrail",
)
)
await processing_obj.common_processing_pre_call_logic(
request=mock_request,
general_settings={},
user_api_key_dict=ProxyUserAPIKeyAuth(api_key="sk-test"),
proxy_logging_obj=mock_proxy_logging_obj,
proxy_config=mock_proxy_config,
route_type="acompletion",
llm_router=self._router_with_free_and_paid_models(),
)
tag_budget_check.assert_not_awaited()
@pytest.mark.asyncio
async def test_common_processing_pre_call_logic_rechecks_guardrail_added_tag_on_fallback_retry(
self, monkeypatch
):
processing_obj, mock_request, mock_proxy_logging_obj, mock_proxy_config, tag_budget_check = (
self._guardrail_tag_budget_harness(
monkeypatch,
request_body={
"model": "paid-model",
"messages": [{"role": "user", "content": "hello"}],
"metadata": {"tags": ["existing-tag"]},
},
guardrail_tags=["guardrail-tag"],
)
)
first_pass_data, _ = await processing_obj.common_processing_pre_call_logic(
request=mock_request,
general_settings={},
user_api_key_dict=ProxyUserAPIKeyAuth(api_key="sk-test"),
proxy_logging_obj=mock_proxy_logging_obj,
proxy_config=mock_proxy_config,
route_type="acompletion",
llm_router=self._router_with_free_and_paid_models(),
)
assert first_pass_data["metadata"]["tags"] == ["existing-tag", "guardrail-tag"]
tag_budget_check.reset_mock()
async def retry_add_litellm_data_to_request(*args, **kwargs):
return first_pass_data
async def idempotent_pre_call_hook(user_api_key_dict, data, call_type):
return data
monkeypatch.setattr(
litellm.proxy.common_request_processing,
"add_litellm_data_to_request",
retry_add_litellm_data_to_request,
)
mock_proxy_logging_obj.pre_call_hook = AsyncMock(side_effect=idempotent_pre_call_hook)
await processing_obj.common_processing_pre_call_logic(
request=mock_request,
general_settings={},
user_api_key_dict=ProxyUserAPIKeyAuth(api_key="sk-test"),
proxy_logging_obj=mock_proxy_logging_obj,
proxy_config=mock_proxy_config,
route_type="acompletion",
llm_router=self._router_with_free_and_paid_models(),
)
tag_budget_check.assert_awaited_once()
_, call_kwargs = tag_budget_check.call_args
assert call_kwargs["tags"] == ("guardrail-tag",)
@pytest.mark.asyncio
async def test_enforce_guardrail_added_tag_budgets_checks_only_added_tags(self, monkeypatch):
from litellm.proxy.common_request_processing import _enforce_guardrail_added_tag_budgets
tag_budget_check = AsyncMock()
monkeypatch.setattr(litellm.proxy.common_request_processing, "tag_max_budget_check_for_tags", tag_budget_check)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MagicMock())
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock())
user_api_key_dict = ProxyUserAPIKeyAuth(api_key="sk-test")
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
await _enforce_guardrail_added_tag_budgets(
data={"metadata": {"tags": ["existing-tag", "guardrail-tag"]}},
tags_before_guardrails=frozenset({"existing-tag"}),
route="/v1/chat/completions",
llm_router=None,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=mock_proxy_logging_obj,
)
tag_budget_check.assert_awaited_once()
_, call_kwargs = tag_budget_check.call_args
assert call_kwargs["tags"] == ("guardrail-tag",)
assert call_kwargs["valid_token"] is user_api_key_dict
tag_budget_check.reset_mock()
await _enforce_guardrail_added_tag_budgets(
data={"metadata": {"tags": ["existing-tag"]}},
tags_before_guardrails=frozenset({"existing-tag"}),
route="/v1/chat/completions",
llm_router=None,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=mock_proxy_logging_obj,
)
tag_budget_check.assert_not_awaited()
@pytest.mark.asyncio
async def test_enforce_guardrail_added_tag_budgets_raises_budget_exceeded_proxy_exception(self, monkeypatch):
from litellm.proxy.common_request_processing import _enforce_guardrail_added_tag_budgets
monkeypatch.setattr(
litellm.proxy.common_request_processing,
"tag_max_budget_check_for_tags",
AsyncMock(
side_effect=litellm.BudgetExceededError(
current_cost=2.0,
max_budget=1.0,
message="Budget has been exceeded! Tag=guardrail-tag Current cost: 2.0, Max budget: 1.0",
entity_type="tag",
entity_id="guardrail-tag",
)
),
)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MagicMock())
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock())
with pytest.raises(ProxyException) as exc_info:
await _enforce_guardrail_added_tag_budgets(
data={"metadata": {"tags": ["guardrail-tag"]}},
tags_before_guardrails=frozenset(),
route="/v1/chat/completions",
llm_router=None,
user_api_key_dict=ProxyUserAPIKeyAuth(api_key="sk-test"),
proxy_logging_obj=MagicMock(spec=ProxyLogging),
)
assert exc_info.value.type == ProxyErrorTypes.budget_exceeded
assert exc_info.value.code == "429"
assert "guardrail-tag" in exc_info.value.message
@pytest.mark.asyncio
async def test_common_processing_pre_call_logic_arms_auto_router_compression_before_guardrails(
self, monkeypatch