mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(s3_v2): URL-encode object keys to prevent SigV4 signature mismatch
S3 object keys containing special characters like '=' (from base64 padding in Responses API composite IDs) caused 403 Invalid signature errors on S3-compatible endpoints (Garage, MinIO). botocore's SigV4Auth URL-encodes the path in the canonical request (turning '=' into '%3D'), but httpx sends the literal '=' on the wire. The server computes its signature using the unencoded path, causing a mismatch. Fix: apply urllib.parse.quote() to the object key before embedding it in the URL, so both the SigV4 signature and the HTTP request agree on the path encoding.
This commit is contained in:
parent
7d7045cbc1
commit
7d42a15965
2 changed files with 138 additions and 9 deletions
|
|
@ -9,6 +9,7 @@ NOTE 1: S3 does not provide a BATCH PUT API endpoint, so we create tasks to uplo
|
|||
import asyncio
|
||||
from datetime import datetime
|
||||
from typing import List, Optional, cast
|
||||
from urllib.parse import quote
|
||||
|
||||
import litellm
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
|
|
@ -346,8 +347,16 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
)
|
||||
verbose_logger.debug(f"s3_v2 logger - s3_verify setting: {self.s3_verify}")
|
||||
|
||||
# URL-encode the object key so that characters like '=' (from
|
||||
# base64 padding) are percent-encoded. This ensures the URL that
|
||||
# httpx sends on the wire matches the canonical path that
|
||||
# SigV4Auth uses when computing the request signature.
|
||||
encoded_key = quote(
|
||||
batch_logging_element.s3_object_key, safe="/"
|
||||
)
|
||||
|
||||
# Prepare the URL
|
||||
url = f"https://{self.s3_bucket_name}.s3.{self.s3_region_name}.amazonaws.com/{batch_logging_element.s3_object_key}"
|
||||
url = f"https://{self.s3_bucket_name}.s3.{self.s3_region_name}.amazonaws.com/{encoded_key}"
|
||||
|
||||
if self.s3_endpoint_url and self.s3_bucket_name:
|
||||
if self.s3_use_virtual_hosted_style:
|
||||
|
|
@ -360,7 +369,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
if self.s3_endpoint_url.startswith("https://")
|
||||
else "http://"
|
||||
)
|
||||
url = f"{protocol}{self.s3_bucket_name}.{endpoint_host}/{batch_logging_element.s3_object_key}"
|
||||
url = f"{protocol}{self.s3_bucket_name}.{endpoint_host}/{encoded_key}"
|
||||
else:
|
||||
# Path-style: endpoint/bucket/key
|
||||
url = (
|
||||
|
|
@ -368,7 +377,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
+ "/"
|
||||
+ self.s3_bucket_name
|
||||
+ "/"
|
||||
+ batch_logging_element.s3_object_key
|
||||
+ encoded_key
|
||||
)
|
||||
|
||||
# Convert JSON to string
|
||||
|
|
@ -520,8 +529,12 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
aws_region_name=self.s3_region_name,
|
||||
)
|
||||
|
||||
encoded_key = quote(
|
||||
batch_logging_element.s3_object_key, safe="/"
|
||||
)
|
||||
|
||||
# Prepare the URL
|
||||
url = f"https://{self.s3_bucket_name}.s3.{self.s3_region_name}.amazonaws.com/{batch_logging_element.s3_object_key}"
|
||||
url = f"https://{self.s3_bucket_name}.s3.{self.s3_region_name}.amazonaws.com/{encoded_key}"
|
||||
|
||||
if self.s3_endpoint_url and self.s3_bucket_name:
|
||||
if self.s3_use_virtual_hosted_style:
|
||||
|
|
@ -534,7 +547,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
if self.s3_endpoint_url.startswith("https://")
|
||||
else "http://"
|
||||
)
|
||||
url = f"{protocol}{self.s3_bucket_name}.{endpoint_host}/{batch_logging_element.s3_object_key}"
|
||||
url = f"{protocol}{self.s3_bucket_name}.{endpoint_host}/{encoded_key}"
|
||||
else:
|
||||
# Path-style: endpoint/bucket/key
|
||||
url = (
|
||||
|
|
@ -542,7 +555,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
+ "/"
|
||||
+ self.s3_bucket_name
|
||||
+ "/"
|
||||
+ batch_logging_element.s3_object_key
|
||||
+ encoded_key
|
||||
)
|
||||
|
||||
# Convert JSON to string
|
||||
|
|
@ -629,8 +642,10 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
f"s3_v2 logger - downloading data from s3 - {s3_object_key}"
|
||||
)
|
||||
|
||||
encoded_key = quote(s3_object_key, safe="/")
|
||||
|
||||
# Prepare the URL
|
||||
url = f"https://{self.s3_bucket_name}.s3.{self.s3_region_name}.amazonaws.com/{s3_object_key}"
|
||||
url = f"https://{self.s3_bucket_name}.s3.{self.s3_region_name}.amazonaws.com/{encoded_key}"
|
||||
|
||||
if self.s3_endpoint_url and self.s3_bucket_name:
|
||||
if self.s3_use_virtual_hosted_style:
|
||||
|
|
@ -643,7 +658,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
if self.s3_endpoint_url.startswith("https://")
|
||||
else "http://"
|
||||
)
|
||||
url = f"{protocol}{self.s3_bucket_name}.{endpoint_host}/{s3_object_key}"
|
||||
url = f"{protocol}{self.s3_bucket_name}.{endpoint_host}/{encoded_key}"
|
||||
else:
|
||||
# Path-style: endpoint/bucket/key
|
||||
url = (
|
||||
|
|
@ -651,7 +666,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
+ "/"
|
||||
+ self.s3_bucket_name
|
||||
+ "/"
|
||||
+ s3_object_key
|
||||
+ encoded_key
|
||||
)
|
||||
|
||||
# Prepare the request for GET operation
|
||||
|
|
|
|||
|
|
@ -25,6 +25,120 @@ class TestS3V2UnitTests:
|
|||
"json.dumps(" not in source_code
|
||||
), "S3 v2 should not use json.dumps directly"
|
||||
|
||||
@patch('asyncio.create_task')
|
||||
@patch('litellm.integrations.s3_v2.CustomBatchLogger.periodic_flush')
|
||||
def test_s3_v2_url_encodes_object_key_with_special_chars(self, mock_periodic_flush, mock_create_task):
|
||||
"""Test that S3 object keys with special characters (e.g. base64 padding '=')
|
||||
are URL-encoded in the request URL.
|
||||
|
||||
Without URL-encoding, SigV4Auth computes the signature over the URL-encoded
|
||||
canonical path (%3D) while httpx sends the literal '=' on the wire. S3-compatible
|
||||
servers (Garage, MinIO) then reject the request with 403 Invalid signature.
|
||||
"""
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from litellm.types.integrations.s3_v2 import s3BatchLoggingElement
|
||||
|
||||
mock_periodic_flush.return_value = None
|
||||
mock_create_task.return_value = None
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
# Object key with base64 padding (== at the end), as produced by
|
||||
# LiteLLM's Responses API composite response IDs
|
||||
test_element = s3BatchLoggingElement(
|
||||
s3_object_key="2025-09-14/time-18-07-17_resp_bGl0ZWxsbTpPUT==.json",
|
||||
payload={"test": "data"},
|
||||
s3_object_download_filename="test-file.json"
|
||||
)
|
||||
|
||||
# Test: path-style with custom endpoint
|
||||
s3_logger = S3Logger(
|
||||
s3_bucket_name="test-bucket",
|
||||
s3_endpoint_url="https://garage.example.com",
|
||||
s3_aws_access_key_id="test-key",
|
||||
s3_aws_secret_access_key="test-secret",
|
||||
s3_region_name="us-east-1"
|
||||
)
|
||||
|
||||
s3_logger.async_httpx_client = AsyncMock()
|
||||
s3_logger.async_httpx_client.put.return_value = mock_response
|
||||
|
||||
asyncio.run(s3_logger.async_upload_data_to_s3(test_element))
|
||||
|
||||
call_args = s3_logger.async_httpx_client.put.call_args
|
||||
assert call_args is not None
|
||||
url = call_args[0][0]
|
||||
# The '=' characters must be percent-encoded as %3D
|
||||
assert "%3D" in url, f"Expected URL-encoded '=' (%3D) in URL, got {url}"
|
||||
assert "==" not in url, f"Literal '==' should not appear in URL path, got {url}"
|
||||
|
||||
# Test: virtual-hosted-style
|
||||
s3_logger_virtual = S3Logger(
|
||||
s3_bucket_name="test-bucket",
|
||||
s3_endpoint_url="https://garage.example.com",
|
||||
s3_aws_access_key_id="test-key",
|
||||
s3_aws_secret_access_key="test-secret",
|
||||
s3_region_name="us-east-1",
|
||||
s3_use_virtual_hosted_style=True,
|
||||
)
|
||||
|
||||
s3_logger_virtual.async_httpx_client = AsyncMock()
|
||||
s3_logger_virtual.async_httpx_client.put.return_value = mock_response
|
||||
|
||||
asyncio.run(s3_logger_virtual.async_upload_data_to_s3(test_element))
|
||||
|
||||
call_args_virtual = s3_logger_virtual.async_httpx_client.put.call_args
|
||||
url_virtual = call_args_virtual[0][0]
|
||||
assert "%3D" in url_virtual, f"Expected URL-encoded '=' in virtual-hosted URL, got {url_virtual}"
|
||||
assert "==" not in url_virtual, f"Literal '==' should not appear in virtual-hosted URL, got {url_virtual}"
|
||||
|
||||
# Test: sync upload method
|
||||
s3_logger_sync = S3Logger(
|
||||
s3_bucket_name="test-bucket",
|
||||
s3_endpoint_url="https://garage.example.com",
|
||||
s3_aws_access_key_id="test-key",
|
||||
s3_aws_secret_access_key="test-secret",
|
||||
s3_region_name="us-east-1"
|
||||
)
|
||||
|
||||
mock_sync_client = MagicMock()
|
||||
mock_sync_client.put.return_value = mock_response
|
||||
|
||||
with patch('litellm.integrations.s3_v2._get_httpx_client', return_value=mock_sync_client):
|
||||
s3_logger_sync.upload_data_to_s3(test_element)
|
||||
|
||||
call_args_sync = mock_sync_client.put.call_args
|
||||
url_sync = call_args_sync[0][0]
|
||||
assert "%3D" in url_sync, f"Expected URL-encoded '=' in sync URL, got {url_sync}"
|
||||
assert "==" not in url_sync, f"Literal '==' should not appear in sync URL, got {url_sync}"
|
||||
|
||||
# Test: download method
|
||||
s3_logger_download = S3Logger(
|
||||
s3_bucket_name="test-bucket",
|
||||
s3_endpoint_url="https://garage.example.com",
|
||||
s3_aws_access_key_id="test-key",
|
||||
s3_aws_secret_access_key="test-secret",
|
||||
s3_region_name="us-east-1"
|
||||
)
|
||||
|
||||
mock_download_response = MagicMock()
|
||||
mock_download_response.status_code = 200
|
||||
mock_download_response.json = MagicMock(return_value={"downloaded": "data"})
|
||||
s3_logger_download.async_httpx_client = AsyncMock()
|
||||
s3_logger_download.async_httpx_client.get.return_value = mock_download_response
|
||||
|
||||
asyncio.run(s3_logger_download._download_object_from_s3(
|
||||
"2025-09-14/time-18-07-17_resp_bGl0ZWxsbTpPUT==.json"
|
||||
))
|
||||
|
||||
call_args_dl = s3_logger_download.async_httpx_client.get.call_args
|
||||
url_dl = call_args_dl[0][0]
|
||||
assert "%3D" in url_dl, f"Expected URL-encoded '=' in download URL, got {url_dl}"
|
||||
assert "==" not in url_dl, f"Literal '==' should not appear in download URL, got {url_dl}"
|
||||
|
||||
@patch('asyncio.create_task')
|
||||
@patch('litellm.integrations.s3_v2.CustomBatchLogger.periodic_flush')
|
||||
def test_s3_v2_endpoint_url(self, mock_periodic_flush, mock_create_task):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue