mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-19 00:01:29 +00:00
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:
parent
252c71c0b2
commit
e79cf49722
2 changed files with 149 additions and 2 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue