Merge pull request #35148 from BerriAI/litellm_bedrock_batch_sse_kms

fix(bedrock): pass SSE-KMS key through to the batch input-file S3 upload
This commit is contained in:
Mateo Wang 2026-08-06 20:52:55 -07:00 committed by GitHub
commit e46721a36c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 119 additions and 14 deletions

View file

@ -12,7 +12,6 @@ from litellm.litellm_core_utils.cloud_storage_security import (
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.llms.base_llm.batches.transformation import BaseBatchesConfig
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.bedrock import (
BedrockCreateBatchRequest,
BedrockCreateBatchResponse,
@ -29,7 +28,7 @@ from litellm.types.llms.openai import (
from litellm.types.utils import LiteLLMBatch, LlmProviders
from ..base_aws_llm import BaseAWSLLM
from ..common_utils import CommonBatchFilesUtils
from ..common_utils import CommonBatchFilesUtils, resolve_s3_encryption_key_id
# Bedrock batch input files are uploaded as
# s3://bucket/litellm-bedrock-files-{model, ":" -> "-"}-{uuid4}.jsonl (see
@ -200,7 +199,10 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
)
# Add optional KMS encryption key ID if provided
s3_encryption_key_id = litellm_params.get("s3_encryption_key_id") or get_secret_str("AWS_S3_ENCRYPTION_KEY_ID")
s3_encryption_key_id = resolve_s3_encryption_key_id(
litellm_params=litellm_params,
optional_params=optional_params,
)
if s3_encryption_key_id:
s3_output_config["s3EncryptionKeyId"] = s3_encryption_key_id

View file

@ -26,7 +26,7 @@ from litellm.llms.base_llm.anthropic_messages.transformation import (
)
from litellm.llms.base_llm.base_utils import BaseLLMModelInfo, BaseTokenCounter
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.secret_managers.main import get_secret
from litellm.secret_managers.main import get_secret, get_secret_str
if TYPE_CHECKING:
from litellm.types.llms.openai import AllMessageValues
@ -1304,6 +1304,23 @@ def get_anthropic_beta_from_headers(headers: dict) -> list[str]:
return []
def resolve_s3_encryption_key_id(
litellm_params: Mapping[str, Any],
optional_params: Mapping[str, Any] | None = None,
) -> str | None:
"""
Resolve the SSE-KMS key configured for Bedrock batch/file S3 objects.
Precedence: `s3_encryption_key_id` in litellm_params, then optional_params
(client-side / request params), then the AWS_S3_ENCRYPTION_KEY_ID env var.
"""
candidates: Final = tuple(
source.get("s3_encryption_key_id") for source in (litellm_params, optional_params) if source is not None
)
explicit: Final = next((value for value in candidates if isinstance(value, str) and value), None)
return explicit or get_secret_str("AWS_S3_ENCRYPTION_KEY_ID")
class CommonBatchFilesUtils:
"""
Common utilities for Bedrock batch and file operations.

View file

@ -46,7 +46,7 @@ from litellm.types.utils import ExtractedFileData, LlmProviders, SpecialEnums
from litellm.utils import get_llm_provider
from ..base_aws_llm import BaseAWSLLM
from ..common_utils import BedrockError
from ..common_utils import BedrockError, resolve_s3_encryption_key_id
# litellm_params key used to hand the SigV4-signed GET headers from
# `transform_file_content_request` to `validate_environment` (the only hook
@ -733,6 +733,10 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
content=file_content,
api_base=api_base,
optional_params=optional_params,
s3_encryption_key_id=resolve_s3_encryption_key_id(
litellm_params=litellm_params,
optional_params=optional_params,
),
)
litellm_params["upload_url"] = api_base
@ -750,6 +754,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
content: str,
api_base: str,
optional_params: dict,
s3_encryption_key_id: str | None = None,
) -> tuple[dict, str]:
"""
Sign S3 PUT request using the same proven logic as S3Logger.
@ -782,12 +787,25 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
content_hash: Final = hashlib.sha256(content.encode("utf-8")).hexdigest()
# Prepare headers with required S3 headers (same as s3_v2.py)
request_headers: Final = {
"Content-Type": "application/json", # JSONL files are JSON content
"x-amz-content-sha256": content_hash, # REQUIRED by S3
"Content-Language": "en",
"Cache-Control": "private, immutable, max-age=31536000, s-maxage=0",
}
sse_headers: Final = (
MappingProxyType(
{
"x-amz-server-side-encryption": "aws:kms",
"x-amz-server-side-encryption-aws-kms-key-id": s3_encryption_key_id,
}
)
if s3_encryption_key_id
else MappingProxyType({})
)
request_headers: Final = MappingProxyType(
{
"Content-Type": "application/json", # JSONL files are JSON content
"x-amz-content-sha256": content_hash, # REQUIRED by S3
"Content-Language": "en",
"Cache-Control": "private, immutable, max-age=31536000, s-maxage=0",
**sse_headers,
}
)
# Use requests.Request to prepare the request (same pattern as s3_v2.py)
req: Final = requests.Request("PUT", api_base, data=content, headers=request_headers)

View file

@ -172,7 +172,7 @@ def test_create_request_omits_kms_key_when_absent(config):
"generate_unique_job_name",
return_value="litellm-batch-1",
), patch.object(config.common_utils, "sign_aws_request") as mock_sign, patch(
"litellm.llms.bedrock.batches.transformation.get_secret_str",
"litellm.llms.bedrock.common_utils.get_secret_str",
return_value=None,
):
mock_sign.return_value = ({}, b"{}")

View file

@ -446,7 +446,7 @@ class TestBedrockFilesTransformation:
captured_optional_params: dict = {}
def fake_sign(content, api_base, optional_params):
def fake_sign(content, api_base, optional_params, s3_encryption_key_id=None):
captured_optional_params.update(optional_params)
return {"Authorization": "fake"}, content
@ -502,7 +502,7 @@ class TestBedrockFilesTransformation:
captured_optional_params: dict = {}
def fake_sign(content, api_base, optional_params):
def fake_sign(content, api_base, optional_params, s3_encryption_key_id=None):
captured_optional_params.update(optional_params)
return {"Authorization": "fake"}, content
@ -518,6 +518,74 @@ class TestBedrockFilesTransformation:
captured_optional_params.get("aws_region_name") == "us-gov-west-1"
), "s3_region_name must override aws_region_name for SigV4 signing"
def _signed_upload_request(self, litellm_params: dict) -> dict:
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
config = BedrockFilesConfig()
jsonl_content = json.dumps(
{
"custom_id": "req-1",
"method": "POST",
"url": "/v1/chat/completions",
"body": {
"model": "bedrock/amazon.nova-pro-v1:0",
"messages": [{"role": "user", "content": "Hello"}],
"max_tokens": 10,
},
}
).encode()
request = config.transform_create_file_request(
model="amazon.nova-pro-v1:0",
create_file_data={
"file": ("batch.jsonl", jsonl_content, "application/jsonl"),
"purpose": "batch",
},
optional_params={
"aws_access_key_id": "test-key-id",
"aws_secret_access_key": "test-secret",
"aws_region_name": "us-west-2",
},
litellm_params={"s3_bucket_name": "litellm-batch-bucket", **litellm_params},
)
assert isinstance(request, dict)
return request
def test_upload_signs_sse_kms_headers_when_key_configured(self, monkeypatch):
"""
Buckets whose policy requires SSE-KMS reject the batch input-file PutObject
unless the upload carries the aws:kms encryption headers; they must also be
covered by SigV4 SignedHeaders or S3 answers SignatureDoesNotMatch.
"""
monkeypatch.delenv("AWS_S3_ENCRYPTION_KEY_ID", raising=False)
kms_key = "arn:aws:kms:us-west-2:1234:key/abcd"
request = self._signed_upload_request({"s3_encryption_key_id": kms_key})
headers = {key.lower(): value for key, value in request["headers"].items()}
assert headers["x-amz-server-side-encryption"] == "aws:kms"
assert headers["x-amz-server-side-encryption-aws-kms-key-id"] == kms_key
signed_headers = headers["authorization"].split("SignedHeaders=")[1].split(",")[0]
assert "x-amz-server-side-encryption" in signed_headers
assert "x-amz-server-side-encryption-aws-kms-key-id" in signed_headers
def test_upload_reads_sse_kms_key_from_env(self, monkeypatch):
monkeypatch.setenv("AWS_S3_ENCRYPTION_KEY_ID", "env-kms-key")
request = self._signed_upload_request({})
headers = {key.lower(): value for key, value in request["headers"].items()}
assert headers["x-amz-server-side-encryption-aws-kms-key-id"] == "env-kms-key"
def test_upload_omits_sse_headers_when_no_key_configured(self, monkeypatch):
monkeypatch.delenv("AWS_S3_ENCRYPTION_KEY_ID", raising=False)
request = self._signed_upload_request({})
headers = {key.lower() for key in request["headers"]}
assert "x-amz-server-side-encryption" not in headers
assert "x-amz-server-side-encryption-aws-kms-key-id" not in headers
def test_openai_passthrough_still_works(self):
"""
Regression test: ensure OpenAI-compatible models (e.g. gpt-oss)