Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 11 additions & 1 deletion litellm/proxy/auth/auth_checks.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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)

Expand Down
30 changes: 30 additions & 0 deletions litellm/proxy/litellm_pre_call_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -1193,6 +1193,36 @@ def add_request_tag_to_metadata(

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,
Expand Down
44 changes: 44 additions & 0 deletions tests/test_litellm/proxy/auth/test_auth_checks.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"""
Expand Down
203 changes: 203 additions & 0 deletions tests/test_litellm/proxy/test_litellm_pre_call_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -4235,6 +4235,209 @@ async def mock_get_current_spend(counter_key, fallback_spend):
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.
Expand Down
Loading