From 523a984f0d40797829613a8cc1e096676f8324c1 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 30 May 2026 03:05:40 +0000 Subject: [PATCH] fix(batch_rate_limiter): resolve provider skip from trusted deployment creds Resolve the batch provider from router deployment credentials instead of the user-supplied custom_llm_provider request field, so an unrestricted key cannot spoof a skip-listed provider to bypass batch rate limiting. Strengthen the provider-skip test to assert the file download and descriptor work were short-circuited, and add a test that a spoofed provider still falls through to rate-limit evaluation. --- litellm/proxy/hooks/batch_rate_limiter.py | 50 +++++++++++---- .../proxy/hooks/test_batch_file_validation.py | 62 +++++++++++++++++-- 2 files changed, 96 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index 787592f4e30..baeadc98c33 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -127,6 +127,36 @@ class _PROXY_BatchRateLimiter(CustomLogger): return None + def _resolve_batch_provider(self, batch_model: Optional[str]) -> Optional[str]: + """Resolve the provider from the deployment that serves ``batch_model``. + + The provider is read from trusted router credentials rather than the + user-supplied ``custom_llm_provider`` request field, so a caller cannot + spoof a skip-listed provider to bypass batch rate limiting. + """ + if not batch_model: + return None + + from litellm.proxy.openai_files_endpoints.common_utils import ( + get_credentials_for_model, + ) + from litellm.proxy.proxy_server import llm_router + + if llm_router is None: + return None + + try: + credentials = get_credentials_for_model( + llm_router=llm_router, + model_id=batch_model, + operation_context="batch input file read (rate limiting)", + ) + except HTTPException: + return None + + provider = credentials.get("custom_llm_provider") + return provider if isinstance(provider, str) and provider else None + def _matches_skip_list(self, value: str, skip_list: List[str]) -> bool: if not skip_list: return False @@ -164,8 +194,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): ``_enforce_batch_file_model_access`` can validate every ``body.model`` entry. Otherwise a restricted key could smuggle unauthorized models into the file via an admin-configured skip (global disable, - per-model, per-provider) or via the user-controlled - ``custom_llm_provider`` field. + per-model, per-provider). Returns ``(should_skip, descriptors)`` where ``descriptors`` is the rate-limit descriptor list computed for the no-limits check, so the @@ -197,16 +226,13 @@ class _PROXY_BatchRateLimiter(CustomLogger): general_settings.get("skip_batch_input_file_rate_limiting_for_providers") or [] ) - custom_llm_provider = data.get("custom_llm_provider") - if ( - isinstance(custom_llm_provider, str) - and custom_llm_provider in skip_providers - ): - verbose_proxy_logger.debug( - "Skipping batch input file processing for " - f"custom_llm_provider={custom_llm_provider}" - ) - return True, None + if skip_providers: + batch_provider = self._resolve_batch_provider(batch_model) + if batch_provider and batch_provider in skip_providers: + verbose_proxy_logger.debug( + f"Skipping batch input file processing for provider={batch_provider}" + ) + return True, None descriptors = self._create_batch_rate_limit_descriptors( user_api_key_dict=user_api_key_dict, diff --git a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py index 515be9df87c..f7e6660f5f0 100644 --- a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py +++ b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py @@ -293,22 +293,76 @@ async def test_pre_call_skips_file_fetch_for_configured_provider(): parallel_request_limiter=MagicMock(), ) user = UserAPIKeyAuth(api_key="sk-ok", user_id="alice", models=["*"]) + data = {"input_file_id": "file-abc123", "model": "my-vllm-model"} - with patch( - "litellm.proxy.proxy_server.general_settings", - {"skip_batch_input_file_rate_limiting_for_providers": ["hosted_vllm"]}, + with ( + patch( + "litellm.proxy.proxy_server.general_settings", + {"skip_batch_input_file_rate_limiting_for_providers": ["hosted_vllm"]}, + ), + patch("litellm.proxy.proxy_server.llm_router", MagicMock()), + patch( + "litellm.proxy.openai_files_endpoints.common_utils.get_credentials_for_model", + return_value={"custom_llm_provider": "hosted_vllm"}, + ), + patch("litellm.afile_content", new=AsyncMock()) as mock_afile_content, ): result = await rate_limiter.async_pre_call_hook( + user_api_key_dict=user, + cache=MagicMock(), + data=data, + call_type="acreate_batch", + ) + + assert result == data + # A real skip must short-circuit before any file download or rate-limit + # work — assert the skip happened rather than the hook's error-recovery + # path (which also returns data unchanged). + mock_afile_content.assert_not_awaited() + rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.assert_not_called() + + +@pytest.mark.asyncio +async def test_pre_call_does_not_skip_for_spoofed_provider(): + """The provider skip is resolved from trusted deployment credentials, so a + user-supplied ``custom_llm_provider`` that is not backed by the routing + deployment must not trigger a skip.""" + from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + + rate_limiter = _PROXY_BatchRateLimiter( + internal_usage_cache=MagicMock(), + parallel_request_limiter=MagicMock(), + ) + rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = ( + [] + ) + user = UserAPIKeyAuth(api_key="sk-ok", user_id="alice", models=["*"]) + + with ( + patch( + "litellm.proxy.proxy_server.general_settings", + {"skip_batch_input_file_rate_limiting_for_providers": ["hosted_vllm"]}, + ), + patch("litellm.proxy.proxy_server.llm_router", MagicMock()), + patch( + "litellm.proxy.openai_files_endpoints.common_utils.get_credentials_for_model", + return_value={"custom_llm_provider": "openai"}, + ), + ): + await rate_limiter.async_pre_call_hook( user_api_key_dict=user, cache=MagicMock(), data={ "input_file_id": "file-abc123", + "model": "my-openai-model", "custom_llm_provider": "hosted_vllm", }, call_type="acreate_batch", ) - assert result["custom_llm_provider"] == "hosted_vllm" + # Reaching descriptor evaluation proves the spoofed provider did not + # short-circuit the skip decision via the provider allow-list. + rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.assert_called_once() @pytest.mark.asyncio