mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
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:
commit
e46721a36c
5 changed files with 119 additions and 14 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"{}")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue