diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 235bd85c5be..1fd81eccfd0 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -76,12 +76,14 @@ def _get_metadata_variable_name(request: Request) -> str: For all /thread or /assistant endpoints we need to call this "litellm_metadata" - For ALL other endpoints we call this "metadata + For ALL other endpoints we call this "metadata" """ - if RouteChecks._is_assistants_api_request(request): + path = request.url.path + + if "thread" in path or "assistant" in path: return "litellm_metadata" - if any(route in request.url.path for route in LITELLM_METADATA_ROUTES): + if any(route in path for route in LITELLM_METADATA_ROUTES): return "litellm_metadata" return "metadata" @@ -832,7 +834,8 @@ async def add_litellm_data_to_request( # noqa: PLR0915 from litellm.proxy.proxy_server import llm_router, premium_user from litellm.types.proxy.litellm_pre_call_utils import SecretFields - _headers = clean_headers( + _raw_headers: Dict[str, str] = dict(request.headers) + _headers: Dict[str, str] = clean_headers( request.headers, litellm_key_header_name=( general_settings.get("litellm_key_header_name") @@ -868,12 +871,10 @@ async def add_litellm_data_to_request( # noqa: PLR0915 ) ) - data.update( - LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( - headers=_headers, - data=data, - _metadata_variable_name=_metadata_variable_name, - ) + LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( + headers=_headers, + data=data, + _metadata_variable_name=_metadata_variable_name, ) # Add headers to metadata for guardrails to access (fixes #17477) @@ -900,7 +901,7 @@ async def add_litellm_data_to_request( # noqa: PLR0915 if "user" not in data: data["user"] = user - data["secret_fields"] = SecretFields(raw_headers=dict(request.headers)) + data["secret_fields"] = SecretFields(raw_headers=_raw_headers) ## Dynamic api version (Azure OpenAI endpoints) ## try: @@ -1044,10 +1045,6 @@ async def add_litellm_data_to_request( # noqa: PLR0915 ] = user_api_key_dict.user_max_budget data[_metadata_variable_name]["user_api_key_metadata"] = user_api_key_dict.metadata - _headers = dict(request.headers) - _headers.pop( - "authorization", None - ) # do not store the original `sk-..` api key in the db data[_metadata_variable_name]["headers"] = _headers data[_metadata_variable_name]["endpoint"] = str(request.url) @@ -1099,7 +1096,7 @@ async def add_litellm_data_to_request( # noqa: PLR0915 # Check if using tag based routing tags = LiteLLMProxyRequestSetup.add_request_tag_to_metadata( llm_router=llm_router, - headers=dict(request.headers), + headers=_headers, data=data, ) @@ -1128,20 +1125,13 @@ async def add_litellm_data_to_request( # noqa: PLR0915 if disabled_callbacks and isinstance(disabled_callbacks, list): data["litellm_disabled_callbacks"] = disabled_callbacks - # Guardrails from key/team metadata + # Guardrails from key/team metadata and policy engine move_guardrails_to_metadata( data=data, _metadata_variable_name=_metadata_variable_name, user_api_key_dict=user_api_key_dict, ) - # Guardrails from policy engine - add_guardrails_from_policy_engine( - data=data, - metadata_variable_name=_metadata_variable_name, - user_api_key_dict=user_api_key_dict, - ) - # Team Model Aliases _update_model_if_team_alias_exists( data=data, @@ -1482,6 +1472,29 @@ def move_guardrails_to_metadata( - Adds guardrails from policies attached to key/team metadata - Adds guardrails from policy engine based on team/key/model context """ + # Early-out: skip all guardrails processing when nothing is configured + key_metadata = user_api_key_dict.metadata + team_metadata = user_api_key_dict.team_metadata + + has_key_config = key_metadata and ( + "guardrails" in key_metadata or "policies" in key_metadata + ) + has_team_config = team_metadata and ( + "guardrails" in team_metadata or "policies" in team_metadata + ) + has_request_config = ( + "guardrails" in data or "guardrail_config" in data or "policies" in data + ) + + # Only check policy engine if no local config (avoid import + registry lookup) + if not (has_key_config or has_team_config or has_request_config): + from litellm.proxy.policy_engine.policy_registry import get_policy_registry + + if not get_policy_registry().is_initialized(): + # Nothing configured anywhere - clean up request body fields and return + data.pop("policies", None) + return + # Check key-level guardrails _add_guardrails_from_key_or_team_metadata( key_metadata=user_api_key_dict.metadata, diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index a8172e85910..e4b8613d204 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -15,6 +15,7 @@ from litellm.proxy.litellm_pre_call_utils import ( LiteLLMProxyRequestSetup, _get_dynamic_logging_metadata, _get_enforced_params, + _get_metadata_variable_name, _update_model_if_key_alias_exists, add_guardrails_from_policy_engine, add_litellm_data_to_request, @@ -47,6 +48,47 @@ def test_check_if_token_is_service_account(): assert check_if_token_is_service_account(other_metadata_token) == False +class TestGetMetadataVariableName: + """Tests for _get_metadata_variable_name()""" + + def _make_request(self, path: str) -> MagicMock: + request = MagicMock(spec=Request) + request.url.path = path + return request + + def test_returns_litellm_metadata_for_thread_routes(self): + request = self._make_request("/v1/threads/thread_123/messages") + assert _get_metadata_variable_name(request) == "litellm_metadata" + + def test_returns_litellm_metadata_for_assistant_routes(self): + request = self._make_request("/v1/assistants/asst_123") + assert _get_metadata_variable_name(request) == "litellm_metadata" + + def test_returns_litellm_metadata_for_batches_route(self): + request = self._make_request("/v1/batches") + assert _get_metadata_variable_name(request) == "litellm_metadata" + + def test_returns_litellm_metadata_for_messages_route(self): + request = self._make_request("/v1/messages") + assert _get_metadata_variable_name(request) == "litellm_metadata" + + def test_returns_litellm_metadata_for_files_route(self): + request = self._make_request("/v1/files") + assert _get_metadata_variable_name(request) == "litellm_metadata" + + def test_returns_metadata_for_chat_completions(self): + request = self._make_request("/chat/completions") + assert _get_metadata_variable_name(request) == "metadata" + + def test_returns_metadata_for_completions(self): + request = self._make_request("/v1/completions") + assert _get_metadata_variable_name(request) == "metadata" + + def test_returns_metadata_for_embeddings(self): + request = self._make_request("/v1/embeddings") + assert _get_metadata_variable_name(request) == "metadata" + + def test_get_enforced_params_for_service_account_settings(): """ Test that service account enforced params are only added to service account keys