add streaming count, fix memory spikes on batch upload

This commit is contained in:
mubashir1osmani 2026-06-06 20:27:47 -07:00
parent 3fbae90839
commit 794ace360b
No known key found for this signature in database
GPG key ID: AB055FF67D0B4D9A
12 changed files with 513 additions and 110 deletions

View file

@ -13,7 +13,9 @@ from litellm import Router, verbose_logger
from litellm._uuid import uuid
from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data
from litellm.litellm_core_utils.prompt_templates.common_utils import (
extract_file_metadata,
)
from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
from litellm.llms.base_llm.managed_resources.isolation import (
build_list_page,
@ -981,9 +983,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
target_model_names_list: List[str],
) -> OpenAIFileObject:
## GET THE FILE TYPE FROM THE CREATE FILE REQUEST
file_data = extract_file_data(create_file_request["file"])
file_type = file_data["content_type"]
_, file_type = extract_file_metadata(create_file_request["file"])
output_file_id = file_objects[0].id
model_id = file_objects[0]._hidden_params.get("model_id")

View file

@ -1,5 +1,5 @@
import json
from typing import Any, List, Literal, Optional, Tuple
from typing import Any, Iterator, List, Literal, Optional, Tuple
import litellm
from litellm._logging import verbose_logger
@ -314,6 +314,46 @@ def _get_file_content_as_dictionary(file_content: bytes) -> List[dict]:
raise e
def _iter_batch_input_entries(file_content: bytes) -> Iterator[dict]:
"""
Yield batch input JSONL entries one at a time without materializing the whole
file as a list, so peak memory stays bounded when counting a large batch file.
"""
start, length, newline = 0, len(file_content), ord("\n")
while start < length:
idx = file_content.find(newline, start)
if idx == -1:
chunk, start = file_content[start:], length
else:
chunk, start = file_content[start:idx], idx + 1
line = chunk.strip()
if line:
yield json.loads(line)
def _count_entry_tokens(
entry: dict,
model_name: Optional[str] = None,
) -> int:
"""Token-count a single batch input entry's body (chat / text / embedding)."""
body = entry.get("body", {}) or {}
model = body.get("model", model_name or "")
messages = body.get("messages")
if messages:
return token_counter(model=model, messages=messages)
prompt = body.get("prompt")
if prompt:
return _count_prompt_or_input_tokens(model=model, value=prompt)
input_data = body.get("input")
if input_data:
return _count_prompt_or_input_tokens(model=model, value=input_data)
return 0
def _get_batch_job_cost_from_file_content(
file_content_dictionary: List[dict],
custom_llm_provider: Literal[
@ -431,27 +471,7 @@ def _get_batch_job_input_file_usage(
completion_tokens: int = 0
for _item in file_content_dictionary:
body = _item.get("body", {})
model = body.get("model", model_name or "")
# Chat completion payloads.
messages = body.get("messages")
if messages:
prompt_tokens += token_counter(model=model, messages=messages)
continue
# Text completion payloads (`prompt`).
prompt = body.get("prompt")
if prompt:
prompt_tokens += _count_prompt_or_input_tokens(model=model, value=prompt)
continue
# Embedding payloads (`input`).
input_data = body.get("input")
if input_data:
prompt_tokens += _count_prompt_or_input_tokens(
model=model, value=input_data
)
prompt_tokens += _count_entry_tokens(_item, model_name=model_name)
return Usage(
total_tokens=prompt_tokens + completion_tokens,

View file

@ -24,6 +24,20 @@ class FilesAPIUtils:
and extracted_file_data.get("content") is not None
)
@staticmethod
def is_batch_jsonl_request(
create_file_data: CreateFileRequest, content_type: Optional[str]
) -> bool:
"""
Batch-jsonl check from metadata only, so the body can stay a streamable
Path/handle instead of being read into memory.
"""
return (
create_file_data.get("purpose") == "batch"
and FilesAPIUtils.valid_content_type(content_type)
and create_file_data.get("file") is not None
)
@staticmethod
def valid_content_type(content_type: Optional[str]) -> bool:
"""

View file

@ -757,6 +757,46 @@ def update_responses_tools_with_model_file_ids(
return updated_tools
def extract_file_metadata(file_data: FileTypes) -> Tuple[Optional[str], Optional[str]]:
"""
Resolve (filename, content_type) without reading the file body.
Mirrors extract_file_data's metadata resolution but never calls .read(), so
it stays O(1) on large uploads. Use this when only metadata is needed (batch
detection, GCS object naming) and the body must remain a streamable Path/handle.
"""
filename: Optional[str] = None
content_type: Optional[str] = None
file_content: Any = None
if isinstance(file_data, tuple):
if len(file_data) == 2:
filename, file_content = file_data
elif len(file_data) == 3:
filename, file_content, content_type = file_data
elif len(file_data) == 4:
filename, file_content, content_type, _ = file_data
elif isinstance(file_data, InMemoryFile):
filename = file_data.name
content_type = file_data.content_type
else:
file_content = file_data
if filename is None:
if isinstance(file_content, PathLike):
filename = Path(file_content).name
elif isinstance(file_content, io.IOBase) and isinstance(
getattr(file_content, "name", None), str
):
filename = Path(file_content.name).name
if not content_type:
guessed = mimetypes.guess_type(filename)[0] if filename else None
content_type = guessed or "application/octet-stream"
return filename, content_type
def extract_file_data(file_data: FileTypes) -> ExtractedFileData:
"""
Extracts and processes file data from various input formats.

View file

@ -22,7 +22,10 @@ from litellm.litellm_core_utils.cloud_storage_security import (
validate_managed_cloud_file_id,
)
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data
from litellm.litellm_core_utils.prompt_templates.common_utils import (
extract_file_data,
extract_file_metadata,
)
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.base_llm.files.transformation import (
BaseFileUploadStream,
@ -47,7 +50,7 @@ from litellm.types.llms.openai import (
)
from litellm.types.files import ResumableChunkedUploadConfig
from litellm.types.llms.vertex_ai import GcsBucketResponse
from litellm.types.utils import ExtractedFileData, LlmProviders, ModelResponse
from litellm.types.utils import LlmProviders, ModelResponse
from ..common_utils import VertexAIError
from ..vertex_llm_base import VertexBase
@ -371,27 +374,21 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
object_name = f"{VERTEX_AI_MANAGED_GCS_PREFIX}{safe_model_path}/{uuid.uuid4()}"
return object_name
def get_object_name(
self, extracted_file_data: ExtractedFileData, purpose: str
) -> str:
def get_object_name(self, file_data: FileTypes, purpose: str) -> str:
"""
Get the object name for the request
Get the object name for the request.
Reads only the first JSONL entry (streamed) for batch files, so a large
upload is never materialized just to derive the GCS object name.
"""
extracted_file_data_content = extracted_file_data.get("content")
if extracted_file_data_content is None:
raise ValueError("file content is required")
if purpose == "batch":
## 1. If jsonl, derive the object name from the first entry's model
first_entry = next(
_iter_openai_jsonl_entries(extracted_file_data_content), None
)
first_entry = next(_iter_openai_jsonl_entries(file_data), None)
if first_entry is not None:
return self._get_gcs_object_name_from_batch_jsonl([first_entry])
## 2. If not jsonl, store under a server-generated managed object name
filename = extracted_file_data.get("filename")
filename, _ = extract_file_metadata(file_data)
return build_managed_cloud_object_name(
prefix=f"{VERTEX_AI_MANAGED_GCS_PREFIX}uploads/",
filename=filename,
@ -424,8 +421,8 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
raise ValueError("file is required")
if purpose is None:
raise ValueError("purpose is required")
extracted_file_data = extract_file_data(file_data)
object_name = self.get_object_name(extracted_file_data, purpose)
_, content_type = extract_file_metadata(file_data)
object_name = self.get_object_name(file_data, purpose)
if object_prefix:
object_name = f"{object_prefix}/{object_name}"
encoded_object_name = encode_gcs_object_name_for_url(object_name)
@ -433,8 +430,8 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
# large uploads); everything else is a single simple-media upload.
upload_type = (
"resumable"
if FilesAPIUtils.is_batch_jsonl_file(
create_file_data=data, extracted_file_data=extracted_file_data
if FilesAPIUtils.is_batch_jsonl_request(
create_file_data=data, content_type=content_type
)
else "media"
)
@ -504,20 +501,16 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
file_data = create_file_data.get("file")
if file_data is None:
raise ValueError("file is required")
extracted_file_data = extract_file_data(file_data)
extracted_file_data_content = extracted_file_data.get("content")
if extracted_file_data_content is None:
raise ValueError("file content is required")
if FilesAPIUtils.is_batch_jsonl_file(
_, content_type = extract_file_metadata(file_data)
if FilesAPIUtils.is_batch_jsonl_request(
create_file_data=create_file_data,
extracted_file_data=extracted_file_data,
content_type=content_type,
):
return {
"resumable_chunked_upload": ResumableChunkedUploadConfig(
body_stream=_OpenAIToVertexBatchUploadStream(
extracted_file_data_content,
file_data,
self._map_openai_to_vertex_params,
),
initiate_headers={
@ -525,10 +518,11 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
},
)
}
elif isinstance(extracted_file_data_content, bytes):
extracted_file_data_content = extract_file_data(file_data).get("content")
if isinstance(extracted_file_data_content, bytes):
return extracted_file_data_content
else:
raise ValueError("Unsupported file content type")
raise ValueError("Unsupported file content type")
def transform_create_file_response(
self,

View file

@ -21,6 +21,7 @@ from typing import (
TYPE_CHECKING,
Any,
Dict,
Iterable,
List,
Literal,
NoReturn,
@ -35,10 +36,9 @@ from pydantic import BaseModel
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.batches.batch_utils import (
_count_entry_tokens,
_extract_file_access_credentials,
_get_batch_job_input_file_usage,
_get_file_content_as_dictionary,
_get_models_from_batch_input_file_content,
_iter_batch_input_entries,
)
from litellm.exceptions import RateLimitErrorCategory
from litellm.integrations.custom_logger import CustomLogger
@ -562,7 +562,19 @@ class _PROXY_BatchRateLimiter(CustomLogger):
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)
# Single streaming pass over the JSONL entries: accumulate request
# count, distinct models, and token total without ever holding all
# entries in a list, so peak memory does not scale with file size.
models: set = set()
total_tokens = 0
request_count = 0
for entry in _iter_batch_input_entries(file_content_bytes):
request_count += 1
model = (entry.get("body") or {}).get("model")
if model:
models.add(model)
total_tokens += _count_entry_tokens(entry)
# Validate every model named in the batch JSONL against the
# caller's per-key model allowlist. Without this, a caller
@ -572,16 +584,11 @@ class _PROXY_BatchRateLimiter(CustomLogger):
if user_api_key_dict is not None:
await self._enforce_batch_file_model_access(
user_api_key_dict=user_api_key_dict,
file_content_as_dict=file_content_as_dict,
models=models,
)
input_file_usage = _get_batch_job_input_file_usage(
file_content_dictionary=file_content_as_dict,
custom_llm_provider=custom_llm_provider,
)
request_count = len(file_content_as_dict)
return BatchFileUsage(
total_tokens=input_file_usage.total_tokens,
total_tokens=total_tokens,
request_count=request_count,
)
@ -607,7 +614,7 @@ class _PROXY_BatchRateLimiter(CustomLogger):
async def _enforce_batch_file_model_access(
self,
user_api_key_dict: UserAPIKeyAuth,
file_content_as_dict: List[dict],
models: Iterable[str],
) -> None:
"""Reject the batch if the caller is not authorized for every
``body.model`` named inside the JSONL.
@ -627,7 +634,6 @@ class _PROXY_BatchRateLimiter(CustomLogger):
from litellm.proxy.proxy_server import proxy_logging_obj
from litellm.proxy.proxy_server import user_api_key_cache
models = _get_models_from_batch_input_file_content(file_content_as_dict)
if not models:
return

View file

@ -7,7 +7,7 @@
import asyncio
import traceback
from typing import Any, Optional, cast, get_args
from typing import Any, BinaryIO, Optional, Union, cast, get_args
import httpx
from fastapi import (
@ -91,16 +91,18 @@ def get_files_provider_config(
return None
def get_first_json_object(file_content_bytes: bytes) -> Optional[dict]:
def get_first_json_object(file_source: Union[bytes, BinaryIO]) -> Optional[dict]:
try:
# Decode the bytes to a string and split into lines
file_content = file_content_bytes.decode("utf-8")
first_line = file_content.splitlines()[0].strip()
# Parse the JSON object from the first line
json_object = json.loads(first_line)
return json_object
except (json.JSONDecodeError, UnicodeDecodeError):
if isinstance(file_source, (bytes, bytearray)):
newline = file_source.find(b"\n")
raw = file_source if newline == -1 else file_source[:newline]
first_line = raw.decode("utf-8")
else:
file_source.seek(0)
first_line = file_source.readline().decode("utf-8")
file_source.seek(0)
return json.loads(first_line.strip())
except (json.JSONDecodeError, UnicodeDecodeError, OSError, ValueError):
return None
@ -321,9 +323,15 @@ async def create_file( # noqa: PLR0915
data: Dict = {}
try:
# Use orjson to parse JSON data, orjson speeds up requests significantly
# Read the file content
file_content = await file.read()
# Batch uploads can be gigabytes. Starlette has already spooled the upload
# to disk, so stream from that handle instead of reading it into memory.
# Other uploads are small and stay in-memory bytes.
file_source: Union[bytes, BinaryIO]
if purpose == "batch":
await file.seek(0)
file_source = file.file
else:
file_source = await file.read()
custom_llm_provider = (
provider
or get_custom_llm_provider_from_request_headers(request=request)
@ -443,13 +451,13 @@ async def create_file( # noqa: PLR0915
)
# Prepare the file data according to FileTypes
file_data = (file.filename, file_content, file.content_type)
file_data = (file.filename, file_source, file.content_type)
## check if model is a loadbalanced model
router_model: Optional[str] = None
is_router_model = False
if litellm.enable_loadbalancing_on_batch_endpoints is True:
json_obj = get_first_json_object(file_content_bytes=file_content)
json_obj = get_first_json_object(file_source)
if json_obj:
router_model = get_model_from_json_obj(json_object=json_obj)
is_router_model = is_known_model(

View file

@ -173,6 +173,9 @@ class CredentialLiteLLMParams(BaseModel):
## UNIFIED PROJECT/REGION ##
region_name: Optional[str] = None
## OBJECT STORAGE (files / batches) ##
bucket_name: Optional[str] = None
## AWS BEDROCK / SAGEMAKER ##
aws_access_key_id: Optional[str] = None
aws_secret_access_key: Optional[str] = None

View file

@ -127,11 +127,11 @@ class TestFileLikeInputNotPartiallyConsumed:
"""
In ``llm_http_handler.create_file`` the object-name step
(get_complete_file_url -> get_object_name) runs before
transform_create_file_request, and each independently calls
``extract_file_data`` on the same create_file_data. When the file is a
tuple-wrapped open handle, the streaming reader must still emit every row
including entry 0: ``extract_file_data`` materializes the handle to bytes and
rewinds it (seek(0)), so neither step consumes the other's cursor. A partial
transform_create_file_request, and both read the same create_file_data
source. When the file is a tuple-wrapped open handle, the streaming reader
must still emit every row including entry 0: ``_iter_openai_jsonl_lines``
rewinds a seekable source (seek(0)) before each pass, so the object-name
step's partial read of the cursor does not consume the upload. A partial
upload missing the first request would be silent and hard to catch, so this
locks the full-payload invariant in.
"""
@ -244,12 +244,9 @@ class TestGetObjectNameLazyParse:
b'{"custom_id": "r-0", "body": {"model": "gemini-2.5-flash"}}\n'
b"garbage line that is not json\n"
)
from litellm.litellm_core_utils.prompt_templates.common_utils import (
extract_file_data,
object_name = cfg.get_object_name(
("batch.jsonl", raw, "application/jsonl"), purpose="batch"
)
extracted = extract_file_data(("batch.jsonl", raw, "application/jsonl"))
object_name = cfg.get_object_name(extracted, purpose="batch")
assert "gemini-2.5-flash" in object_name
@ -300,22 +297,111 @@ class TestStreamingPeakMemory:
def test_get_object_name_does_not_scale_with_payload(self):
cfg = VertexAIFilesConfig()
from litellm.litellm_core_utils.prompt_templates.common_utils import (
extract_file_data,
)
raw = _make_openai_jsonl_bytes(8000)
extracted = extract_file_data(("batch.jsonl", raw, "application/jsonl"))
file_data = ("batch.jsonl", raw, "application/jsonl")
# The payload bytes already exist before measurement starts, so a lazy
# first-row parse should allocate only a small fraction of the payload;
# parsing every row would blow past this bound.
peak = self._measure(lambda: cfg.get_object_name(extracted, purpose="batch"))
peak = self._measure(lambda: cfg.get_object_name(file_data, purpose="batch"))
assert (
peak / len(raw) < 2.0
), "get_object_name should not copy the whole payload"
class TestPathSourcedStreaming:
"""
The proxy spools large batch uploads to a temp file and passes a pathlib.Path
as the file content instead of pre-reading bytes, so the transform streams
from disk. These lock in that a Path source yields identical output, keeps
every row, stays memory-bounded, and is re-iterable (multi-model uploads).
"""
def _write_jsonl(self, tmp_path, n_rows, padding=400):
raw = _make_openai_jsonl_bytes(n_rows, padding=padding)
path = tmp_path / "batch.jsonl"
path.write_bytes(raw)
return path, raw
def _batch_request(self, path) -> CreateFileRequest:
return {"file": ("batch.jsonl", path, "application/jsonl"), "purpose": "batch"}
def test_transform_from_path_matches_legacy_and_keeps_all_rows(self, tmp_path):
cfg = VertexAIFilesConfig()
n_rows = 200
path, raw = self._write_jsonl(tmp_path, n_rows)
data = self._batch_request(path)
url = cfg.get_complete_file_url(
api_base=None,
api_key=None,
model="",
optional_params={},
litellm_params={"bucket_name": "test-bucket"},
data=data,
)
assert "uploadType=resumable" in url
out = cfg.transform_create_file_request(
model="", create_file_data=data, optional_params={}, litellm_params={}
)
assert isinstance(out, dict) and "resumable_chunked_upload" in out
body = _join_upload_body(out).decode("utf-8")
assert body == _legacy_vertex_jsonl_string(cfg, raw.decode("utf-8"))
lines = body.splitlines()
assert len(lines) == n_rows, "no batch row may be dropped from a Path source"
first_labels = json.loads(lines[0])["request"]["labels"]
assert _get_litellm_batch_custom_id_from_labels(first_labels) == "request-0"
def test_path_source_peak_stays_below_payload(self, tmp_path):
cfg = VertexAIFilesConfig()
path, raw = self._write_jsonl(tmp_path, 8000)
data = self._batch_request(path)
def run():
cfg.get_complete_file_url(
api_base=None,
api_key=None,
model="",
optional_params={},
litellm_params={"bucket_name": "test-bucket"},
data=data,
)
out = cfg.transform_create_file_request(
model="", create_file_data=data, optional_params={}, litellm_params={}
)
for _ in _resumable_stream(out).iter_bytes():
pass # drain without accumulating
gc.collect()
tracemalloc.start()
try:
run()
_, peak = tracemalloc.get_traced_memory()
finally:
tracemalloc.stop()
# Streaming from disk must not materialize the payload. Reading the whole
# file into bytes (the pre-fix path) would push peak past the file size.
assert peak < len(raw) * 0.3, (
f"peak {peak} not bounded vs payload {len(raw)} "
f"(ratio {peak / len(raw):.2f})"
)
def test_path_source_stream_is_reiterable(self, tmp_path):
cfg = VertexAIFilesConfig()
path, _ = self._write_jsonl(tmp_path, 50)
data = self._batch_request(path)
out = cfg.transform_create_file_request(
model="", create_file_data=data, optional_params={}, litellm_params={}
)
stream = _resumable_stream(out)
first = b"".join(stream.iter_bytes())
second = b"".join(stream.iter_bytes())
assert first == second and len(first) > 0
_GCS_OBJECT_JSON = {
"id": "test-bucket/litellm-vertex-files/x/123",
"name": "litellm-vertex-files/x",

View file

@ -14,6 +14,17 @@ from fastapi import HTTPException
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
def _models(file_content_as_dict):
"""Distinct body.model values, mirroring how the rate limiter collects the
models from a streamed batch file before the access check."""
return [
entry["body"]["model"]
for entry in file_content_as_dict
if (entry.get("body") or {}).get("model")
]
# ---------------------------------------------------------------------------
# Token counter — covers all three batch payload shapes
# ---------------------------------------------------------------------------
@ -211,7 +222,7 @@ async def test_pre_call_rejects_unauthorized_model_in_batch_file():
with pytest.raises(HTTPException) as exc:
await rate_limiter._enforce_batch_file_model_access(
user_api_key_dict=user,
file_content_as_dict=file_dict,
models=_models(file_dict),
)
assert exc.value.status_code == 403
@ -250,7 +261,7 @@ async def test_pre_call_allows_all_team_models_key_when_model_in_team_allowlist(
with patch("litellm.proxy.proxy_server.llm_router", None):
await rate_limiter._enforce_batch_file_model_access(
user_api_key_dict=user,
file_content_as_dict=file_dict,
models=_models(file_dict),
)
@ -297,7 +308,7 @@ async def test_pre_call_uses_current_team_allowlist_for_all_team_models_key():
):
await rate_limiter._enforce_batch_file_model_access(
user_api_key_dict=user,
file_content_as_dict=file_dict,
models=_models(file_dict),
)
assert exc_info.value.status_code == 403
@ -358,7 +369,7 @@ async def test_pre_call_allows_all_team_models_key_via_current_team_object():
):
await rate_limiter._enforce_batch_file_model_access(
user_api_key_dict=user,
file_content_as_dict=file_dict,
models=_models(file_dict),
)
mock_get_team_object.assert_awaited_once()
@ -421,7 +432,7 @@ async def test_pre_call_denies_all_team_models_key_via_member_scope():
):
await rate_limiter._enforce_batch_file_model_access(
user_api_key_dict=user,
file_content_as_dict=file_dict,
models=_models(file_dict),
)
assert exc_info.value.status_code == 403
@ -479,7 +490,7 @@ async def test_pre_call_fails_closed_when_current_team_fetch_fails_for_all_team_
):
await rate_limiter._enforce_batch_file_model_access(
user_api_key_dict=user,
file_content_as_dict=file_dict,
models=_models(file_dict),
)
assert exc_info.value.status_code == expected_status
@ -524,7 +535,7 @@ async def test_pre_call_allows_authorized_model_in_batch_file():
# Should not raise
await rate_limiter._enforce_batch_file_model_access(
user_api_key_dict=user,
file_content_as_dict=file_dict,
models=_models(file_dict),
)
@ -744,7 +755,7 @@ async def test_pre_call_allows_stripped_provider_model_when_key_has_proxy_alias(
):
await rate_limiter._enforce_batch_file_model_access(
user_api_key_dict=user,
file_content_as_dict=file_dict,
models=_models(file_dict),
)
can_key_call_model.assert_awaited_once()
@ -768,11 +779,11 @@ async def test_pre_call_skips_check_when_no_models_present():
# entirely.
await rate_limiter._enforce_batch_file_model_access(
user_api_key_dict=user,
file_content_as_dict=[],
models=_models([]),
)
await rate_limiter._enforce_batch_file_model_access(
user_api_key_dict=user,
file_content_as_dict=[{"body": {}}],
models=_models([{"body": {}}]),
)
@ -1295,3 +1306,117 @@ async def test_count_input_file_usage_raises_on_non_bytes_content():
user_api_key_dict=UserAPIKeyAuth(api_key="sk", models=["*"]),
data={},
)
# Streaming input counting — peak memory must not scale with a full dict list
# ---------------------------------------------------------------------------
def _make_batch_input_bytes(n_rows: int, padding: int = 200) -> bytes:
import json as _json
pad = "x" * padding
rows = []
for i in range(n_rows):
rows.append(
_json.dumps(
{
"custom_id": f"request-{i}",
"method": "POST",
"url": "/v1/chat/completions",
"body": {
"model": "gpt-4o" if i % 2 else "gpt-3.5-turbo",
"messages": [{"role": "user", "content": f"{pad} {i}"}],
},
}
)
)
return ("\n".join(rows)).encode("utf-8")
def test_iter_batch_input_entries_matches_dict_list():
from litellm.batches.batch_utils import (
_get_file_content_as_dictionary,
_iter_batch_input_entries,
)
raw = _make_batch_input_bytes(50)
streamed = list(_iter_batch_input_entries(raw))
assert streamed == _get_file_content_as_dictionary(raw)
assert streamed[0]["custom_id"] == "request-0"
# tolerant of blank lines and a missing trailing newline
assert list(_iter_batch_input_entries(raw + b"\n\n")) == streamed
def test_streaming_count_peak_below_dict_list():
import gc
import tracemalloc
from litellm.batches.batch_utils import (
_get_file_content_as_dictionary,
_iter_batch_input_entries,
)
raw = _make_batch_input_bytes(8000)
def _measure(fn):
gc.collect()
tracemalloc.start()
try:
fn()
_, peak = tracemalloc.get_traced_memory()
finally:
tracemalloc.stop()
return peak
def _stream():
count = 0
models: set = set()
for entry in _iter_batch_input_entries(raw):
count += 1
model = (entry.get("body") or {}).get("model")
if model:
models.add(model)
return count
def _build_list():
return len(_get_file_content_as_dictionary(raw))
stream_peak = _measure(_stream)
list_peak = _measure(_build_list)
assert stream_peak < list_peak * 0.5, (
f"streaming count peak {stream_peak} is not a clear win over the dict "
f"list {list_peak} (ratio {stream_peak / list_peak:.2f})"
)
@pytest.mark.asyncio
async def test_count_input_file_usage_streams_without_building_list():
"""count_input_file_usage must count requests/tokens in one streaming pass.
Mocks the download; asserts the count is correct and that the dict-list
helper is never called (a revert to the list approach would call it)."""
from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter
rate_limiter = _PROXY_BatchRateLimiter(
internal_usage_cache=MagicMock(),
parallel_request_limiter=MagicMock(),
)
raw = _make_batch_input_bytes(10)
fake_content = MagicMock()
fake_content.content = raw
with (
patch("litellm.afile_content", new=AsyncMock(return_value=fake_content)),
patch(
"litellm.batches.batch_utils._get_file_content_as_dictionary"
) as mock_dict_list,
):
usage = await rate_limiter.count_input_file_usage(
file_id="file-not-managed",
custom_llm_provider="openai",
user_api_key_dict=None,
)
assert usage.request_count == 10
assert usage.total_tokens > 0
mock_dict_list.assert_not_called()

View file

@ -378,6 +378,83 @@ def test_mock_create_audio_file(mocker: MockerFixture, monkeypatch, llm_router:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
def test_create_file_batch_streams_from_upload_spool(monkeypatch, llm_router: Router):
"""
Batch uploads must be passed downstream as the upload's streamable file handle
(Starlette's already-spooled file), not read into an in-memory bytes object, so
the proxy never buffers the whole payload. Non-batch uploads keep the bytes path.
"""
import litellm.proxy.proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.openai_files_endpoints import files_endpoints as fe
from litellm.types.llms.openai import OpenAIFileObject
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", llm_router)
setup_proxy_logging_object(monkeypatch, llm_router)
captured: dict = {}
async def fake_route_create_file(*, _create_file_request, **kwargs):
file_elem = _create_file_request["file"][1]
captured["file_elem"] = file_elem
if hasattr(file_elem, "read") and hasattr(file_elem, "seek"):
file_elem.seek(0)
captured["streamed_content"] = file_elem.read()
return OpenAIFileObject(
id="dummy-id",
object="file",
bytes=0,
created_at=1234567890,
filename="batch.jsonl",
purpose="batch",
status="uploaded",
)
monkeypatch.setattr(fe, "route_create_file", fake_route_create_file)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test-user"
)
content = (
b'{"custom_id":"r-0","method":"POST","url":"/v1/chat/completions",'
b'"body":{"model":"gpt-3.5-turbo","messages":[{"role":"user","content":"hi"}]}}\n'
)
try:
resp = client.post(
"/v1/files",
files={"file": ("batch.jsonl", content, "application/jsonl")},
data={"purpose": "batch"},
headers={"Authorization": "Bearer test-key"},
)
assert resp.status_code == 200, resp.text
file_elem = captured["file_elem"]
assert not isinstance(
file_elem, (bytes, bytearray)
), "batch upload must be a streamable handle, not in-memory bytes"
assert hasattr(file_elem, "read") and hasattr(
file_elem, "seek"
), "batch upload must be a seekable file handle"
assert (
captured["streamed_content"] == content
), "the handle must stream the uploaded bytes"
captured.clear()
resp = client.post(
"/v1/files",
files={"file": ("data.jsonl", content, "application/jsonl")},
data={"purpose": "user_data"},
headers={"Authorization": "Bearer test-key"},
)
assert resp.status_code == 200, resp.text
assert isinstance(
captured["file_elem"], (bytes, bytearray)
), "non-batch upload must stay in-memory bytes"
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.flaky(retries=3, delay=2)
def test_target_storage_invokes_storage_backend(
mocker: MockerFixture, monkeypatch, llm_router: Router

View file

@ -2778,6 +2778,36 @@ def test_get_deployment_credentials_with_provider_aws_bedrock_runtime_endpoint()
assert credentials["custom_llm_provider"] == "bedrock"
def test_get_deployment_credentials_with_provider_includes_bucket_name():
"""
Regression: bucket_name must survive the CredentialLiteLLMParams filter so
managed-files batch retrieval can resolve the GCS/S3 bucket. Previously it was
dropped, causing "GCS bucket_name is required" when fetching batch output files.
"""
router = litellm.Router(
model_list=[
{
"model_name": "vertex-gemini",
"litellm_params": {
"model": "vertex_ai/gemini-3.5-flash",
"vertex_project": "my-project",
"vertex_location": "global",
"bucket_name": "my-batch-bucket",
},
}
],
)
credentials = router.get_deployment_credentials_with_provider(
model_id="vertex-gemini"
)
assert credentials is not None
assert credentials["bucket_name"] == "my-batch-bucket"
assert credentials["vertex_project"] == "my-project"
assert credentials["custom_llm_provider"] == "vertex_ai"
def test_get_deployment_credentials_with_provider_resolves_credential_name():
"""
Test that get_deployment_credentials_with_provider correctly resolves