mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-21 00:21:49 +00:00
Merge pull request #41904 from BerriAI/litellm_bedrock_batch_retrieve_sigv4_over_env_bearer
fix(bedrock): sign batch retrieve and cancel with deployment credentials when AWS_BEARER_TOKEN_BEDROCK is set
This commit is contained in:
commit
fc0b37ff5d
2 changed files with 69 additions and 0 deletions
|
|
@ -10,6 +10,8 @@ from litellm.types.llms.bedrock import AwsAuthParams, AwsSessionTag
|
|||
from litellm.types.utils import LiteLLMBatch
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from botocore.config import Config
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
# AWS Bedrock model-invocation-job statuses → OpenAI Batch statuses.
|
||||
|
|
@ -31,6 +33,12 @@ _BEDROCK_MIJ_STATUS_TO_OPENAI: Final = {
|
|||
_CANCEL_IDEMPOTENT_STATUSES: Final = frozenset({"cancelling", "cancelled", "completed", "failed", "expired"})
|
||||
|
||||
|
||||
def _sigv4_config() -> "Config":
|
||||
from botocore.config import Config
|
||||
|
||||
return Config(signature_version="v4")
|
||||
|
||||
|
||||
def _extract_region_from_bedrock_arn(arn: str) -> str | None:
|
||||
"""ARN shape: ``arn:aws:bedrock:<region>:<account>:<type>/<id>``"""
|
||||
try:
|
||||
|
|
@ -150,6 +158,7 @@ class BedrockBatchesHandler:
|
|||
aws_access_key_id=creds.access_key,
|
||||
aws_secret_access_key=creds.secret_key,
|
||||
aws_session_token=creds.token,
|
||||
config=_sigv4_config(),
|
||||
)
|
||||
|
||||
def job_status() -> "LiteLLMBatch":
|
||||
|
|
@ -309,6 +318,7 @@ class BedrockBatchesHandler:
|
|||
aws_access_key_id=creds.access_key,
|
||||
aws_secret_access_key=creds.secret_key,
|
||||
aws_session_token=creds.token,
|
||||
config=_sigv4_config(),
|
||||
)
|
||||
|
||||
if logging_obj is not None:
|
||||
|
|
|
|||
|
|
@ -8,10 +8,14 @@ the tests don't hit AWS.
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Iterator, Mapping
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Final
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from botocore.awsrequest import AWSPreparedRequest, AWSResponse
|
||||
|
||||
|
||||
from litellm.llms.bedrock.batches.handler import ( # noqa: E402
|
||||
|
|
@ -570,3 +574,58 @@ def test_cancel_batch_stops_and_polls_the_job_with_the_tagged_session(monkeypatc
|
|||
fake_bedrock.stop_model_invocation_job.assert_called_once_with(jobIdentifier=JOB_ARN)
|
||||
assert batch.status == "cancelled"
|
||||
assert [kwargs["aws_access_key_id"] for kwargs in bedrock_client_kwargs] == ["ASIABATCHCANCELTAGGED"] * 2
|
||||
|
||||
|
||||
class _JsonBody:
|
||||
def __init__(self, payload: bytes) -> None:
|
||||
self._payload: Final = payload
|
||||
|
||||
def stream(self) -> Iterator[bytes]:
|
||||
return iter((self._payload,))
|
||||
|
||||
|
||||
class _AuthorizationRecorder:
|
||||
def __init__(self, body: Mapping[str, object]) -> None:
|
||||
self._payload: Final = json.dumps(body, default=str).encode()
|
||||
self.authorization_headers: tuple[str, ...] = ()
|
||||
|
||||
def send(self, request: AWSPreparedRequest) -> AWSResponse:
|
||||
raw_authorization: Final = request.headers["Authorization"]
|
||||
authorization: Final = (
|
||||
raw_authorization.decode() if isinstance(raw_authorization, bytes) else str(raw_authorization)
|
||||
)
|
||||
self.authorization_headers = (*self.authorization_headers, authorization)
|
||||
return AWSResponse(request.url, 200, {"content-type": "application/json"}, _JsonBody(self._payload))
|
||||
|
||||
|
||||
def test_retrieve_signs_with_deployment_credentials_when_env_bearer_token_is_set(monkeypatch):
|
||||
"""A proxy-wide AWS_BEARER_TOKEN_BEDROCK must not override the deployment's own SigV4 credentials."""
|
||||
monkeypatch.setenv("AWS_BEARER_TOKEN_BEDROCK", "env-bearer-token")
|
||||
recorder: Final = _AuthorizationRecorder(_fake_boto3_response())
|
||||
|
||||
with patch("botocore.httpsession.URLLib3Session.send", recorder.send):
|
||||
batch = BedrockBatchesHandler._handle_model_invocation_job_status(
|
||||
batch_id=JOB_ARN,
|
||||
aws_access_key_id="AKIADEPLOYMENTKEY",
|
||||
aws_secret_access_key="deployment-secret",
|
||||
)
|
||||
|
||||
assert batch.status == "completed"
|
||||
assert len(recorder.authorization_headers) == 1
|
||||
assert recorder.authorization_headers[0].startswith("AWS4-HMAC-SHA256 Credential=AKIADEPLOYMENTKEY/")
|
||||
|
||||
|
||||
def test_cancel_signs_with_deployment_credentials_when_env_bearer_token_is_set(monkeypatch):
|
||||
monkeypatch.setenv("AWS_BEARER_TOKEN_BEDROCK", "env-bearer-token")
|
||||
recorder: Final = _AuthorizationRecorder(_fake_boto3_response(status="Stopped"))
|
||||
|
||||
with patch("botocore.httpsession.URLLib3Session.send", recorder.send):
|
||||
batch = BedrockBatchesHandler.cancel_batch(
|
||||
batch_id=JOB_ARN,
|
||||
aws_access_key_id="AKIADEPLOYMENTKEY",
|
||||
aws_secret_access_key="deployment-secret",
|
||||
)
|
||||
|
||||
assert batch.status == "cancelled"
|
||||
assert len(recorder.authorization_headers) == 2
|
||||
assert all(h.startswith("AWS4-HMAC-SHA256 Credential=AKIADEPLOYMENTKEY/") for h in recorder.authorization_headers)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue