fix cloud storage file guards

This commit is contained in:
user 2026-05-01 15:34:07 -07:00
parent d07cdd4481
commit ba1188117d
14 changed files with 658 additions and 201 deletions

View file

@ -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)

View file

@ -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,
)

View file

@ -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

View file

@ -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(

View file

@ -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(

View file

@ -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:

View file

@ -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)

View file

@ -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(

View file

@ -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

View file

@ -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():

View file

@ -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",
)

View file

@ -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

View file

@ -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,

View file

@ -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}"
)