mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-21 00:21:49 +00:00
test(bedrock): type the SigV4 request recorder and drop caller-owned mutation
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
37da5b6f4d
commit
6b082d3a01
1 changed files with 31 additions and 20 deletions
|
|
@ -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
|
||||
|
|
@ -572,26 +576,36 @@ def test_cancel_batch_stops_and_polls_the_job_with_the_tagged_session(monkeypatc
|
|||
assert [kwargs["aws_access_key_id"] for kwargs in bedrock_client_kwargs] == ["ASIABATCHCANCELTAGGED"] * 2
|
||||
|
||||
|
||||
def _sigv4_capture_send(sent_headers: list[dict[str, str]], body: dict):
|
||||
import json
|
||||
class _JsonBody:
|
||||
def __init__(self, payload: bytes) -> None:
|
||||
self._payload: Final = payload
|
||||
|
||||
from botocore.awsrequest import AWSResponse
|
||||
def stream(self) -> Iterator[bytes]:
|
||||
return iter((self._payload,))
|
||||
|
||||
def send(_self, request):
|
||||
sent_headers.append({k: v.decode() if isinstance(v, bytes) else v for k, v in request.headers.items()})
|
||||
raw = MagicMock()
|
||||
raw.stream.return_value = iter([json.dumps(body, default=str).encode()])
|
||||
return AWSResponse(request.url, 200, {"content-type": "application/json"}, raw)
|
||||
|
||||
return send
|
||||
class _AuthorizationRecorder:
|
||||
"""Stands in for botocore's HTTP session and records the Authorization header of every request it receives."""
|
||||
|
||||
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")
|
||||
sent_headers: list[dict[str, str]] = []
|
||||
recorder: Final = _AuthorizationRecorder(_fake_boto3_response())
|
||||
|
||||
with patch("botocore.httpsession.URLLib3Session.send", _sigv4_capture_send(sent_headers, _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",
|
||||
|
|
@ -599,18 +613,15 @@ def test_retrieve_signs_with_deployment_credentials_when_env_bearer_token_is_set
|
|||
)
|
||||
|
||||
assert batch.status == "completed"
|
||||
assert len(sent_headers) == 1
|
||||
assert sent_headers[0]["Authorization"].startswith("AWS4-HMAC-SHA256 Credential=AKIADEPLOYMENTKEY/")
|
||||
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")
|
||||
sent_headers: list[dict[str, str]] = []
|
||||
recorder: Final = _AuthorizationRecorder(_fake_boto3_response(status="Stopped"))
|
||||
|
||||
with patch(
|
||||
"botocore.httpsession.URLLib3Session.send",
|
||||
_sigv4_capture_send(sent_headers, _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",
|
||||
|
|
@ -618,5 +629,5 @@ def test_cancel_signs_with_deployment_credentials_when_env_bearer_token_is_set(m
|
|||
)
|
||||
|
||||
assert batch.status == "cancelled"
|
||||
assert len(sent_headers) == 2
|
||||
assert all(h["Authorization"].startswith("AWS4-HMAC-SHA256 Credential=AKIADEPLOYMENTKEY/") for h in sent_headers)
|
||||
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