mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-19 00:01:29 +00:00
Merge pull request #40842 from BerriAI/litellm_guardrail_tag_budget_enforcement
fix(proxy): enforce tag budgets for tags added by guardrails
This commit is contained in:
commit
617a40bb1c
4 changed files with 421 additions and 19 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -51,6 +51,8 @@ from litellm.proxy.auth.auth_checks import (
|
|||
get_key_object,
|
||||
get_user_object,
|
||||
invalidate_team_member_spend_state,
|
||||
request_skips_budget_checks,
|
||||
route_skips_budget_checks,
|
||||
vector_store_access_check,
|
||||
)
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
|
|
@ -8447,3 +8449,15 @@ async def test_access_group_model_fallback_uses_the_injected_database(channel: s
|
|||
llm_router=None, prisma_client=client,
|
||||
) is True
|
||||
reader.assert_awaited_once_with(where={"access_group_id": "group-a"})
|
||||
|
||||
|
||||
def test_route_skips_budget_checks_marks_only_spend_free_routes() -> None:
|
||||
assert route_skips_budget_checks(route="/v1/models") is True
|
||||
assert route_skips_budget_checks(route="/spend/logs") is True
|
||||
assert route_skips_budget_checks(route="/health") is False
|
||||
assert route_skips_budget_checks(route="/v1/chat/completions") is False
|
||||
|
||||
|
||||
def test_request_skips_budget_checks_extends_route_rule_with_zero_cost_models() -> None:
|
||||
assert request_skips_budget_checks(route="/v1/models", model=None, llm_router=None) is True
|
||||
assert request_skips_budget_checks(route="/v1/chat/completions", model=None, llm_router=None) is False
|
||||
|
|
|
|||
|
|
@ -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,316 @@ 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_enforce_guardrail_added_tag_budgets_still_checks_when_model_is_unparseable(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())
|
||||
|
||||
await _enforce_guardrail_added_tag_budgets(
|
||||
data={"model": 5, "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),
|
||||
)
|
||||
|
||||
tag_budget_check.assert_awaited_once()
|
||||
|
||||
@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