mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
add streaming count, fix memory spikes on batch upload
This commit is contained in:
parent
3fbae90839
commit
794ace360b
12 changed files with 513 additions and 110 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue