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.
This commit is contained in:
mateo-berri 2026-05-30 03:05:40 +00:00
parent 82c1a362b9
commit 523a984f0d
No known key found for this signature in database
2 changed files with 96 additions and 16 deletions

View file

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

View file

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