fix(batches): skip unnecessary batch input file reads

Skip expensive pre-read of batch input files when no batch limits apply and model allowlist checks are not required, and decode model-embedded file IDs before file-content fetches to prevent upstream 404s.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Sameer Kankute 2026-05-28 11:45:17 +05:30
parent 5699a06413
commit 7afd657c46
No known key found for this signature in database
2 changed files with 315 additions and 5 deletions

View file

@ -17,7 +17,7 @@ Quick summary:
- async_log_success_event() fires on GET /v1/batches/{id} (batch completion)
"""
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Union
from fastapi import HTTPException
from pydantic import BaseModel
@ -25,12 +25,13 @@ from pydantic import BaseModel
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.batches.batch_utils import (
_extract_file_access_credentials,
_get_batch_job_input_file_usage,
_get_file_content_as_dictionary,
_get_models_from_batch_input_file_content,
)
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy._types import SpecialModelNames, UserAPIKeyAuth
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
@ -98,6 +99,189 @@ 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
input_file_id = data.get("input_file_id")
if not isinstance(input_file_id, str) or not input_file_id:
return None
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
decode_model_from_file_id,
get_models_from_unified_file_id,
)
model_from_file_id = decode_model_from_file_id(input_file_id)
if model_from_file_id:
return model_from_file_id
unified_file_id = _is_base64_encoded_unified_file_id(input_file_id)
if unified_file_id:
target_model_names = get_models_from_unified_file_id(unified_file_id)
if target_model_names:
return target_model_names[0]
return None
def _matches_skip_list(self, value: str, skip_list: List[str]) -> bool:
if not skip_list:
return False
for entry in skip_list:
if not isinstance(entry, str) or not entry:
continue
if value == entry or value.startswith(f"{entry}/"):
return True
return False
def _should_skip_batch_input_file_processing(
self,
data: Dict,
user_api_key_dict: UserAPIKeyAuth,
) -> bool:
"""
Skip downloading batch input files when configured or when there is
nothing to enforce (no applicable rate limits and no model allowlist).
"""
from litellm.proxy.proxy_server import general_settings
if general_settings.get("disable_batch_input_file_rate_limiting") is True:
return True
litellm_metadata = data.get("litellm_metadata") or {}
if litellm_metadata.get("skip_batch_input_file_rate_limiting") is True:
return True
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):
verbose_proxy_logger.debug(
f"Skipping batch input file processing for model={batch_model}"
)
return True
skip_providers = (
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
if not self._key_requires_batch_model_access_check(user_api_key_dict):
descriptors = self.parallel_request_limiter._create_rate_limit_descriptors(
user_api_key_dict=user_api_key_dict,
data=data,
rpm_limit_type=None,
tpm_limit_type=None,
model_has_failures=False,
)
if not self._has_applicable_batch_rate_limits(descriptors):
verbose_proxy_logger.debug(
"Skipping batch input file processing: no rate limits configured"
)
return True
return False
@staticmethod
def _key_requires_batch_model_access_check(
user_api_key_dict: UserAPIKeyAuth,
) -> bool:
"""True when the key may only call a subset of models (JSONL must be checked)."""
models = user_api_key_dict.models or []
if user_api_key_dict.access_group_ids:
return True
if not models:
return False
if "*" in models:
return False
if SpecialModelNames.all_proxy_models.value in models:
return False
return True
@staticmethod
def _has_applicable_batch_rate_limits(
descriptors: List["RateLimitDescriptor"],
) -> bool:
for descriptor in descriptors:
rate_limit = descriptor.get("rate_limit") or {}
if (
rate_limit.get("requests_per_unit") is not None
or rate_limit.get("tokens_per_unit") is not None
or rate_limit.get("max_parallel_requests") is not None
):
return True
return False
def _resolve_batch_input_file_fetch_params(
self,
file_id: str,
custom_llm_provider: str,
data: Dict,
) -> Tuple[str, Dict[str, Any]]:
"""
Map proxy-facing file IDs to provider file IDs and credentials.
Model-embedded IDs (``file-<base64>``) are not unified managed-file IDs;
without decoding them, ``afile_content`` is called with the encoded ID
and the upstream provider returns 404.
"""
from litellm.proxy.openai_files_endpoints.common_utils import (
decode_model_from_file_id,
get_credentials_for_model,
get_original_file_id,
)
from litellm.proxy.proxy_server import llm_router
fetch_kwargs: Dict[str, Any] = {
"custom_llm_provider": custom_llm_provider,
}
model_from_file_id = decode_model_from_file_id(file_id)
if model_from_file_id:
credentials = get_credentials_for_model(
llm_router=llm_router,
model_id=model_from_file_id,
operation_context="batch input file read (rate limiting)",
)
fetch_kwargs.update(_extract_file_access_credentials(credentials))
fetch_kwargs["model"] = model_from_file_id
provider = credentials.get("custom_llm_provider")
if provider:
fetch_kwargs["custom_llm_provider"] = provider
return get_original_file_id(file_id), fetch_kwargs
request_model = data.get("model")
if isinstance(request_model, str) and request_model and llm_router is not None:
try:
credentials = get_credentials_for_model(
llm_router=llm_router,
model_id=request_model,
operation_context="batch input file read (rate limiting)",
)
fetch_kwargs.update(_extract_file_access_credentials(credentials))
fetch_kwargs["model"] = request_model
provider = credentials.get("custom_llm_provider")
if provider:
fetch_kwargs["custom_llm_provider"] = provider
except HTTPException:
pass
return file_id, fetch_kwargs
def _raise_rate_limit_error(
self,
status: "RateLimitStatus",
@ -211,6 +395,7 @@ class _PROXY_BatchRateLimiter(CustomLogger):
file_id: str,
custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai",
user_api_key_dict: Optional[UserAPIKeyAuth] = None,
data: Optional[Dict] = None,
) -> BatchFileUsage:
"""
Count number of requests and tokens in a batch input file.
@ -238,14 +423,27 @@ class _PROXY_BatchRateLimiter(CustomLogger):
user_api_key_dict=user_api_key_dict,
)
else:
provider_file_id, fetch_kwargs = (
self._resolve_batch_input_file_fetch_params(
file_id=file_id,
custom_llm_provider=custom_llm_provider,
data=data or {},
)
)
# For non-managed files, use the standard litellm.afile_content
file_content = await litellm.afile_content(
file_id=file_id,
custom_llm_provider=custom_llm_provider,
file_id=provider_file_id,
user_api_key_dict=user_api_key_dict,
**fetch_kwargs,
)
file_content_as_dict = _get_file_content_as_dictionary(file_content.content)
file_content_bytes = getattr(file_content, "content", None)
if not isinstance(file_content_bytes, bytes):
raise ValueError(
f"Expected bytes content from file retrieval for {file_id}, "
f"got {type(file_content_bytes)}"
)
file_content_as_dict = _get_file_content_as_dictionary(file_content_bytes)
# Validate every model named in the batch JSONL against the
# caller's per-key model allowlist. Without this, a caller
@ -435,6 +633,11 @@ class _PROXY_BatchRateLimiter(CustomLogger):
)
return data
if self._should_skip_batch_input_file_processing(
data=data, user_api_key_dict=user_api_key_dict
):
return data
# Get custom_llm_provider for token counting
custom_llm_provider = data.get("custom_llm_provider", "openai")
@ -446,6 +649,7 @@ class _PROXY_BatchRateLimiter(CustomLogger):
file_id=input_file_id,
custom_llm_provider=custom_llm_provider,
user_api_key_dict=user_api_key_dict,
data=data,
)
verbose_proxy_logger.debug(

View file

@ -260,6 +260,112 @@ async def test_pre_call_allows_authorized_model_in_batch_file():
)
@pytest.mark.asyncio
async def test_pre_call_skips_file_fetch_when_disabled_in_general_settings():
from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter
rate_limiter = _PROXY_BatchRateLimiter(
internal_usage_cache=MagicMock(),
parallel_request_limiter=MagicMock(),
)
user = UserAPIKeyAuth(api_key="sk-ok", user_id="alice", models=["gpt-4o"])
with patch(
"litellm.proxy.proxy_server.general_settings",
{"disable_batch_input_file_rate_limiting": True},
):
result = await rate_limiter.async_pre_call_hook(
user_api_key_dict=user,
cache=MagicMock(),
data={"input_file_id": "file-abc123"},
call_type="acreate_batch",
)
assert result == {"input_file_id": "file-abc123"}
rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.assert_not_called()
@pytest.mark.asyncio
async def test_pre_call_skips_file_fetch_for_configured_provider():
from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter
rate_limiter = _PROXY_BatchRateLimiter(
internal_usage_cache=MagicMock(),
parallel_request_limiter=MagicMock(),
)
user = UserAPIKeyAuth(api_key="sk-ok", user_id="alice", models=["gpt-4o"])
with patch(
"litellm.proxy.proxy_server.general_settings",
{"skip_batch_input_file_rate_limiting_for_providers": ["hosted_vllm"]},
):
result = await rate_limiter.async_pre_call_hook(
user_api_key_dict=user,
cache=MagicMock(),
data={
"input_file_id": "file-abc123",
"custom_llm_provider": "hosted_vllm",
},
call_type="acreate_batch",
)
assert result["custom_llm_provider"] == "hosted_vllm"
@pytest.mark.asyncio
async def test_count_input_file_usage_decodes_model_embedded_file_id():
import base64
from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter
original_file_id = "file-provider-xyz"
encoded_payload = (
base64.urlsafe_b64encode(
f"litellm:{original_file_id};model,my-vllm-batch".encode()
)
.decode()
.rstrip("=")
)
encoded_file_id = f"file-{encoded_payload}"
rate_limiter = _PROXY_BatchRateLimiter(
internal_usage_cache=MagicMock(),
parallel_request_limiter=MagicMock(),
)
mock_content = MagicMock()
mock_content.content = b'{"custom_id": "1", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "my-vllm-batch", "messages": [{"role": "user", "content": "hi"}]}}\n'
with (
patch(
"litellm.afile_content",
new=AsyncMock(return_value=mock_content),
) as mock_afile_content,
patch(
"litellm.proxy.proxy_server.llm_router",
MagicMock(),
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_credentials_for_model",
return_value={
"api_key": "test-key",
"api_base": "http://vllm:8000/v1",
"custom_llm_provider": "hosted_vllm",
},
),
):
await rate_limiter.count_input_file_usage(
file_id=encoded_file_id,
custom_llm_provider="openai",
user_api_key_dict=UserAPIKeyAuth(api_key="sk-ok", user_id="alice"),
data={},
)
mock_afile_content.assert_awaited_once()
assert mock_afile_content.await_args.kwargs["file_id"] == original_file_id
assert mock_afile_content.await_args.kwargs["custom_llm_provider"] == "hosted_vllm"
@pytest.mark.asyncio
async def test_pre_call_skips_check_when_no_models_present():
"""Files without any `body.model` (corrupt or empty) must not 500;