load from s3

This commit is contained in:
Ishaan Jaff 2025-07-15 14:25:56 -07:00
parent 1ae72f6957
commit dd71655b2f
2 changed files with 344 additions and 0 deletions

View file

@ -1,4 +1,6 @@
import asyncio
import importlib
import importlib.util
import os
from typing import Any, Callable, Literal, Optional, get_type_hints
@ -7,6 +9,10 @@ def get_instance_fn(value: str, config_file_path: Optional[str] = None) -> Any:
module_name = value
instance_name = None
try:
# Check if value starts with s3:// or gcs://
if value.startswith("s3://") or value.startswith("gcs://"):
return _load_instance_from_remote_storage(value, config_file_path)
# Split the path by dots to separate module from instance
parts = value.split(".")
@ -20,12 +26,22 @@ def get_instance_fn(value: str, config_file_path: Optional[str] = None) -> Any:
module_file_path = os.path.join(directory, *module_name.split("."))
module_file_path += ".py"
# Check if the file exists before trying to load it
if not os.path.exists(module_file_path):
raise ImportError(
f"Could not find module file {module_file_path}"
)
spec = importlib.util.spec_from_file_location(module_name, module_file_path) # type: ignore
if spec is None:
raise ImportError(
f"Could not find a module specification for {module_file_path}"
)
module = importlib.util.module_from_spec(spec) # type: ignore
if spec.loader is None:
raise ImportError(
f"Could not find a module loader for {module_file_path}"
)
spec.loader.exec_module(module) # type: ignore
else:
# Dynamically import the module
@ -47,6 +63,122 @@ def get_instance_fn(value: str, config_file_path: Optional[str] = None) -> Any:
raise e
def _load_instance_from_remote_storage(remote_url: str, config_file_path: Optional[str] = None) -> Any:
"""
Load custom logger instance from S3 or GCS URL.
Expected format:
- s3://bucket-name/path/to/module.instance_name
- gcs://bucket-name/path/to/module.instance_name
Args:
remote_url (str): The s3:// or gcs:// URL
config_file_path (str): Optional config file path for temp directory context
Returns:
Any: The loaded instance
"""
try:
from litellm._logging import verbose_proxy_logger
# Parse the URL
if remote_url.startswith("s3://"):
storage_type = "s3"
url_without_prefix = remote_url[5:] # Remove 's3://'
elif remote_url.startswith("gcs://"):
storage_type = "gcs"
url_without_prefix = remote_url[6:] # Remove 'gcs://'
else:
raise ValueError(f"Unsupported URL scheme in {remote_url}")
# Split bucket and path
parts = url_without_prefix.split("/", 1)
if len(parts) < 2:
raise ValueError(f"Invalid URL format: {remote_url}. Expected: {storage_type}://bucket-name/path/to/module.instance")
bucket_name = parts[0]
path_and_module = parts[1]
# Extract module path and instance name
# Example: "loggers/custom_callbacks.proxy_handler_instance"
# Split by last dot to separate module from instance
module_parts = path_and_module.split(".")
if len(module_parts) < 2:
raise ValueError(f"Invalid module specification in {remote_url}. Expected: path/to/module.instance_name")
instance_name = module_parts[-1]
module_path = ".".join(module_parts[:-1])
# Create object key (file path in bucket)
object_key = f"{module_path}.py"
verbose_proxy_logger.debug(
f"Loading custom logger from {storage_type}: bucket={bucket_name}, "
f"object_key={object_key}, instance={instance_name}"
)
# Create temporary file for the downloaded module
temp_dir = "/tmp"
if config_file_path:
temp_dir = os.path.dirname(config_file_path)
# Create a unique filename to avoid conflicts
import uuid
temp_filename = f"remote_logger_{uuid.uuid4().hex[:8]}.py"
local_file_path = os.path.join(temp_dir, temp_filename)
# Download the file
if storage_type == "s3":
from litellm.proxy.common_utils.load_config_utils import (
download_python_file_from_s3,
)
success = download_python_file_from_s3(bucket_name, object_key, local_file_path)
else: # gcs
success = asyncio.run(_download_gcs_file_wrapper(bucket_name, object_key, local_file_path))
if not success:
raise ImportError(f"Failed to download {object_key} from {storage_type} bucket {bucket_name}")
# Load the module from the downloaded file
module_name = f"remote_logger_{uuid.uuid4().hex[:8]}"
spec = importlib.util.spec_from_file_location(module_name, local_file_path)
if spec is None or spec.loader is None:
raise ImportError(f"Could not create module spec for {local_file_path}")
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
# Get the instance
instance = getattr(module, instance_name)
# Clean up the temporary file
try:
os.remove(local_file_path)
except Exception as cleanup_error:
verbose_proxy_logger.warning(f"Could not clean up temporary file {local_file_path}: {cleanup_error}")
verbose_proxy_logger.info(f"Successfully loaded custom logger from {remote_url}")
return instance
except Exception as e:
raise ImportError(f"Failed to load custom logger from {remote_url}: {str(e)}") from e
async def _download_gcs_file_wrapper(bucket_name: str, object_key: str, local_file_path: str) -> bool:
"""Wrapper for GCS download to handle async properly"""
try:
from litellm.proxy.common_utils.load_config_utils import (
download_python_file_from_gcs,
)
return await download_python_file_from_gcs(bucket_name, object_key, local_file_path)
except Exception as e:
from litellm._logging import verbose_proxy_logger
verbose_proxy_logger.error(f"Error downloading from GCS: {str(e)}")
return False
def validate_custom_validate_return_type(
fn: Optional[Callable[..., Any]],
) -> Optional[Callable[..., Literal[True]]]:

View file

@ -0,0 +1,212 @@
import pytest
import os
import tempfile
import importlib.util
from unittest.mock import patch, MagicMock
from litellm.proxy.types_utils.utils import get_instance_fn, _load_instance_from_remote_storage
class TestCustomLoggerS3GCS:
"""Test custom logger loading from S3/GCS using URL prefixes"""
@pytest.fixture
def sample_custom_logger_content(self):
"""Sample custom logger file content"""
return '''
from litellm.integrations.custom_logger import CustomLogger
class TestCustomLogger(CustomLogger):
def __init__(self):
super().__init__()
self.initialized = True
def log_pre_api_call(self, model, messages, kwargs):
print(f"Pre-API call to {model}")
# Instance to be imported
test_logger_instance = TestCustomLogger()
'''
@pytest.fixture
def temp_config_dir(self):
"""Create a temporary directory for config files"""
with tempfile.TemporaryDirectory() as temp_dir:
yield temp_dir
def test_local_file_loading_still_works(self, temp_config_dir, sample_custom_logger_content):
"""Test that local file loading continues to work (no URL prefix)"""
# Create a local custom logger file
custom_logger_path = os.path.join(temp_config_dir, "test_custom_logger.py")
with open(custom_logger_path, 'w') as f:
f.write(sample_custom_logger_content)
# Create a dummy config file
config_path = os.path.join(temp_config_dir, "config.yaml")
with open(config_path, 'w') as f:
f.write("model_list: []")
# Test loading the custom logger (traditional way)
instance = get_instance_fn("test_custom_logger.test_logger_instance", config_path)
assert instance is not None
assert hasattr(instance, 'initialized')
assert instance.initialized is True
def test_s3_url_parsing(self):
"""Test S3 URL parsing"""
test_url = "s3://my-bucket/loggers/custom_callbacks.proxy_handler_instance"
# Mock the download function to avoid actual S3 calls
with patch('litellm.proxy.common_utils.load_config_utils.download_python_file_from_s3') as mock_download:
mock_download.return_value = False # Will cause failure, but we just want to test parsing
with pytest.raises(ImportError, match="Failed to download"):
_load_instance_from_remote_storage(test_url)
# Verify the download was called with correct parameters
mock_download.assert_called_once()
args = mock_download.call_args[0]
assert args[0] == "my-bucket" # bucket_name
assert args[1] == "loggers/custom_callbacks.py" # object_key
def test_gcs_url_parsing(self):
"""Test GCS URL parsing"""
test_url = "gcs://my-bucket/custom_logger.my_instance"
# Mock the download function
with patch('litellm.proxy.types_utils.utils._download_gcs_file_wrapper') as mock_download:
mock_download.return_value = False # Will cause failure
with pytest.raises(ImportError, match="Failed to download"):
_load_instance_from_remote_storage(test_url)
# Verify the download was called with correct parameters
mock_download.assert_called_once()
args = mock_download.call_args[0]
assert args[0] == "my-bucket" # bucket_name
assert args[1] == "custom_logger.py" # object_key
@patch('litellm.proxy.common_utils.load_config_utils.download_python_file_from_s3')
def test_s3_download_success(self, mock_s3_download, sample_custom_logger_content):
"""Test successful S3 download and loading"""
# Configure S3 download to succeed and create the file
def mock_download(bucket, key, local_path):
with open(local_path, 'w') as f:
f.write(sample_custom_logger_content)
return True
mock_s3_download.side_effect = mock_download
# Test loading with S3 URL
test_url = "s3://test-bucket/test_custom_logger.test_logger_instance"
instance = get_instance_fn(test_url)
assert instance is not None
assert hasattr(instance, 'initialized')
assert instance.initialized is True
# Verify S3 download was called with correct parameters
mock_s3_download.assert_called_once()
call_args = mock_s3_download.call_args
assert call_args[0][0] == 'test-bucket' # bucket_name
assert call_args[0][1] == 'test_custom_logger.py' # object_key
@patch('litellm.proxy.types_utils.utils._download_gcs_file_wrapper')
def test_gcs_download_success(self, mock_gcs_download, sample_custom_logger_content):
"""Test successful GCS download and loading"""
# Configure GCS download to succeed and create the file
def mock_download(bucket, key, local_path):
with open(local_path, 'w') as f:
f.write(sample_custom_logger_content)
return True
mock_gcs_download.side_effect = mock_download
# Test loading with GCS URL
test_url = "gcs://test-bucket/test_custom_logger.test_logger_instance"
instance = get_instance_fn(test_url)
assert instance is not None
assert hasattr(instance, 'initialized')
assert instance.initialized is True
def test_nested_path_parsing(self):
"""Test parsing of nested paths in URLs"""
test_url = "s3://my-bucket/loggers/production/advanced_logger.handler_instance"
with patch('litellm.proxy.common_utils.load_config_utils.download_python_file_from_s3') as mock_download:
mock_download.return_value = False
with pytest.raises(ImportError):
_load_instance_from_remote_storage(test_url)
# Verify correct object key was generated
call_args = mock_download.call_args
assert call_args[0][1] == "loggers/production/advanced_logger.py"
def test_invalid_url_schemes(self):
"""Test error handling for invalid URL schemes"""
# URLs that look like URLs but aren't s3:// or gcs:// will be treated as module names
# and fail with regular ImportError
with pytest.raises(ImportError):
get_instance_fn("http://bucket/module.instance")
with pytest.raises(ImportError):
get_instance_fn("ftp://bucket/module.instance")
def test_invalid_url_format(self):
"""Test error handling for invalid URL formats"""
# Missing bucket
with pytest.raises(ImportError, match="Invalid URL format"):
get_instance_fn("s3://")
# Missing path
with pytest.raises(ImportError, match="Invalid URL format"):
get_instance_fn("s3://bucket-only")
# Missing instance name
with pytest.raises(ImportError, match="Invalid module specification"):
get_instance_fn("s3://bucket/module-only")
@patch('litellm.proxy.common_utils.load_config_utils.download_python_file_from_s3')
def test_download_failure_handling(self, mock_s3_download):
"""Test handling of download failures"""
mock_s3_download.return_value = False
test_url = "s3://test-bucket/failing_logger.instance"
with pytest.raises(ImportError, match="Failed to download"):
get_instance_fn(test_url)
@patch('litellm.proxy.common_utils.load_config_utils.download_python_file_from_s3')
def test_file_cleanup(self, mock_s3_download, sample_custom_logger_content):
"""Test that temporary files are cleaned up"""
created_files = []
def mock_download(bucket, key, local_path):
created_files.append(local_path)
with open(local_path, 'w') as f:
f.write(sample_custom_logger_content)
return True
mock_s3_download.side_effect = mock_download
test_url = "s3://test-bucket/test_custom_logger.test_logger_instance"
instance = get_instance_fn(test_url)
assert instance is not None
# Verify file was created and then cleaned up
assert len(created_files) == 1
temp_file = created_files[0]
assert not os.path.exists(temp_file), f"Temporary file {temp_file} was not cleaned up"
def test_no_url_prefix_fallback(self, temp_config_dir):
"""Test fallback when no URL prefix is used and local file doesn't exist"""
config_path = os.path.join(temp_config_dir, "config.yaml")
with open(config_path, 'w') as f:
f.write("model_list: []")
# Test that it tries local loading when no URL prefix is used
with pytest.raises(ImportError, match="Could not import instance from nonexistent_logger"):
get_instance_fn("nonexistent_logger.instance", config_path)