mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
[Bug Fix] s3 config.yaml file - ensure yaml safe load is used (#12373)
* use yaml safe load * test_get_file_contents_from_s3_no_temp_file_creation
This commit is contained in:
parent
32f3887a01
commit
0a36f89009
2 changed files with 95 additions and 14 deletions
|
|
@ -6,8 +6,6 @@ from litellm._logging import verbose_proxy_logger
|
|||
def get_file_contents_from_s3(bucket_name, object_key):
|
||||
try:
|
||||
# v0 rely on boto3 for authentication - allowing boto3 to handle IAM credentials etc
|
||||
import tempfile
|
||||
|
||||
import boto3
|
||||
from botocore.credentials import Credentials
|
||||
|
||||
|
|
@ -26,21 +24,14 @@ def get_file_contents_from_s3(bucket_name, object_key):
|
|||
response = s3_client.get_object(Bucket=bucket_name, Key=object_key)
|
||||
verbose_proxy_logger.debug(f"Response: {response}")
|
||||
|
||||
# Read the file contents
|
||||
# Read the file contents and directly parse YAML
|
||||
file_contents = response["Body"].read().decode("utf-8")
|
||||
verbose_proxy_logger.debug("File contents retrieved from S3")
|
||||
|
||||
# Create a temporary file with YAML extension
|
||||
with tempfile.NamedTemporaryFile(delete=False, suffix=".yaml") as temp_file:
|
||||
temp_file.write(file_contents.encode("utf-8"))
|
||||
temp_file_path = temp_file.name
|
||||
verbose_proxy_logger.debug(f"File stored temporarily at: {temp_file_path}")
|
||||
|
||||
# Load the YAML file content
|
||||
with open(temp_file_path, "r") as yaml_file:
|
||||
config = yaml.safe_load(yaml_file)
|
||||
|
||||
|
||||
# Parse YAML directly from string
|
||||
config = yaml.safe_load(file_contents)
|
||||
return config
|
||||
|
||||
except ImportError as e:
|
||||
# this is most likely if a user is not using the litellm docker container
|
||||
verbose_proxy_logger.error(f"ImportError: {str(e)}")
|
||||
|
|
|
|||
|
|
@ -0,0 +1,90 @@
|
|||
from unittest.mock import MagicMock, mock_open, patch
|
||||
|
||||
import pytest
|
||||
import yaml
|
||||
|
||||
from litellm.proxy.common_utils.load_config_utils import get_file_contents_from_s3
|
||||
|
||||
|
||||
class TestGetFileContentsFromS3:
|
||||
"""Test suite for S3 config loading functionality."""
|
||||
|
||||
@patch('boto3.client')
|
||||
@patch('litellm.main.bedrock_converse_chat_completion')
|
||||
@patch('yaml.safe_load')
|
||||
def test_get_file_contents_from_s3_no_temp_file_creation(
|
||||
self, mock_yaml_load, mock_bedrock, mock_boto3_client
|
||||
):
|
||||
"""
|
||||
Test that get_file_contents_from_s3 doesn't create temporary files
|
||||
and uses yaml.safe_load directly on the S3 response content.
|
||||
|
||||
Note: It's critical that yaml.safe_load is used
|
||||
|
||||
Relevant issue/PR: https://github.com/BerriAI/litellm/pull/12078
|
||||
"""
|
||||
# Mock credentials
|
||||
mock_credentials = MagicMock()
|
||||
mock_credentials.access_key = "test_access_key"
|
||||
mock_credentials.secret_key = "test_secret_key"
|
||||
mock_credentials.token = "test_token"
|
||||
mock_bedrock.get_credentials.return_value = mock_credentials
|
||||
|
||||
# Mock S3 client and response
|
||||
mock_s3_client = MagicMock()
|
||||
mock_boto3_client.return_value = mock_s3_client
|
||||
|
||||
# Mock S3 response with YAML content
|
||||
yaml_content = """
|
||||
model_list:
|
||||
- model_name: gpt-3.5-turbo
|
||||
litellm_params:
|
||||
model: gpt-3.5-turbo
|
||||
"""
|
||||
mock_response_body = MagicMock()
|
||||
mock_response_body.read.return_value = yaml_content.encode('utf-8')
|
||||
mock_s3_response = {
|
||||
'Body': mock_response_body
|
||||
}
|
||||
mock_s3_client.get_object.return_value = mock_s3_response
|
||||
|
||||
# Mock yaml.safe_load to return parsed config
|
||||
expected_config = {
|
||||
'model_list': [{
|
||||
'model_name': 'gpt-3.5-turbo',
|
||||
'litellm_params': {
|
||||
'model': 'gpt-3.5-turbo'
|
||||
}
|
||||
}]
|
||||
}
|
||||
mock_yaml_load.return_value = expected_config
|
||||
|
||||
# Call the function
|
||||
bucket_name = "test-bucket"
|
||||
object_key = "config.yaml"
|
||||
result = get_file_contents_from_s3(bucket_name, object_key)
|
||||
|
||||
# Assertions
|
||||
assert result == expected_config
|
||||
|
||||
# Verify S3 client was created with correct credentials
|
||||
mock_boto3_client.assert_called_once_with(
|
||||
"s3",
|
||||
aws_access_key_id="test_access_key",
|
||||
aws_secret_access_key="test_secret_key",
|
||||
aws_session_token="test_token"
|
||||
)
|
||||
|
||||
# Verify S3 get_object was called with correct parameters
|
||||
mock_s3_client.get_object.assert_called_once_with(
|
||||
Bucket=bucket_name,
|
||||
Key=object_key
|
||||
)
|
||||
|
||||
# Verify the response body was read and decoded
|
||||
mock_response_body.read.assert_called_once()
|
||||
|
||||
# Verify yaml.safe_load was called with the decoded content
|
||||
mock_yaml_load.assert_called_once_with(yaml_content)
|
||||
|
||||
|
||||
Loading…
Add table
Reference in a new issue