mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix cloud storage file guards
This commit is contained in:
parent
d07cdd4481
commit
ba1188117d
14 changed files with 658 additions and 201 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
140
litellm/litellm_core_utils/cloud_storage_security.py
Normal file
140
litellm/litellm_core_utils/cloud_storage_security.py
Normal 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
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue