Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 5 additions & 3 deletions litellm/llms/bedrock/batches/transformation.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,6 @@
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,
Expand All @@ -29,7 +28,7 @@
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
Expand Down Expand Up @@ -200,7 +199,10 @@ def transform_create_batch_request(
)

# 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

Expand Down
19 changes: 18 additions & 1 deletion litellm/llms/bedrock/common_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,7 @@
)
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
Expand Down Expand Up @@ -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.
Expand Down
32 changes: 25 additions & 7 deletions litellm/llms/bedrock/files/transformation.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,7 +46,7 @@
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
Expand Down Expand Up @@ -733,6 +733,10 @@ def transform_create_file_request(
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
Expand All @@ -750,6 +754,7 @@ def _sign_s3_request(
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.
Expand Down Expand Up @@ -782,12 +787,25 @@ def _sign_s3_request(
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)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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"{}")
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -446,7 +446,7 @@ def test_transform_create_file_request_injects_s3_region_for_signing(self):

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

Expand Down Expand Up @@ -502,7 +502,7 @@ def test_s3_region_name_wins_over_aws_region_name_for_signing(self):

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

Expand All @@ -518,6 +518,74 @@ def fake_sign(content, api_base, optional_params):
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)
Expand Down
Loading