diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index 8f5126b907c..ec59f2e5e99 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -99,12 +99,15 @@ class _PROXY_BatchRateLimiter(CustomLogger): self.internal_usage_cache = internal_usage_cache self.parallel_request_limiter = parallel_request_limiter - def _get_batch_routing_model(self, data: Dict) -> Optional[str]: - """Resolve the deployment/model used for this batch from request data.""" - model = data.get("model") - if isinstance(model, str) and model: - return model + def _get_file_bound_batch_model(self, data: Dict) -> Optional[str]: + """Resolve the model bound to the batch input file ID. + The model embedded in a ``file-`` ID or a unified managed + file's target model is fixed when the file is created, so it reflects + the model the batch will actually run. Unlike the client-supplied + top-level ``model`` field, it cannot be swapped per request to point a + skip decision at a deployment the JSONL never routes to. + """ input_file_id = data.get("input_file_id") if not isinstance(input_file_id, str) or not input_file_id: return None @@ -127,6 +130,14 @@ class _PROXY_BatchRateLimiter(CustomLogger): return None + def _get_batch_routing_model(self, data: Dict) -> Optional[str]: + """Resolve the deployment/model used for this batch from request data.""" + model = data.get("model") + if isinstance(model, str) and model: + return model + + return self._get_file_bound_batch_model(data) + def _resolve_batch_provider(self, batch_model: Optional[str]) -> Optional[str]: """Resolve the provider from the deployment that serves ``batch_model``. @@ -196,6 +207,12 @@ class _PROXY_BatchRateLimiter(CustomLogger): into the file via an admin-configured skip (global disable, per-model, per-provider). + The per-model skip is matched against the file-bound model only, never + the client-supplied top-level ``model``. The latter selects routing + credentials but not the models the batch runs (those are the JSONL + ``body.model`` entries), so honoring it would let a caller name a + skip-listed deployment while routing a different, rate-limited model. + Returns ``(should_skip, descriptors)`` where ``descriptors`` is the rate-limit descriptor list computed for the no-limits check, so the caller can reuse it for counter enforcement without recomputing. @@ -208,13 +225,13 @@ class _PROXY_BatchRateLimiter(CustomLogger): if general_settings.get("disable_batch_input_file_rate_limiting") is True: return True, None - batch_model = self._get_batch_routing_model(data) skip_models = ( general_settings.get("skip_batch_input_file_rate_limiting_for_models") or [] ) - if batch_model and self._matches_skip_list(batch_model, skip_models): + file_bound_model = self._get_file_bound_batch_model(data) + if file_bound_model and self._matches_skip_list(file_bound_model, skip_models): verbose_proxy_logger.debug( - f"Skipping batch input file processing for model={batch_model}" + f"Skipping batch input file processing for model={file_bound_model}" ) return True, None @@ -223,7 +240,9 @@ class _PROXY_BatchRateLimiter(CustomLogger): or [] ) if skip_providers: - batch_provider = self._resolve_batch_provider(batch_model) + batch_provider = self._resolve_batch_provider( + self._get_batch_routing_model(data) + ) if batch_provider and batch_provider in skip_providers: verbose_proxy_logger.debug( f"Skipping batch input file processing for provider={batch_provider}" 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 504202d64fa..386f7373b11 100644 --- a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py +++ b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py @@ -661,8 +661,41 @@ def test_should_skip_ignores_client_supplied_metadata_flag(): def test_should_skip_honors_per_model_skip_list(): + """The per-model skip fires for a model bound to the input file ID (here a + model-embedded ``file-``), which reflects the model the batch runs.""" + import base64 + rate_limiter = _make_rate_limiter() user = UserAPIKeyAuth(api_key="sk", models=["*"]) + encoded = ( + base64.urlsafe_b64encode(b"litellm:file-xyz;model,gpt-4o-mini") + .decode() + .rstrip("=") + ) + with patch( + "litellm.proxy.proxy_server.general_settings", + {"skip_batch_input_file_rate_limiting_for_models": ["gpt-4o-mini"]}, + ): + should_skip, descriptors = ( + rate_limiter._should_skip_batch_input_file_processing( + data={"input_file_id": f"file-{encoded}"}, + user_api_key_dict=user, + ) + ) + assert should_skip is True + assert descriptors is None + + +def test_should_not_skip_per_model_for_spoofed_top_level_model(): + """A caller must not bypass batch rate limits by naming a skip-listed model + in the top-level ``model`` while routing a different model through the JSONL + ``body.model`` entries. The per-model skip only trusts the file-bound model, + so a skip-listed top-level model over a plain file still gets processed.""" + rate_limiter = _make_rate_limiter() + rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [ + {"rate_limit": {"requests_per_unit": 5}} + ] + user = UserAPIKeyAuth(api_key="sk", models=["*"]) with patch( "litellm.proxy.proxy_server.general_settings", {"skip_batch_input_file_rate_limiting_for_models": ["gpt-4o-mini"]}, @@ -673,8 +706,7 @@ def test_should_skip_honors_per_model_skip_list(): user_api_key_dict=user, ) ) - assert should_skip is True - assert descriptors is None + assert should_skip is False def test_should_skip_when_no_rate_limits_configured():