diff --git a/litellm/integrations/gcs_bucket/gcs_bucket.py b/litellm/integrations/gcs_bucket/gcs_bucket.py index 65296bafcf3..90057984235 100644 --- a/litellm/integrations/gcs_bucket/gcs_bucket.py +++ b/litellm/integrations/gcs_bucket/gcs_bucket.py @@ -6,12 +6,14 @@ import time from litellm._uuid import uuid from datetime import datetime, timedelta, timezone from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple -from urllib.parse import quote from litellm._logging import verbose_logger from litellm.constants import LITELLM_ASYNCIO_QUEUE_MAXSIZE from litellm.integrations.additional_logging_utils import AdditionalLoggingUtils from litellm.integrations.gcs_bucket.gcs_bucket_base import GCSBucketBase +from litellm.litellm_core_utils.cloud_storage_security import ( + sanitize_cloud_object_component, +) from litellm.proxy._types import CommonProxyErrors from litellm.types.integrations.base_health_check import IntegrationHealthCheckStatus from litellm.types.integrations.gcs_bucket import * @@ -335,7 +337,11 @@ class GCSBucketLogger(GCSBucketBase, AdditionalLoggingUtils): _litellm_params = kwargs.get("litellm_params", None) or {} _metadata = _litellm_params.get("metadata", None) or {} if "gcs_log_id" in _metadata: - object_name = _metadata["gcs_log_id"] + safe_log_id = sanitize_cloud_object_component( + _metadata.get("gcs_log_id"), fallback="" + ) + if safe_log_id: + object_name = f"{current_date}/custom-{uuid.uuid4().hex}-{safe_log_id}" return object_name @@ -367,8 +373,7 @@ class GCSBucketLogger(GCSBucketBase, AdditionalLoggingUtils): request_date_str=date_str, response_id=request_id, ) - encoded_object_name = quote(object_name, safe="") - response = await self.download_gcs_object(encoded_object_name) + response = await self.download_gcs_object(object_name) if response is not None: loaded_response = json.loads(response) diff --git a/litellm/integrations/gcs_bucket/gcs_bucket_base.py b/litellm/integrations/gcs_bucket/gcs_bucket_base.py index e84b37e689b..1c5e30777a2 100644 --- a/litellm/integrations/gcs_bucket/gcs_bucket_base.py +++ b/litellm/integrations/gcs_bucket/gcs_bucket_base.py @@ -11,6 +11,10 @@ from litellm.integrations.gcs_bucket.gcs_bucket_mock_client import ( from litellm._logging import verbose_logger from litellm.integrations.custom_batch_logger import CustomBatchLogger +from litellm.litellm_core_utils.cloud_storage_security import ( + encode_gcs_object_name_for_url, + split_configured_cloud_bucket_name, +) from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, @@ -133,8 +137,8 @@ class GCSBucketBase(CustomBatchLogger): - Returns: bucket_name="my-bucket", object_name="my-folder/dev/my-object" """ - if "/" in bucket_name: - bucket_name, prefix = bucket_name.split("/", 1) + bucket_name, prefix = split_configured_cloud_bucket_name(bucket_name) + if prefix: object_name = f"{prefix}/{object_name}" return bucket_name, object_name return bucket_name, object_name @@ -248,6 +252,7 @@ class GCSBucketBase(CustomBatchLogger): bucket_name=bucket_name, object_name=object_name, ) + object_name = encode_gcs_object_name_for_url(object_name) url = f"https://storage.googleapis.com/storage/v1/b/{bucket_name}/o/{object_name}?alt=media" @@ -288,6 +293,7 @@ class GCSBucketBase(CustomBatchLogger): bucket_name=bucket_name, object_name=object_name, ) + object_name = encode_gcs_object_name_for_url(object_name) url = f"https://storage.googleapis.com/storage/v1/b/{bucket_name}/o/{object_name}" @@ -334,10 +340,11 @@ class GCSBucketBase(CustomBatchLogger): bucket_name=bucket_name, object_name=object_name, ) + encoded_object_name = encode_gcs_object_name_for_url(object_name) response = await self.async_httpx_client.post( headers=headers, - url=f"https://storage.googleapis.com/upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={object_name}", + url=f"https://storage.googleapis.com/upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={encoded_object_name}", data=json_logged_payload, ) diff --git a/litellm/litellm_core_utils/cloud_storage_security.py b/litellm/litellm_core_utils/cloud_storage_security.py new file mode 100644 index 00000000000..3c60e0cd523 --- /dev/null +++ b/litellm/litellm_core_utils/cloud_storage_security.py @@ -0,0 +1,140 @@ +import posixpath +import re +from typing import Optional, Sequence, Tuple +from urllib.parse import quote, unquote + +from litellm._uuid import uuid + +VERTEX_AI_MANAGED_GCS_PREFIX = "litellm-vertex-files/" +BEDROCK_MANAGED_S3_BATCH_PREFIX = "litellm-bedrock-files-" +BEDROCK_MANAGED_S3_UPLOAD_PREFIX = "litellm-bedrock-files/" +BEDROCK_MANAGED_S3_OUTPUT_PREFIX = "litellm-batch-outputs/" +BEDROCK_MANAGED_S3_PREFIXES = ( + BEDROCK_MANAGED_S3_BATCH_PREFIX, + BEDROCK_MANAGED_S3_UPLOAD_PREFIX, + BEDROCK_MANAGED_S3_OUTPUT_PREFIX, +) + +_SAFE_OBJECT_COMPONENT_PATTERN = re.compile(r"[^A-Za-z0-9._-]+") + + +def sanitize_cloud_object_component( + value: Optional[str], fallback: str = "file" +) -> str: + if not isinstance(value, str): + return fallback + + component = posixpath.basename(value.replace("\\", "/")).strip() + if component in {"", ".", ".."}: + return fallback + + component = "".join( + "_" if ord(char) < 32 or ord(char) == 127 else char for char in component + ) + component = _SAFE_OBJECT_COMPONENT_PATTERN.sub("_", component) + component = component.strip("._") + if not component: + return fallback + return component[:255] + + +def sanitize_cloud_object_path(value: Optional[str], fallback: str = "file") -> str: + if not isinstance(value, str): + return fallback + + segments = [] + for segment in value.replace("\\", "/").split("/"): + sanitized_segment = sanitize_cloud_object_component(segment, fallback="") + if sanitized_segment: + segments.append(sanitized_segment) + + if not segments: + return fallback + return "/".join(segments) + + +def build_managed_cloud_object_name( + prefix: str, filename: Optional[str], fallback_filename: str = "file" +) -> str: + safe_filename = sanitize_cloud_object_component( + filename, fallback=fallback_filename + ) + return f"{prefix}{uuid.uuid4().hex}-{safe_filename}" + + +def _validate_cloud_object_path(object_name: str) -> None: + if not object_name: + raise ValueError("Cloud storage object name is required") + if object_name.startswith("/"): + raise ValueError("Cloud storage object name must be relative") + if any(ord(char) < 32 or ord(char) == 127 for char in object_name): + raise ValueError("Cloud storage object name contains control characters") + if any(segment in {".", ".."} for segment in object_name.split("/")): + raise ValueError("Cloud storage object name contains an invalid path segment") + + +def split_configured_cloud_bucket_name(bucket_name: str) -> Tuple[str, str]: + if not isinstance(bucket_name, str) or not bucket_name.strip(): + raise ValueError("Cloud storage bucket name is required") + + bucket_name = bucket_name.strip() + if "://" in bucket_name or "?" in bucket_name or "#" in bucket_name: + raise ValueError( + "Cloud storage bucket name must not include a URI scheme or query" + ) + if any(ord(char) < 32 or ord(char) == 127 for char in bucket_name): + raise ValueError("Cloud storage bucket name contains control characters") + + bucket, _, prefix = bucket_name.partition("/") + if not bucket: + raise ValueError("Cloud storage bucket name is required") + if "\\" in bucket: + raise ValueError("Cloud storage bucket name contains an invalid separator") + + prefix = prefix.strip("/") + if prefix: + _validate_cloud_object_path(prefix) + + return bucket, prefix + + +def encode_gcs_object_name_for_url(object_name: str) -> str: + return quote(unquote(object_name), safe="") + + +def encode_s3_object_key_for_url(object_key: str) -> str: + return quote(unquote(object_key), safe="/") + + +def validate_managed_cloud_file_id( + file_id: str, + scheme: str, + configured_bucket_name: str, + allowed_object_prefixes: Sequence[str], +) -> Tuple[str, str]: + decoded_file_id = unquote(file_id) + if not decoded_file_id.startswith(scheme): + raise ValueError(f"file_id must be a {scheme} URI") + + full_path = decoded_file_id[len(scheme) :] + if "/" not in full_path: + raise ValueError("file_id must include a cloud storage object name") + + bucket_name, object_name = full_path.split("/", 1) + configured_bucket, configured_prefix = split_configured_cloud_bucket_name( + configured_bucket_name + ) + if bucket_name != configured_bucket: + raise ValueError("file_id bucket does not match the configured storage bucket") + + _validate_cloud_object_path(object_name) + allowed_prefixes = tuple(allowed_object_prefixes) + if configured_prefix: + allowed_prefixes = tuple( + f"{configured_prefix.rstrip('/')}/{prefix}" for prefix in allowed_prefixes + ) + + if not object_name.startswith(allowed_prefixes): + raise ValueError("file_id must reference a LiteLLM-managed storage object") + + return bucket_name, object_name diff --git a/litellm/litellm_core_utils/initialize_dynamic_callback_params.py b/litellm/litellm_core_utils/initialize_dynamic_callback_params.py index ffb6436f387..f7e3470bab2 100644 --- a/litellm/litellm_core_utils/initialize_dynamic_callback_params.py +++ b/litellm/litellm_core_utils/initialize_dynamic_callback_params.py @@ -57,6 +57,11 @@ _supported_callback_params = [ "lunary_public_key", ] +_request_blocked_callback_params = { + "gcs_bucket_name", + "gcs_path_service_account", +} + def initialize_standard_callback_dynamic_params( kwargs: Optional[Dict] = None, @@ -71,6 +76,8 @@ def initialize_standard_callback_dynamic_params( if kwargs: # 1. Check top-level kwargs for param in _supported_callback_params: + if param in _request_blocked_callback_params: + continue if param in kwargs: _param_value = kwargs.get(param) validate_no_callback_env_reference( @@ -86,6 +93,8 @@ def initialize_standard_callback_dynamic_params( if isinstance(metadata, dict): for param in _supported_callback_params: + if param in _request_blocked_callback_params: + continue if param not in standard_callback_dynamic_params and param in metadata: _param_value = metadata.get(param) validate_no_callback_env_reference( diff --git a/litellm/llms/bedrock/files/handler.py b/litellm/llms/bedrock/files/handler.py index 13bd87a1f01..6061fc8c0c9 100644 --- a/litellm/llms/bedrock/files/handler.py +++ b/litellm/llms/bedrock/files/handler.py @@ -1,10 +1,15 @@ import asyncio import base64 +import os from typing import Any, Coroutine, Optional, Tuple, Union import httpx from litellm import LlmProviders +from litellm.litellm_core_utils.cloud_storage_security import ( + BEDROCK_MANAGED_S3_PREFIXES, + validate_managed_cloud_file_id, +) from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.types.llms.openai import ( FileContentRequest, @@ -35,7 +40,7 @@ class BedrockFilesHandler(BaseAWSLLM): The file ID can be in two formats: 1. Base64-encoded unified file ID containing: llm_output_file_id,s3://bucket/path - 2. Direct S3 URI: s3://bucket/path + 2. Direct S3 URI: s3://bucket/litellm-managed-prefix/path Args: file_id: Encoded file ID or direct S3 URI @@ -58,14 +63,16 @@ class BedrockFilesHandler(BaseAWSLLM): except Exception: pass - # If not base64 encoded or doesn't contain llm_output_file_id, assume it's already an S3 URI + # If not base64 encoded or doesn't contain llm_output_file_id, accept only + # explicit S3 URIs. Bucket and key validation happens before any S3 call. if file_id.startswith("s3://"): return file_id - # If it doesn't start with s3://, assume it's a direct S3 URI and add the prefix - return f"s3://{file_id}" + raise ValueError("file_id must be a managed LiteLLM S3 file id") - def _parse_s3_uri(self, s3_uri: str) -> Tuple[str, str]: + def _parse_s3_uri( + self, s3_uri: str, configured_bucket_name: str + ) -> Tuple[str, str]: """ Parse S3 URI to extract bucket name and object key. @@ -75,21 +82,22 @@ class BedrockFilesHandler(BaseAWSLLM): Returns: Tuple of (bucket_name, object_key) """ - if not s3_uri.startswith("s3://"): + return validate_managed_cloud_file_id( + file_id=s3_uri, + scheme="s3://", + configured_bucket_name=configured_bucket_name, + allowed_object_prefixes=BEDROCK_MANAGED_S3_PREFIXES, + ) + + def _get_configured_s3_bucket_name(self, optional_params: dict) -> str: + bucket_name = optional_params.get("s3_bucket_name") or os.getenv( + "AWS_S3_BUCKET_NAME" + ) + if not bucket_name: raise ValueError( - f"Invalid S3 URI format: {s3_uri}. Expected format: s3://bucket-name/path/to/file" + "S3 bucket_name is required. Set 's3_bucket_name' or AWS_S3_BUCKET_NAME." ) - - # Remove 's3://' prefix - path = s3_uri[5:] - - if "/" in path: - bucket_name, object_key = path.split("/", 1) - else: - bucket_name = path - object_key = "" - - return bucket_name, object_key + return bucket_name async def afile_content( self, @@ -119,7 +127,11 @@ class BedrockFilesHandler(BaseAWSLLM): # Extract S3 URI from file ID s3_uri = self._extract_s3_uri_from_file_id(file_id) - bucket_name, object_key = self._parse_s3_uri(s3_uri) + configured_bucket_name = self._get_configured_s3_bucket_name(optional_params) + bucket_name, object_key = self._parse_s3_uri( + s3_uri=s3_uri, + configured_bucket_name=configured_bucket_name, + ) # Get AWS credentials aws_region_name = self._get_aws_region_name( diff --git a/litellm/llms/bedrock/files/transformation.py b/litellm/llms/bedrock/files/transformation.py index 3007b54808c..6669363093b 100644 --- a/litellm/llms/bedrock/files/transformation.py +++ b/litellm/llms/bedrock/files/transformation.py @@ -2,6 +2,7 @@ import json import os import time from typing import Any, Dict, List, Optional, Tuple, Union +from urllib.parse import unquote import httpx from httpx import Headers, Response @@ -10,6 +11,14 @@ from openai.types.file_deleted import FileDeleted from litellm._logging import verbose_logger from litellm._uuid import uuid from litellm.files.utils import FilesAPIUtils +from litellm.litellm_core_utils.cloud_storage_security import ( + BEDROCK_MANAGED_S3_BATCH_PREFIX, + BEDROCK_MANAGED_S3_UPLOAD_PREFIX, + build_managed_cloud_object_name, + encode_s3_object_key_for_url, + sanitize_cloud_object_component, + split_configured_cloud_bucket_name, +) from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.base_llm.files.transformation import ( @@ -116,10 +125,13 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): if _model.startswith("bedrock/"): _model = _model[8:] - # Replace colons with hyphens for Bedrock S3 URI compliance - _model = _model.replace(":", "-") + safe_model = sanitize_cloud_object_component( + _model.replace(":", "-"), fallback="model" + ) - object_name = f"litellm-bedrock-files-{_model}-{uuid.uuid4()}.jsonl" + object_name = ( + f"{BEDROCK_MANAGED_S3_BATCH_PREFIX}{safe_model}-{uuid.uuid4()}.jsonl" + ) return object_name def get_object_name( @@ -146,12 +158,13 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): if len(openai_jsonl_content) > 0: return self._get_s3_object_name_from_batch_jsonl(openai_jsonl_content) - ## 2. If not jsonl, return the filename + ## 2. If not jsonl, store under a server-generated managed object name filename = extracted_file_data.get("filename") - if filename: - return filename - ## 3. If no file name, return timestamp - return str(int(time.time())) + return build_managed_cloud_object_name( + prefix=BEDROCK_MANAGED_S3_UPLOAD_PREFIX, + filename=filename, + fallback_filename="file", + ) def get_complete_file_url( self, @@ -172,6 +185,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): raise ValueError( "S3 bucket_name is required. Set 's3_bucket_name' in litellm_params or AWS_S3_BUCKET_NAME env var" ) + bucket_name, object_prefix = split_configured_cloud_bucket_name(bucket_name) s3_region_name = litellm_params.get("s3_region_name") or optional_params.get( "s3_region_name" @@ -188,14 +202,17 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): raise ValueError("purpose is required") extracted_file_data = extract_file_data(file_data) object_name = self.get_object_name(extracted_file_data, purpose) + if object_prefix: + object_name = f"{object_prefix}/{object_name}" + encoded_object_name = encode_s3_object_key_for_url(object_name) # S3 endpoint URL format s3_endpoint_url = ( optional_params.get("s3_endpoint_url") or f"https://s3.{aws_region_name}.amazonaws.com" - ) + ).rstrip("/") - return f"{s3_endpoint_url}/{bucket_name}/{object_name}" + return f"{s3_endpoint_url}/{bucket_name}/{encoded_object_name}" def get_supported_openai_params( self, model: str @@ -532,10 +549,12 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): if match1: # Pattern: https://s3.region.amazonaws.com/bucket/key region, bucket, key = match1.groups() + key = unquote(key) s3_uri = f"s3://{bucket}/{key}" elif match2: # Pattern: https://bucket.s3.region.amazonaws.com/key bucket, region, key = match2.groups() + key = unquote(key) s3_uri = f"s3://{bucket}/{key}" else: # Fallback: try to extract bucket and key from URL path @@ -545,6 +564,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): path_parts = parsed.path.lstrip("/").split("/", 1) if len(path_parts) >= 2: bucket, key = path_parts[0], path_parts[1] + key = unquote(key) s3_uri = f"s3://{bucket}/{key}" else: raise ValueError(f"Unable to parse S3 URL: {https_url}") @@ -722,7 +742,12 @@ class BedrockJsonlFilesTransformation: # Remove bedrock/ prefix if present if _model.startswith("bedrock/"): _model = _model[8:] - object_name = f"litellm-bedrock-files-{_model}-{uuid.uuid4()}.jsonl" + safe_model = sanitize_cloud_object_component( + _model.replace(":", "-"), fallback="model" + ) + object_name = ( + f"{BEDROCK_MANAGED_S3_BATCH_PREFIX}{safe_model}-{uuid.uuid4()}.jsonl" + ) return object_name def _get_content_from_openai_file(self, openai_file_content: FileTypes) -> str: diff --git a/litellm/llms/vertex_ai/files/handler.py b/litellm/llms/vertex_ai/files/handler.py index 6636bccd6a3..48683ba41ae 100644 --- a/litellm/llms/vertex_ai/files/handler.py +++ b/litellm/llms/vertex_ai/files/handler.py @@ -1,5 +1,5 @@ import asyncio -import urllib.parse +from urllib.parse import unquote from typing import Any, Coroutine, Optional, Tuple, Union import httpx @@ -9,6 +9,10 @@ from litellm.integrations.gcs_bucket.gcs_bucket_base import ( GCSBucketBase, GCSLoggingConfig, ) +from litellm.litellm_core_utils.cloud_storage_security import ( + VERTEX_AI_MANAGED_GCS_PREFIX, + validate_managed_cloud_file_id, +) from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.types.llms.openai import ( CreateFileRequest, @@ -112,34 +116,25 @@ class VertexAIFilesHandler(GCSBucketBase): ) ) - def _extract_bucket_and_object_from_file_id(self, file_id: str) -> Tuple[str, str]: + def _extract_bucket_and_object_from_file_id( + self, file_id: str, configured_bucket_name: str + ) -> Tuple[str, str]: """ - Extract bucket name and object path from URL-encoded file_id. + Validate and extract bucket name and object path from file_id. - Expected format: gs%3A%2F%2Fbucket-name%2Fpath%2Fto%2Ffile - Which decodes to: gs://bucket-name/path/to/file + Expected format: gs://bucket-name/litellm-vertex-files/path/to/file Returns: - tuple: (bucket_name, url_encoded_object_path) + tuple: (bucket_name, object_path) - bucket_name: "bucket-name" - - url_encoded_object_path: "path%2Fto%2Ffile" + - object_path: "litellm-vertex-files/path/to/file" """ - decoded_path = urllib.parse.unquote(file_id) - - if decoded_path.startswith("gs://"): - full_path = decoded_path[5:] # Remove 'gs://' prefix - else: - full_path = decoded_path - - if "/" in full_path: - bucket_name, object_path = full_path.split("/", 1) - else: - bucket_name = full_path - object_path = "" - - encoded_object_path = urllib.parse.quote(object_path, safe="") - - return bucket_name, encoded_object_path + return validate_managed_cloud_file_id( + file_id=file_id, + scheme="gs://", + configured_bucket_name=configured_bucket_name, + allowed_object_prefixes=(VERTEX_AI_MANAGED_GCS_PREFIX,), + ) async def afile_content( self, @@ -168,28 +163,34 @@ class VertexAIFilesHandler(GCSBucketBase): if not file_id: raise ValueError("file_id is required in file_content_request") - bucket_name, encoded_object_path = self._extract_bucket_and_object_from_file_id( - file_id + gcs_logging_config: GCSLoggingConfig = await self.get_gcs_logging_config( + kwargs={} + ) + bucket_name, object_path = self._extract_bucket_and_object_from_file_id( + file_id=file_id, + configured_bucket_name=gcs_logging_config["bucket_name"], ) download_kwargs = { - "standard_callback_dynamic_params": {"gcs_bucket_name": bucket_name} + "standard_callback_dynamic_params": { + "gcs_bucket_name": bucket_name, + "gcs_path_service_account": gcs_logging_config["path_service_account"], + } } file_content = await self.download_gcs_object( - object_name=encoded_object_path, **download_kwargs + object_name=object_path, **download_kwargs ) + decoded_file_id = unquote(file_id) if file_content is None: - decoded_path = urllib.parse.unquote(file_id) - raise ValueError(f"Failed to download file from GCS: {decoded_path}") + raise ValueError(f"Failed to download file from GCS: {decoded_file_id}") - decoded_path = urllib.parse.unquote(file_id) mock_response = httpx.Response( status_code=200, content=file_content, headers={"content-type": "application/octet-stream"}, - request=httpx.Request(method="GET", url=decoded_path), + request=httpx.Request(method="GET", url=decoded_file_id), ) return HttpxBinaryResponseContent(response=mock_response) diff --git a/litellm/llms/vertex_ai/files/transformation.py b/litellm/llms/vertex_ai/files/transformation.py index 070ec508283..b864ad43136 100644 --- a/litellm/llms/vertex_ai/files/transformation.py +++ b/litellm/llms/vertex_ai/files/transformation.py @@ -1,6 +1,5 @@ import json import os -import time from typing import Any, Dict, List, Optional, Tuple, Union from httpx import Headers, Response @@ -8,6 +7,14 @@ from openai.types.file_deleted import FileDeleted from litellm._uuid import uuid from litellm.files.utils import FilesAPIUtils +from litellm.litellm_core_utils.cloud_storage_security import ( + VERTEX_AI_MANAGED_GCS_PREFIX, + build_managed_cloud_object_name, + encode_gcs_object_name_for_url, + sanitize_cloud_object_path, + split_configured_cloud_bucket_name, + validate_managed_cloud_file_id, +) from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.base_llm.files.transformation import ( @@ -119,7 +126,8 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): _model = openai_jsonl_content[0].get("body", {}).get("model", "") if "publishers/google/models" not in _model: _model = f"publishers/google/models/{_model}" - object_name = f"litellm-vertex-files/{_model}/{uuid.uuid4()}" + safe_model_path = sanitize_cloud_object_path(_model, fallback="model") + object_name = f"{VERTEX_AI_MANAGED_GCS_PREFIX}{safe_model_path}/{uuid.uuid4()}" return object_name def get_object_name( @@ -146,12 +154,19 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): if len(openai_jsonl_content) > 0: return self._get_gcs_object_name_from_batch_jsonl(openai_jsonl_content) - ## 2. If not jsonl, return the filename + ## 2. If not jsonl, store under a server-generated managed object name filename = extracted_file_data.get("filename") - if filename: - return filename - ## 3. If no file name, return timestamp - return str(int(time.time())) + return build_managed_cloud_object_name( + prefix=f"{VERTEX_AI_MANAGED_GCS_PREFIX}uploads/", + filename=filename, + fallback_filename="file", + ) + + def _get_configured_bucket_name(self, litellm_params: Dict) -> str: + bucket_name = litellm_params.get("bucket_name") or os.getenv("GCS_BUCKET_NAME") + if not bucket_name: + raise ValueError("GCS bucket_name is required") + return bucket_name def get_complete_file_url( self, @@ -165,13 +180,8 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): """ Get the complete url for the request """ - bucket_name = ( - litellm_params.get("bucket_name") - or litellm_params.get("litellm_metadata", {}).pop("gcs_bucket_name", None) - or os.getenv("GCS_BUCKET_NAME") - ) - if not bucket_name: - raise ValueError("GCS bucket_name is required") + bucket_name = self._get_configured_bucket_name(litellm_params) + bucket_name, object_prefix = split_configured_cloud_bucket_name(bucket_name) file_data = data.get("file") purpose = data.get("purpose") if file_data is None: @@ -180,9 +190,10 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): raise ValueError("purpose is required") extracted_file_data = extract_file_data(file_data) object_name = self.get_object_name(extracted_file_data, purpose) - endpoint = ( - f"upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={object_name}" - ) + if object_prefix: + object_name = f"{object_prefix}/{object_name}" + encoded_object_name = encode_gcs_object_name_for_url(object_name) + endpoint = f"upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={encoded_object_name}" api_base = api_base or "https://storage.googleapis.com" if not api_base: raise ValueError("api_base is required") @@ -339,27 +350,20 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): status_code=status_code, message=error_message, headers=headers ) - def _parse_gcs_uri(self, file_id: str) -> Tuple[str, str]: + def _parse_gcs_uri( + self, file_id: str, litellm_params: Optional[Dict] = None + ) -> Tuple[str, str]: """ - Parse a GCS URI (gs://bucket/path/to/object) into (bucket, url-encoded-object-path). - Handles both raw and URL-encoded input. + Validate a managed GCS file_id and return (bucket, url-encoded-object-path). """ - import urllib.parse - - decoded = urllib.parse.unquote(file_id) - if decoded.startswith("gs://"): - full_path = decoded[5:] - else: - full_path = decoded - - if "/" in full_path: - bucket_name, object_path = full_path.split("/", 1) - else: - bucket_name = full_path - object_path = "" - - encoded_object = urllib.parse.quote(object_path, safe="") - return bucket_name, encoded_object + configured_bucket_name = self._get_configured_bucket_name(litellm_params or {}) + bucket_name, object_path = validate_managed_cloud_file_id( + file_id=file_id, + scheme="gs://", + configured_bucket_name=configured_bucket_name, + allowed_object_prefixes=(VERTEX_AI_MANAGED_GCS_PREFIX,), + ) + return bucket_name, encode_gcs_object_name_for_url(object_path) def transform_retrieve_file_request( self, @@ -367,7 +371,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): optional_params: dict, litellm_params: dict, ) -> tuple[str, dict]: - bucket, encoded_object = self._parse_gcs_uri(file_id) + bucket, encoded_object = self._parse_gcs_uri(file_id, litellm_params) url = f"https://storage.googleapis.com/storage/v1/b/{bucket}/o/{encoded_object}" return url, {} @@ -399,7 +403,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): optional_params: dict, litellm_params: dict, ) -> tuple[str, dict]: - bucket, encoded_object = self._parse_gcs_uri(file_id) + bucket, encoded_object = self._parse_gcs_uri(file_id, litellm_params) url = f"https://storage.googleapis.com/storage/v1/b/{bucket}/o/{encoded_object}" return url, {} @@ -443,7 +447,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): litellm_params: dict, ) -> tuple[str, dict]: file_id = file_content_request.get("file_id", "") - bucket, encoded_object = self._parse_gcs_uri(file_id) + bucket, encoded_object = self._parse_gcs_uri(file_id, litellm_params) url = f"https://storage.googleapis.com/storage/v1/b/{bucket}/o/{encoded_object}?alt=media" return url, {} @@ -528,7 +532,8 @@ class VertexAIJsonlFilesTransformation(VertexGeminiConfig): _model = openai_jsonl_content[0].get("body", {}).get("model", "") if "publishers/google/models" not in _model: _model = f"publishers/google/models/{_model}" - object_name = f"litellm-vertex-files/{_model}/{uuid.uuid4()}" + safe_model_path = sanitize_cloud_object_path(_model, fallback="model") + object_name = f"{VERTEX_AI_MANAGED_GCS_PREFIX}{safe_model_path}/{uuid.uuid4()}" return object_name def _map_openai_to_vertex_params( diff --git a/tests/test_litellm/integrations/gcs_bucket/test_gcs_bucket_base.py b/tests/test_litellm/integrations/gcs_bucket/test_gcs_bucket_base.py index d7b2a842bc0..a4e16500aee 100644 --- a/tests/test_litellm/integrations/gcs_bucket/test_gcs_bucket_base.py +++ b/tests/test_litellm/integrations/gcs_bucket/test_gcs_bucket_base.py @@ -1,5 +1,9 @@ import os -from unittest.mock import patch +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from litellm.integrations.gcs_bucket.gcs_bucket import GCSBucketLogger from litellm.integrations.gcs_bucket.gcs_bucket_base import GCSBucketBase @@ -86,3 +90,41 @@ class TestGCSBucketBase: project_id=None, # Should be None when no env var is set custom_llm_provider="vertex_ai", ) + + @pytest.mark.asyncio + async def test_log_json_data_on_gcs_url_encodes_object_name(self): + handler = GCSBucketBase(bucket_name="test-bucket") + handler.async_httpx_client = AsyncMock() + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = {"name": "logs/object"} + handler.async_httpx_client.post.return_value = mock_response + + await handler._log_json_data_on_gcs( + headers={"Authorization": "Bearer token"}, + bucket_name="test-bucket", + object_name="logs/object?uploadType=media&name=evil", + logging_payload={"ok": True}, + ) + + post_url = handler.async_httpx_client.post.call_args.kwargs["url"] + assert "name=logs%2Fobject%3FuploadType%3Dmedia%26name%3Devil" in post_url + assert "name=logs/object?" not in post_url + + def test_gcs_log_id_is_only_used_as_sanitized_hint(self): + logger = GCSBucketLogger.__new__(GCSBucketLogger) + + object_name = logger._get_object_name( + kwargs={ + "litellm_params": { + "metadata": {"gcs_log_id": "../../target?uploadType=media"} + } + }, + logging_payload={"id": "payload"}, + response_obj={"id": "response-id"}, + ) + + assert "/custom-" in object_name + assert object_name.endswith("-target_uploadType_media") + assert ".." not in object_name + assert "?" not in object_name diff --git a/tests/test_litellm/litellm_core_utils/test_initialize_dynamic_callback_params.py b/tests/test_litellm/litellm_core_utils/test_initialize_dynamic_callback_params.py index 55f3c2ba3aa..f63216b96d4 100644 --- a/tests/test_litellm/litellm_core_utils/test_initialize_dynamic_callback_params.py +++ b/tests/test_litellm/litellm_core_utils/test_initialize_dynamic_callback_params.py @@ -64,7 +64,7 @@ def test_env_reference_in_metadata_raises_with_guidance(): assert "metadata" in message -def test_env_reference_in_litellm_params_metadata_raises(): +def test_gcs_bucket_name_in_litellm_params_metadata_is_ignored(): kwargs = { "litellm_params": { "metadata": { @@ -73,10 +73,21 @@ def test_env_reference_in_litellm_params_metadata_raises(): } } - with pytest.raises(ValueError) as exc_info: - initialize_standard_callback_dynamic_params(kwargs) + params = initialize_standard_callback_dynamic_params(kwargs) - assert "gcs_bucket_name" in str(exc_info.value) + assert params.get("gcs_bucket_name") is None + + +def test_gcs_callback_params_are_not_extracted_from_request_kwargs(): + kwargs = { + "gcs_bucket_name": "server-bucket", + "gcs_path_service_account": "/path/to/service-account.json", + } + + params = initialize_standard_callback_dynamic_params(kwargs) + + assert params.get("gcs_bucket_name") is None + assert params.get("gcs_path_service_account") is None def test_non_string_values_are_not_flagged(): diff --git a/tests/test_litellm/llms/bedrock/files/test_bedrock_files_handler.py b/tests/test_litellm/llms/bedrock/files/test_bedrock_files_handler.py new file mode 100644 index 00000000000..b8375748295 --- /dev/null +++ b/tests/test_litellm/llms/bedrock/files/test_bedrock_files_handler.py @@ -0,0 +1,85 @@ +import base64 + +import pytest + +from litellm.llms.bedrock.files.handler import BedrockFilesHandler +from litellm.types.utils import SpecialEnums + + +def _encode_unified_file_id(s3_uri: str) -> str: + unified_file_id = SpecialEnums.LITELLM_MANAGED_FILE_COMPLETE_STR.value.format( + "application/json", + "unified-id", + "", + s3_uri, + "model-id", + ) + return base64.urlsafe_b64encode(unified_file_id.encode()).decode().rstrip("=") + + +class TestBedrockFilesHandler: + def setup_method(self): + self.handler = BedrockFilesHandler() + + def test_should_parse_direct_managed_s3_uri(self): + bucket, key = self.handler._parse_s3_uri( + s3_uri="s3://safe-bucket/litellm-bedrock-files-model-id-abc.jsonl", + configured_bucket_name="safe-bucket", + ) + + assert bucket == "safe-bucket" + assert key == "litellm-bedrock-files-model-id-abc.jsonl" + + def test_should_parse_managed_batch_output_uri(self): + bucket, key = self.handler._parse_s3_uri( + s3_uri="s3://safe-bucket/litellm-batch-outputs/job/", + configured_bucket_name="safe-bucket", + ) + + assert bucket == "safe-bucket" + assert key == "litellm-batch-outputs/job/" + + def test_should_reject_arbitrary_bucket(self): + with pytest.raises(ValueError, match="configured storage bucket"): + self.handler._parse_s3_uri( + s3_uri="s3://other-bucket/litellm-bedrock-files-model-id-abc.jsonl", + configured_bucket_name="safe-bucket", + ) + + def test_should_reject_unmanaged_same_bucket_key(self): + with pytest.raises(ValueError, match="LiteLLM-managed"): + self.handler._parse_s3_uri( + s3_uri="s3://safe-bucket/private/output.jsonl", + configured_bucket_name="safe-bucket", + ) + + def test_should_reject_dot_segment_key(self): + with pytest.raises(ValueError, match="invalid path segment"): + self.handler._parse_s3_uri( + s3_uri="s3://safe-bucket/litellm-bedrock-files/../secret.jsonl", + configured_bucket_name="safe-bucket", + ) + + def test_should_extract_unified_managed_s3_uri(self): + file_id = _encode_unified_file_id( + "s3://safe-bucket/litellm-batch-outputs/job/output.jsonl" + ) + + assert ( + self.handler._extract_s3_uri_from_file_id(file_id) + == "s3://safe-bucket/litellm-batch-outputs/job/output.jsonl" + ) + + def test_should_reject_file_id_without_s3_scheme(self): + with pytest.raises(ValueError, match="managed LiteLLM S3 file id"): + self.handler._extract_s3_uri_from_file_id("safe-bucket/private.jsonl") + + def test_should_reject_unified_unmanaged_s3_uri(self): + file_id = _encode_unified_file_id("s3://safe-bucket/private/output.jsonl") + s3_uri = self.handler._extract_s3_uri_from_file_id(file_id) + + with pytest.raises(ValueError, match="LiteLLM-managed"): + self.handler._parse_s3_uri( + s3_uri=s3_uri, + configured_bucket_name="safe-bucket", + ) diff --git a/tests/test_litellm/llms/bedrock/files/test_bedrock_files_transformation.py b/tests/test_litellm/llms/bedrock/files/test_bedrock_files_transformation.py index d9a2ddefd34..5245612e9d3 100644 --- a/tests/test_litellm/llms/bedrock/files/test_bedrock_files_transformation.py +++ b/tests/test_litellm/llms/bedrock/files/test_bedrock_files_transformation.py @@ -4,9 +4,7 @@ Test bedrock files transformation functionality import json import os -from typing import Any, Dict, List - -import pytest +from urllib.parse import unquote, urlparse from litellm.llms.bedrock.files.transformation import BedrockJsonlFilesTransformation @@ -43,19 +41,6 @@ class TestBedrockFilesTransformation: ) ) - # Print the transformation results for validation - print("\n=== INPUT (OpenAI format) ===") - for i, content in enumerate(openai_jsonl_content): - print(f"Record {i+1}:") - print(json.dumps(content, indent=2)) - print() - - print("\n=== OUTPUT (Bedrock format) ===") - for i, content in enumerate(bedrock_jsonl_content): - print(f"Record {i+1}:") - print(json.dumps(content, indent=2)) - print() - # Basic validation assert len(bedrock_jsonl_content) == len( openai_jsonl_content @@ -88,17 +73,6 @@ class TestBedrockFilesTransformation: "max_tokens" in model_input ), f"Record {i+1} should have max_tokens" - # Write expected output to file for reference - expected_output_path = os.path.join( - os.path.dirname(__file__), "expected_bedrock_batch_completions.jsonl" - ) - - with open(expected_output_path, "w") as f: - for record in bedrock_jsonl_content: - f.write(json.dumps(record) + "\n") - - print(f"\n=== Expected output written to: {expected_output_path} ===") - def test_nova_text_only_uses_converse_format(self): """ Test that Nova models produce Converse API format in batch modelInput. @@ -327,6 +301,44 @@ class TestBedrockFilesTransformation: ), f"us-west-2 must not appear when s3_region_name is set, got: {url}" assert "litellm-batch-352026" in url + def test_get_complete_file_url_sanitizes_untrusted_filename(self): + from litellm.llms.bedrock.files.transformation import BedrockFilesConfig + + config = BedrockFilesConfig() + create_file_data = { + "file": ("../../owned.jsonl?acl=public", b"hello", "application/jsonl"), + "purpose": "assistants", + } + + url = config.get_complete_file_url( + api_base=None, + api_key=None, + model="amazon.nova-pro-v1:0", + optional_params={"aws_region_name": "us-west-2"}, + litellm_params={"s3_bucket_name": "safe-bucket"}, + data=create_file_data, + ) + + parsed_url = urlparse(url) + object_key = unquote(parsed_url.path).split("/safe-bucket/", 1)[1] + assert object_key.startswith("litellm-bedrock-files/") + assert object_key.endswith("-owned.jsonl_acl_public") + assert ".." not in object_key + assert parsed_url.query == "" + + def test_batch_object_name_sanitizes_model_path(self): + from litellm.llms.bedrock.files.transformation import BedrockFilesConfig + + config = BedrockFilesConfig() + object_name = config._get_s3_object_name_from_batch_jsonl( + [{"body": {"model": "bedrock/../../secret:model"}}] + ) + + assert object_name.startswith("litellm-bedrock-files-") + assert object_name.endswith(".jsonl") + assert "/" not in object_name + assert ".." not in object_name + def test_transform_create_file_request_injects_s3_region_for_signing(self): """ When s3_region_name is provided, transform_create_file_request must pass diff --git a/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_handler.py b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_handler.py index ea056a2fc80..2b71e6c18dc 100644 --- a/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_handler.py +++ b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_handler.py @@ -12,6 +12,14 @@ from litellm.llms.vertex_ai.files.handler import VertexAIFilesHandler from litellm.types.llms.openai import FileContentRequest, HttpxBinaryResponseContent +def _mock_gcs_logging_config(bucket_name: str = "test-bucket"): + return { + "bucket_name": bucket_name, + "path_service_account": None, + "vertex_instance": None, + } + + class TestVertexAIFilesHandler: """Test Vertex AI files handler""" @@ -22,57 +30,69 @@ class TestVertexAIFilesHandler: def test_extract_bucket_and_object_from_file_id_standard_path(self): """Test extraction of bucket and object from URL-encoded file_id with standard path""" # Sample file_id with nested folder structure - file_id = "gs%3A%2F%2Ftest-bucket%2Ftest-folder" "%2Fsub-folder%2Ftest-file.txt" + file_id = ( + "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files" + "%2Ftest-folder%2Fsub-folder%2Ftest-file.txt" + ) - bucket_name, encoded_object_path = ( - self.handler._extract_bucket_and_object_from_file_id(file_id) + bucket_name, object_path = self.handler._extract_bucket_and_object_from_file_id( + file_id=file_id, + configured_bucket_name="test-bucket", ) # Verify bucket name extraction assert bucket_name == "test-bucket" - # Verify object path encoding - expected_encoded_object = "test-folder%2Fsub-folder%2Ftest-file.txt" - assert encoded_object_path == expected_encoded_object + expected_object = "litellm-vertex-files/test-folder/sub-folder/test-file.txt" + assert object_path == expected_object - def test_extract_bucket_and_object_from_file_id_bucket_only(self): + def test_extract_bucket_and_object_from_file_id_rejects_bucket_only(self): """Test extraction when only bucket name is provided""" file_id = "gs%3A%2F%2Ftest-bucket" - bucket_name, encoded_object_path = ( - self.handler._extract_bucket_and_object_from_file_id(file_id) - ) + with pytest.raises(ValueError, match="object name"): + self.handler._extract_bucket_and_object_from_file_id( + file_id=file_id, + configured_bucket_name="test-bucket", + ) - assert bucket_name == "test-bucket" - assert encoded_object_path == "" - - def test_extract_bucket_and_object_from_file_id_simple_path(self): + def test_extract_bucket_and_object_from_file_id_rejects_unmanaged_path(self): """Test extraction with simple path""" file_id = "gs%3A%2F%2Ftest-bucket%2Ftest-file.txt" - bucket_name, encoded_object_path = ( - self.handler._extract_bucket_and_object_from_file_id(file_id) - ) + with pytest.raises(ValueError, match="LiteLLM-managed"): + self.handler._extract_bucket_and_object_from_file_id( + file_id=file_id, + configured_bucket_name="test-bucket", + ) - assert bucket_name == "test-bucket" - assert encoded_object_path == "test-file.txt" - - def test_extract_bucket_and_object_from_file_id_no_gs_prefix(self): + def test_extract_bucket_and_object_from_file_id_rejects_no_gs_prefix(self): """Test extraction when gs:// prefix is missing""" - file_id = "test-bucket%2Ftest-file.txt" + file_id = "test-bucket%2Flitellm-vertex-files%2Ftest-file.txt" - bucket_name, encoded_object_path = ( - self.handler._extract_bucket_and_object_from_file_id(file_id) - ) + with pytest.raises(ValueError, match="gs://"): + self.handler._extract_bucket_and_object_from_file_id( + file_id=file_id, + configured_bucket_name="test-bucket", + ) - assert bucket_name == "test-bucket" - assert encoded_object_path == "test-file.txt" + def test_extract_bucket_and_object_from_file_id_rejects_wrong_bucket(self): + file_id = "gs%3A%2F%2Fother-bucket%2Flitellm-vertex-files%2Ftest-file.txt" + + with pytest.raises(ValueError, match="configured storage bucket"): + self.handler._extract_bucket_and_object_from_file_id( + file_id=file_id, + configured_bucket_name="test-bucket", + ) @pytest.mark.asyncio async def test_afile_content_success(self): """Test successful async file content retrieval""" # Setup test data - file_id = "gs%3A%2F%2Ftest-bucket%2Ftest-file.txt" + file_id = ( + "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files" + "%2Fuploads%2Fabc-test-file.txt" + ) expected_content = b"test file content" file_content_request = FileContentRequest( @@ -80,9 +100,17 @@ class TestVertexAIFilesHandler: ) # Mock the download_gcs_object method - with patch.object( - self.handler, "download_gcs_object", new_callable=AsyncMock - ) as mock_download: + with ( + patch.object( + self.handler, "download_gcs_object", new_callable=AsyncMock + ) as mock_download, + patch.object( + self.handler, + "get_gcs_logging_config", + new_callable=AsyncMock, + return_value=_mock_gcs_logging_config(), + ), + ): mock_download.return_value = expected_content # Call the method @@ -104,7 +132,10 @@ class TestVertexAIFilesHandler: # Verify the download was called with correct parameters mock_download.assert_called_once() call_args = mock_download.call_args - assert call_args.kwargs["object_name"] == "test-file.txt" + assert ( + call_args.kwargs["object_name"] + == "litellm-vertex-files/uploads/abc-test-file.txt" + ) assert "standard_callback_dynamic_params" in call_args.kwargs assert ( call_args.kwargs["standard_callback_dynamic_params"]["gcs_bucket_name"] @@ -132,22 +163,33 @@ class TestVertexAIFilesHandler: @pytest.mark.asyncio async def test_afile_content_download_failure(self): """Test async file content retrieval when download fails""" - file_id = "gs%3A%2F%2Ftest-bucket%2Ftest-file.txt" + file_id = ( + "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files" + "%2Fuploads%2Fabc-test-file.txt" + ) file_content_request = FileContentRequest( file_id=file_id, extra_headers=None, extra_body=None ) # Mock download to return None (failure) - with patch.object( - self.handler, "download_gcs_object", new_callable=AsyncMock - ) as mock_download: + with ( + patch.object( + self.handler, "download_gcs_object", new_callable=AsyncMock + ) as mock_download, + patch.object( + self.handler, + "get_gcs_logging_config", + new_callable=AsyncMock, + return_value=_mock_gcs_logging_config(), + ), + ): mock_download.return_value = None # Should raise ValueError for failed download with pytest.raises( ValueError, - match="Failed to download file from GCS: gs://test-bucket/test-file.txt", + match="Failed to download file from GCS: gs://test-bucket/litellm-vertex-files/uploads/abc-test-file.txt", ): await self.handler.afile_content( file_content_request=file_content_request, diff --git a/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_transformation.py b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_transformation.py index 596726cdb4b..4a7a8733ae9 100644 --- a/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_transformation.py @@ -2,8 +2,8 @@ Tests for VertexAIFilesConfig transformation methods (Issues 5-7). """ -import json import urllib.parse +from urllib.parse import parse_qs, urlparse import httpx import pytest @@ -23,13 +23,20 @@ class TestParseGcsUri: """Tests for the _parse_gcs_uri helper used by retrieve / content / delete.""" def test_should_parse_standard_gs_uri(self, config): - bucket, encoded = config._parse_gcs_uri("gs://my-bucket/path/to/object.jsonl") + file_id = "gs://my-bucket/litellm-vertex-files/path/to/object.jsonl" + bucket, encoded = config._parse_gcs_uri( + file_id, litellm_params={"bucket_name": "my-bucket"} + ) assert bucket == "my-bucket" - assert encoded == urllib.parse.quote("path/to/object.jsonl", safe="") + assert encoded == urllib.parse.quote( + "litellm-vertex-files/path/to/object.jsonl", safe="" + ) def test_should_parse_uri_with_nested_publisher_path(self, config): uri = "gs://litellm-local/litellm-vertex-files/publishers/google/models/gemini-2.0-flash-001/abc-123" - bucket, encoded = config._parse_gcs_uri(uri) + bucket, encoded = config._parse_gcs_uri( + uri, litellm_params={"bucket_name": "litellm-local"} + ) assert bucket == "litellm-local" expected_path = ( "litellm-vertex-files/publishers/google/models/gemini-2.0-flash-001/abc-123" @@ -37,30 +44,82 @@ class TestParseGcsUri: assert encoded == urllib.parse.quote(expected_path, safe="") def test_should_handle_url_encoded_input(self, config): - encoded_uri = urllib.parse.quote("gs://my-bucket/some/path", safe="") - bucket, encoded = config._parse_gcs_uri(encoded_uri) + encoded_uri = urllib.parse.quote( + "gs://my-bucket/litellm-vertex-files/some/path", safe="" + ) + bucket, encoded = config._parse_gcs_uri( + encoded_uri, litellm_params={"bucket_name": "my-bucket"} + ) assert bucket == "my-bucket" - assert encoded == urllib.parse.quote("some/path", safe="") + assert encoded == urllib.parse.quote("litellm-vertex-files/some/path", safe="") - def test_should_handle_bucket_only(self, config): - bucket, encoded = config._parse_gcs_uri("gs://my-bucket") - assert bucket == "my-bucket" - assert encoded == "" + def test_should_reject_bucket_only(self, config): + with pytest.raises(ValueError, match="object name"): + config._parse_gcs_uri( + "gs://my-bucket", litellm_params={"bucket_name": "my-bucket"} + ) - def test_should_handle_no_gs_prefix(self, config): - bucket, encoded = config._parse_gcs_uri("my-bucket/object.txt") - assert bucket == "my-bucket" - assert encoded == "object.txt" + def test_should_reject_no_gs_prefix(self, config): + with pytest.raises(ValueError, match="gs://"): + config._parse_gcs_uri( + "my-bucket/litellm-vertex-files/object.txt", + litellm_params={"bucket_name": "my-bucket"}, + ) + + def test_should_reject_unmanaged_object_path(self, config): + with pytest.raises(ValueError, match="LiteLLM-managed"): + config._parse_gcs_uri( + "gs://my-bucket/private/object.txt", + litellm_params={"bucket_name": "my-bucket"}, + ) + + def test_should_reject_unconfigured_bucket(self, config): + with pytest.raises(ValueError, match="configured storage bucket"): + config._parse_gcs_uri( + "gs://other-bucket/litellm-vertex-files/object.txt", + litellm_params={"bucket_name": "my-bucket"}, + ) + + +class TestCreateFileUrl: + def test_should_ignore_request_metadata_bucket_and_sanitize_filename(self, config): + url = config.get_complete_file_url( + api_base=None, + api_key=None, + model="", + optional_params={}, + litellm_params={ + "bucket_name": "safe-bucket", + "litellm_metadata": {"gcs_bucket_name": "attacker-bucket"}, + }, + data={ + "file": ("../../owned.jsonl?alt=media", b"{}", "application/jsonl"), + "purpose": "assistants", + }, + ) + + parsed_url = urlparse(url) + object_name = parse_qs(parsed_url.query)["name"][0] + assert "/b/safe-bucket/" in parsed_url.path + assert "attacker-bucket" not in url + assert object_name.startswith("litellm-vertex-files/uploads/") + assert object_name.endswith("-owned.jsonl_alt_media") + assert ".." not in object_name + assert "?" not in object_name class TestTransformRetrieveFile: def test_should_build_correct_gcs_metadata_url(self, config): - file_id = "gs://my-bucket/path/to/file.jsonl" + file_id = "gs://my-bucket/litellm-vertex-files/path/to/file.jsonl" url, params = config.transform_retrieve_file_request( - file_id=file_id, optional_params={}, litellm_params={} + file_id=file_id, + optional_params={}, + litellm_params={"bucket_name": "my-bucket"}, + ) + expected_encoded = urllib.parse.quote( + "litellm-vertex-files/path/to/file.jsonl", safe="" ) - expected_encoded = urllib.parse.quote("path/to/file.jsonl", safe="") assert ( url == f"https://storage.googleapis.com/storage/v1/b/my-bucket/o/{expected_encoded}" @@ -113,13 +172,13 @@ class TestTransformRetrieveFile: class TestTransformFileContent: def test_should_build_gcs_media_download_url(self, config): - file_id = "gs://my-bucket/path/to/file.jsonl" + file_id = "gs://my-bucket/litellm-vertex-files/path/to/file.jsonl" url, params = config.transform_file_content_request( file_content_request={"file_id": file_id}, optional_params={}, - litellm_params={}, + litellm_params={"bucket_name": "my-bucket"}, ) - encoded = urllib.parse.quote("path/to/file.jsonl", safe="") + encoded = urllib.parse.quote("litellm-vertex-files/path/to/file.jsonl", safe="") assert ( url == f"https://storage.googleapis.com/storage/v1/b/my-bucket/o/{encoded}?alt=media" @@ -146,11 +205,13 @@ class TestTransformFileContent: class TestTransformDeleteFile: def test_should_build_correct_gcs_delete_url(self, config): - file_id = "gs://my-bucket/path/to/file.jsonl" + file_id = "gs://my-bucket/litellm-vertex-files/path/to/file.jsonl" url, params = config.transform_delete_file_request( - file_id=file_id, optional_params={}, litellm_params={} + file_id=file_id, + optional_params={}, + litellm_params={"bucket_name": "my-bucket"}, ) - encoded = urllib.parse.quote("path/to/file.jsonl", safe="") + encoded = urllib.parse.quote("litellm-vertex-files/path/to/file.jsonl", safe="") assert ( url == f"https://storage.googleapis.com/storage/v1/b/my-bucket/o/{encoded}" )