add _load_instance_from_remote_storage

This commit is contained in:
Ishaan Jaff 2025-07-15 15:40:28 -07:00
parent a078e43d44
commit 77321e76bb

View file

@ -101,6 +101,16 @@ def _load_instance_from_remote_storage(remote_url: str, config_file_path: Option
# Extract module path and instance name
# Example: "loggers/custom_callbacks.proxy_handler_instance"
# Handle case where user accidentally includes .py extension
if path_and_module.endswith('.py'):
module_name_without_py = path_and_module[:-3] # Remove .py
raise ValueError(
f"Invalid URL format in {remote_url}. "
f"Don't include '.py' extension and you must specify the instance name. "
f"Expected format: {storage_type}://{bucket_name}/{module_name_without_py}.instance_name "
f"(e.g., {storage_type}://{bucket_name}/{module_name_without_py}.proxy_handler_instance)"
)
# Split by last dot to separate module from instance
module_parts = path_and_module.split(".")
if len(module_parts) < 2:
@ -116,32 +126,32 @@ def _load_instance_from_remote_storage(remote_url: str, config_file_path: Option
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)
import tempfile
# Create temporary file for the downloaded module using the actual module name
temp_file = tempfile.NamedTemporaryFile(suffix='.py', delete=False)
local_file_path = temp_file.name
temp_file.close() # Close the file so we can write to it
# 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)
success = download_python_file_from_s3(
bucket_name=bucket_name,
object_key=object_key,
local_file_path=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)
# Load the module from the downloaded file using the actual module name
spec = importlib.util.spec_from_file_location(module_path, local_file_path)
if spec is None or spec.loader is None:
raise ImportError(f"Could not create module spec for {local_file_path}")