mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
82c1a362b9
commit
523a984f0d
2 changed files with 96 additions and 16 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue