diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index ae5905f9cdf..8722ba4b429 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -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") diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index 74e753b09ea..f1fec83525f 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -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, diff --git a/litellm/files/utils.py b/litellm/files/utils.py index a2b9a42c154..9639042caab 100644 --- a/litellm/files/utils.py +++ b/litellm/files/utils.py @@ -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: """ diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index fe34731759f..aa65cdc5fe0 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -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. diff --git a/litellm/llms/vertex_ai/files/transformation.py b/litellm/llms/vertex_ai/files/transformation.py index 3966bd9a018..09299b366f4 100644 --- a/litellm/llms/vertex_ai/files/transformation.py +++ b/litellm/llms/vertex_ai/files/transformation.py @@ -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, diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index 3957e3a7fbb..aed35b34b0e 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -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 diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index 378cbbda89c..1d2cc726154 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -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( diff --git a/litellm/types/router.py b/litellm/types/router.py index ef7eb05d087..ca7a9295450 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -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 diff --git a/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_streaming.py b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_streaming.py index 65b12191d0a..32f4eaaaa7b 100644 --- a/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_streaming.py +++ b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_streaming.py @@ -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", diff --git a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py index ae71d10b378..894e2e9171e 100644 --- a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py +++ b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py @@ -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() diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py index 5fc36b71f2b..7564844d54a 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py @@ -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 diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index cd235d8de67..81b6108d044 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -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