mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
load from s3
This commit is contained in:
parent
1ae72f6957
commit
dd71655b2f
2 changed files with 344 additions and 0 deletions
|
|
@ -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]]]:
|
||||
|
|
|
|||
212
tests/proxy_unit_tests/test_custom_logger_s3_gcs.py
Normal file
212
tests/proxy_unit_tests/test_custom_logger_s3_gcs.py
Normal 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)
|
||||
Loading…
Add table
Reference in a new issue