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:
Sameer Kankute 2026-05-29 00:09:02 +05:30 committed by GitHub
parent d5d6b26a72
commit eef1ec3e8d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 288 additions and 1 deletions

View file

@ -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)

View file

@ -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,

View file

@ -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"""

View file

@ -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.