diff --git a/litellm/proxy/types_utils/utils.py b/litellm/proxy/types_utils/utils.py index e159da49549..643fd1b7f4b 100644 --- a/litellm/proxy/types_utils/utils.py +++ b/litellm/proxy/types_utils/utils.py @@ -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]]]: diff --git a/tests/proxy_unit_tests/test_custom_logger_s3_gcs.py b/tests/proxy_unit_tests/test_custom_logger_s3_gcs.py new file mode 100644 index 00000000000..beefc929e39 --- /dev/null +++ b/tests/proxy_unit_tests/test_custom_logger_s3_gcs.py @@ -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) \ No newline at end of file