mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
fix(proxy): enforce tag budgets for key-level tags (#29108)
* fix(proxy): enforce tag budgets for key-level tags Merge API key metadata.tags into request_data before _tag_max_budget_check so per-tag budgets apply when tags are set on the key at creation time. Co-authored-by: Cursor <cursoragent@cursor.com> * fix(auth): avoid false reject for key-inherited tags Run reject_clientside_metadata_tags before key-tag injection, then inject key metadata tags immediately before tag budget checks so key tags still enforce budgets without being treated as client-supplied tags. Co-authored-by: Cursor <cursoragent@cursor.com> --------- Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
d5d6b26a72
commit
eef1ec3e8d
4 changed files with 288 additions and 1 deletions
|
|
@ -619,6 +619,9 @@ async def common_checks( # noqa: PLR0915
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
# Run before apply_key_tags_pre_auth injects key metadata.tags into request_body.
|
||||
_reject_clientside_metadata_tags_check(general_settings, request_body, route)
|
||||
|
||||
# If this is a free model, skip all budget checks
|
||||
if not skip_budget_checks:
|
||||
# 3. If team is in budget
|
||||
|
|
@ -660,6 +663,14 @@ async def common_checks( # noqa: PLR0915
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
if valid_token is not None:
|
||||
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
|
||||
|
||||
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
|
||||
request_data=request_body,
|
||||
user_api_key_dict=valid_token,
|
||||
)
|
||||
|
||||
with tracer.trace("litellm.proxy.auth.common_checks.tag_max_budget_check"):
|
||||
await _tag_max_budget_check(
|
||||
request_body=request_body,
|
||||
|
|
@ -709,7 +720,6 @@ async def common_checks( # noqa: PLR0915
|
|||
await _check_end_user_budget(end_user_obj=end_user_object, route=route)
|
||||
|
||||
_enforce_user_param_check(general_settings, request, request_body, route)
|
||||
_reject_clientside_metadata_tags_check(general_settings, request_body, route)
|
||||
_global_proxy_budget_check(global_proxy_spend, skip_budget_checks, route)
|
||||
_guardrail_modification_check(request_body, team_object)
|
||||
|
||||
|
|
|
|||
|
|
@ -1193,6 +1193,36 @@ class LiteLLMProxyRequestSetup:
|
|||
|
||||
return tags
|
||||
|
||||
@staticmethod
|
||||
def apply_key_tags_pre_auth(
|
||||
request_data: dict,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> None:
|
||||
"""Merge key metadata tags into request_data before _tag_max_budget_check."""
|
||||
key_metadata = user_api_key_dict.metadata
|
||||
if not key_metadata:
|
||||
return
|
||||
|
||||
key_tags = key_metadata.get("tags")
|
||||
if not key_tags or not isinstance(key_tags, list):
|
||||
return
|
||||
|
||||
_metadata_variable_name = get_metadata_variable_name_from_kwargs(request_data)
|
||||
metadata = request_data.get(_metadata_variable_name)
|
||||
if isinstance(metadata, str):
|
||||
parsed = safe_json_loads(metadata)
|
||||
metadata = parsed if isinstance(parsed, dict) else {}
|
||||
request_data[_metadata_variable_name] = metadata
|
||||
elif not isinstance(metadata, dict):
|
||||
metadata = {}
|
||||
request_data[_metadata_variable_name] = metadata
|
||||
|
||||
existing_tags = metadata.get("tags")
|
||||
metadata["tags"] = LiteLLMProxyRequestSetup._merge_tags(
|
||||
request_tags=existing_tags if isinstance(existing_tags, list) else None,
|
||||
tags_to_add=key_tags,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def apply_client_tag_policy_pre_auth(
|
||||
request: Request,
|
||||
|
|
|
|||
|
|
@ -1625,6 +1625,50 @@ async def test_reject_clientside_metadata_tags_non_llm_route():
|
|||
assert result is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reject_clientside_metadata_tags_allows_key_tags_without_client_tags():
|
||||
"""Key metadata.tags are injected after the reject check; requests without
|
||||
client metadata.tags must not be blocked when reject_clientside_metadata_tags is on."""
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy.auth.auth_checks import common_checks
|
||||
|
||||
request_body = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [{"role": "user", "content": "test"}],
|
||||
}
|
||||
|
||||
general_settings = {"reject_clientside_metadata_tags": True}
|
||||
mock_request = MagicMock(spec=Request)
|
||||
valid_token = UserAPIKeyAuth(
|
||||
token="test-token",
|
||||
models=["gpt-3.5-turbo"],
|
||||
metadata={"tags": ["engineering"]},
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.auth.auth_checks.get_tag_objects_batch",
|
||||
new_callable=AsyncMock,
|
||||
return_value={},
|
||||
):
|
||||
result = await common_checks(
|
||||
request_body=request_body,
|
||||
team_object=None,
|
||||
user_object=None,
|
||||
end_user_object=None,
|
||||
global_proxy_spend=None,
|
||||
general_settings=general_settings,
|
||||
route="/chat/completions",
|
||||
llm_router=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
valid_token=valid_token,
|
||||
request=mock_request,
|
||||
)
|
||||
|
||||
assert result is True
|
||||
assert request_body["metadata"]["tags"] == ["engineering"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_virtual_key_soft_budget_check_with_user_obj():
|
||||
"""Test _virtual_key_soft_budget_check includes user_email when user_obj is provided"""
|
||||
|
|
|
|||
|
|
@ -4235,6 +4235,209 @@ class TestApplyClientTagPolicyPreAuth:
|
|||
assert exc_info.value.max_budget == 0.10
|
||||
|
||||
|
||||
class TestApplyKeyTagsPreAuth:
|
||||
def test_merges_key_tags_into_metadata(self):
|
||||
data = {"model": "gpt-3.5-turbo"}
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="hashed-key",
|
||||
metadata={"tags": ["engineering", "production"]},
|
||||
team_metadata={},
|
||||
)
|
||||
|
||||
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
|
||||
request_data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
assert data["metadata"]["tags"] == ["engineering", "production"]
|
||||
|
||||
def test_unions_key_tags_with_existing_request_tags(self):
|
||||
data = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"metadata": {"tags": ["request-tag"]},
|
||||
}
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="hashed-key",
|
||||
metadata={"tags": ["key-tag", "request-tag"]},
|
||||
team_metadata={},
|
||||
)
|
||||
|
||||
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
|
||||
request_data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
# request-tag deduplicated; key-tag appended
|
||||
assert data["metadata"]["tags"] == ["request-tag", "key-tag"]
|
||||
|
||||
def test_no_key_tags_no_mutation(self):
|
||||
data = {"model": "gpt-3.5-turbo"}
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="hashed-key",
|
||||
metadata={},
|
||||
team_metadata={},
|
||||
)
|
||||
|
||||
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
|
||||
request_data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
assert "metadata" not in data or "tags" not in data.get("metadata", {})
|
||||
|
||||
def test_empty_key_metadata_no_mutation(self):
|
||||
data = {"model": "gpt-3.5-turbo"}
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="hashed-key",
|
||||
metadata={},
|
||||
team_metadata={},
|
||||
)
|
||||
|
||||
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
|
||||
request_data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
assert "metadata" not in data
|
||||
|
||||
def test_uses_litellm_metadata_when_present(self):
|
||||
data = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"litellm_metadata": {"foo": "bar"},
|
||||
}
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="hashed-key",
|
||||
metadata={"tags": ["key-tag"]},
|
||||
team_metadata={},
|
||||
)
|
||||
|
||||
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
|
||||
request_data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
assert data["litellm_metadata"]["tags"] == ["key-tag"]
|
||||
assert "tags" not in data.get("metadata", {})
|
||||
|
||||
def test_string_metadata_parsed_before_merge(self):
|
||||
data = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"metadata": '{"tags": ["existing"]}',
|
||||
}
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="hashed-key",
|
||||
metadata={"tags": ["key-tag"]},
|
||||
team_metadata={},
|
||||
)
|
||||
|
||||
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
|
||||
request_data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
assert isinstance(data["metadata"], dict)
|
||||
assert data["metadata"]["tags"] == ["existing", "key-tag"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_tags_visible_to_tag_max_budget_check(self):
|
||||
from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_TagTable
|
||||
from litellm.proxy.auth.auth_checks import _tag_max_budget_check
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
||||
data = {"model": "gpt-3.5-turbo"}
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="hashed-key",
|
||||
metadata={"tags": ["engineering"]},
|
||||
team_metadata={},
|
||||
)
|
||||
|
||||
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
|
||||
request_data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
tag_object = LiteLLM_TagTable(
|
||||
tag_name="engineering",
|
||||
spend=0.0,
|
||||
litellm_budget_table=LiteLLM_BudgetTable(max_budget=0.10),
|
||||
)
|
||||
|
||||
async def mock_get_current_spend(counter_key, fallback_spend):
|
||||
if counter_key == "spend:tag:engineering":
|
||||
return 0.50
|
||||
return fallback_spend
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.get_current_spend",
|
||||
mock_get_current_spend,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.auth.auth_checks.get_tag_objects_batch",
|
||||
new_callable=AsyncMock,
|
||||
return_value={"engineering": tag_object},
|
||||
),
|
||||
):
|
||||
with pytest.raises(litellm.BudgetExceededError) as exc_info:
|
||||
await _tag_max_budget_check(
|
||||
request_body=data,
|
||||
prisma_client=MagicMock(),
|
||||
user_api_key_cache=MagicMock(),
|
||||
proxy_logging_obj=ProxyLogging(user_api_key_cache=None),
|
||||
valid_token=UserAPIKeyAuth(token="test-token"),
|
||||
)
|
||||
assert exc_info.value.current_cost == 0.50
|
||||
assert exc_info.value.max_budget == 0.10
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_tags_within_budget_passes_check(self):
|
||||
from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_TagTable
|
||||
from litellm.proxy.auth.auth_checks import _tag_max_budget_check
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
||||
data = {"model": "gpt-3.5-turbo"}
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="hashed-key",
|
||||
metadata={"tags": ["engineering"]},
|
||||
team_metadata={},
|
||||
)
|
||||
|
||||
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
|
||||
request_data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
tag_object = LiteLLM_TagTable(
|
||||
tag_name="engineering",
|
||||
spend=0.05,
|
||||
litellm_budget_table=LiteLLM_BudgetTable(max_budget=0.10),
|
||||
)
|
||||
|
||||
async def mock_get_current_spend(counter_key, fallback_spend):
|
||||
if counter_key == "spend:tag:engineering":
|
||||
return 0.05
|
||||
return fallback_spend
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.get_current_spend",
|
||||
mock_get_current_spend,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.auth.auth_checks.get_tag_objects_batch",
|
||||
new_callable=AsyncMock,
|
||||
return_value={"engineering": tag_object},
|
||||
),
|
||||
):
|
||||
await _tag_max_budget_check(
|
||||
request_body=data,
|
||||
prisma_client=MagicMock(),
|
||||
user_api_key_cache=MagicMock(),
|
||||
proxy_logging_obj=ProxyLogging(user_api_key_cache=None),
|
||||
valid_token=UserAPIKeyAuth(token="test-token"),
|
||||
)
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Tests for #27516: provider hint resolution from deployment when the
|
||||
# user-facing model name has no provider prefix.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue