mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(bedrock): stop stale SigV4 headers clobbering fresh signature on strip-and-retry re-sign (#32371)
* fix(bedrock): stop stale SigV4 headers clobbering fresh signature on re-sign When the Anthropic /v1/messages strip-thinking-and-retry path re-signs a Bedrock request, _sign_request received attempt 1's already-signed headers and copied the old Authorization and X-Amz-Date back over the freshly computed SigV4 signature, so the retry POSTed the stripped body with a signature for the original body and AWS returned 403 SignatureDoesNotMatch. Skip SigV4-computed headers (authorization, x-amz-date, x-amz-security-token, date) when restoring caller headers after signing, and only preserve a caller-supplied Authorization that is not itself a SigV4 header so bearer-token setups keep working. * fix(bedrock): apply the same stale-header guard to get_request_headers
This commit is contained in:
parent
06a97c83bd
commit
c8b78d49dd
3 changed files with 288 additions and 6 deletions
|
|
@ -50,6 +50,8 @@ _STS_REGION_FROM_ENDPOINT_PATTERN = re.compile(
|
|||
r"(?:^|\.)sts(?:-fips)?\.([a-z0-9-]+)\.(?:amazonaws\.com(?:\.cn)?|vpce\.amazonaws\.com)"
|
||||
)
|
||||
|
||||
SIGV4_COMPUTED_HEADERS = frozenset({"authorization", "x-amz-date", "x-amz-security-token", "date"})
|
||||
|
||||
|
||||
class Boto3CredentialsInfo(BaseModel):
|
||||
credentials: Credentials
|
||||
|
|
@ -1400,11 +1402,13 @@ class BaseAWSLLM:
|
|||
|
||||
# Add back all original headers (including forwarded ones) after signature calculation
|
||||
for header_name, header_value in headers.items():
|
||||
if header_value is not None:
|
||||
if header_value is not None and header_name.lower() not in SIGV4_COMPUTED_HEADERS:
|
||||
request.headers[header_name] = header_value
|
||||
|
||||
if (
|
||||
extra_headers is not None and "Authorization" in extra_headers
|
||||
extra_headers is not None
|
||||
and "Authorization" in extra_headers
|
||||
and not extra_headers["Authorization"].startswith("AWS4-HMAC-SHA256")
|
||||
): # prevent sigv4 from overwriting the auth header
|
||||
request.headers["Authorization"] = extra_headers["Authorization"]
|
||||
prepped = request.prepare()
|
||||
|
|
@ -1527,9 +1531,15 @@ class BaseAWSLLM:
|
|||
# Add back original headers after signing. Only headers in SignedHeaders
|
||||
# are integrity-protected; forwarded headers (x-forwarded-*) must remain unsigned.
|
||||
for header_name, header_value in headers.items():
|
||||
if header_value is not None:
|
||||
if header_value is not None and header_name.lower() not in SIGV4_COMPUTED_HEADERS:
|
||||
request_headers_dict[header_name] = header_value
|
||||
if headers is not None and "Authorization" in headers: # prevent sigv4 from overwriting the auth header
|
||||
request_headers_dict["Authorization"] = headers["Authorization"]
|
||||
incoming_authorization = next(
|
||||
(value for name, value in headers.items() if name.lower() == "authorization" and value is not None),
|
||||
None,
|
||||
)
|
||||
if incoming_authorization is not None and not incoming_authorization.startswith(
|
||||
"AWS4-HMAC-SHA256"
|
||||
): # prevent sigv4 from overwriting the auth header
|
||||
request_headers_dict["Authorization"] = incoming_authorization
|
||||
|
||||
return request_headers_dict, request.body
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ sys.path.insert(
|
|||
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any, Dict
|
||||
from typing import Any, Dict, Optional
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from botocore.awsrequest import AWSPreparedRequest, AWSRequest
|
||||
|
|
@ -2653,3 +2653,184 @@ class TestGetBedrockModelIdArnHandling:
|
|||
"""invoke/ prefix stripping still works after the fix."""
|
||||
model_id = self._call("invoke/anthropic.claude-3-sonnet-20240229-v1:0")
|
||||
assert model_id == "anthropic.claude-3-sonnet-20240229-v1:0"
|
||||
|
||||
|
||||
def _recomputed_sigv4_signature(url: str, secret_key: str, authorization: str, headers: Dict[str, Any], body) -> str:
|
||||
import hashlib
|
||||
import hmac
|
||||
from urllib.parse import urlparse
|
||||
|
||||
parsed = urlparse(url)
|
||||
credential_scope = authorization.split("Credential=")[1].split(",")[0].split("/", 1)[1]
|
||||
signed_header_names = authorization.split("SignedHeaders=")[1].split(",")[0].split(";")
|
||||
header_lookup = {name.lower(): str(value) for name, value in headers.items()}
|
||||
header_lookup["host"] = parsed.netloc
|
||||
body_bytes = body if isinstance(body, bytes) else str(body).encode()
|
||||
canonical_request = "\n".join(
|
||||
[
|
||||
"POST",
|
||||
parsed.path or "/",
|
||||
"",
|
||||
"".join(f"{name}:{header_lookup[name]}\n" for name in signed_header_names),
|
||||
";".join(signed_header_names),
|
||||
hashlib.sha256(body_bytes).hexdigest(),
|
||||
]
|
||||
)
|
||||
string_to_sign = "\n".join(
|
||||
[
|
||||
"AWS4-HMAC-SHA256",
|
||||
header_lookup["x-amz-date"],
|
||||
credential_scope,
|
||||
hashlib.sha256(canonical_request.encode()).hexdigest(),
|
||||
]
|
||||
)
|
||||
key = f"AWS4{secret_key}".encode()
|
||||
for scope_part in credential_scope.split("/"):
|
||||
key = hmac.new(key, scope_part.encode(), hashlib.sha256).digest()
|
||||
return hmac.new(key, string_to_sign.encode(), hashlib.sha256).hexdigest()
|
||||
|
||||
|
||||
class TestSignRequestResign:
|
||||
"""Regression: retrying a Bedrock request with headers from a previous SigV4 sign
|
||||
(e.g. the /v1/messages strip-thinking-and-retry path) must produce a fresh
|
||||
Authorization / X-Amz-Date for the new body, not inherit the stale ones and 403."""
|
||||
|
||||
URL = "https://bedrock-runtime.us-east-1.amazonaws.com/model/test-model/invoke"
|
||||
ACCESS_KEY = "AKIAIOSFODNN7EXAMPLE"
|
||||
SECRET_KEY = "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY"
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clean_aws_env(self, monkeypatch):
|
||||
for env_var in ("AWS_BEARER_TOKEN_BEDROCK", "AWS_SESSION_TOKEN", "AWS_PROFILE"):
|
||||
monkeypatch.delenv(env_var, raising=False)
|
||||
|
||||
def _optional_params(self) -> Dict[str, Any]:
|
||||
return {
|
||||
"aws_access_key_id": self.ACCESS_KEY,
|
||||
"aws_secret_access_key": self.SECRET_KEY,
|
||||
"aws_region_name": "us-east-1",
|
||||
}
|
||||
|
||||
def _sign(self, headers: Dict[str, Any], request_data: Dict[str, Any]):
|
||||
return BaseAWSLLM()._sign_request(
|
||||
service_name="bedrock",
|
||||
headers=headers,
|
||||
optional_params=self._optional_params(),
|
||||
request_data=request_data,
|
||||
api_base=self.URL,
|
||||
)
|
||||
|
||||
def test_resign_with_previously_signed_headers_replaces_stale_sigv4_headers(self):
|
||||
original_body = {
|
||||
"messages": [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [{"type": "thinking", "thinking": "x", "signature": ""}],
|
||||
}
|
||||
]
|
||||
}
|
||||
first_headers, _ = self._sign(headers={"Content-Type": "application/json"}, request_data=original_body)
|
||||
assert first_headers["Authorization"].startswith("AWS4-HMAC-SHA256")
|
||||
|
||||
stale_headers = {**first_headers, "X-Amz-Date": "20200101T000000Z"}
|
||||
stripped_body = {"messages": [{"role": "user", "content": "hi"}]}
|
||||
second_headers, second_signed_body = self._sign(headers=stale_headers, request_data=stripped_body)
|
||||
|
||||
assert second_headers["X-Amz-Date"] != "20200101T000000Z"
|
||||
assert second_headers["Authorization"] != stale_headers["Authorization"]
|
||||
assert second_headers["Authorization"].split("Signature=")[1] == _recomputed_sigv4_signature(
|
||||
url=self.URL,
|
||||
secret_key=self.SECRET_KEY,
|
||||
authorization=second_headers["Authorization"],
|
||||
headers=second_headers,
|
||||
body=second_signed_body,
|
||||
)
|
||||
|
||||
def test_forwarded_headers_still_added_back_after_signing(self):
|
||||
signed_headers, _ = self._sign(
|
||||
headers={"Content-Type": "application/json", "anthropic-version": "bedrock-2023-05-31"},
|
||||
request_data={"messages": []},
|
||||
)
|
||||
assert signed_headers["anthropic-version"] == "bedrock-2023-05-31"
|
||||
assert signed_headers["Content-Type"] == "application/json"
|
||||
|
||||
def test_caller_supplied_bearer_authorization_survives_signing(self):
|
||||
signed_headers, _ = self._sign(
|
||||
headers={"Content-Type": "application/json", "Authorization": "Bearer caller-token"},
|
||||
request_data={"messages": []},
|
||||
)
|
||||
assert signed_headers["Authorization"] == "Bearer caller-token"
|
||||
|
||||
|
||||
class TestGetRequestHeadersResign:
|
||||
"""Regression: get_request_headers (invoke/converse/embed/image paths) must not let
|
||||
stale SigV4 values present in the input headers clobber the freshly computed signature."""
|
||||
|
||||
URL = "https://bedrock-runtime.us-east-1.amazonaws.com/model/test-model/converse"
|
||||
ACCESS_KEY = "AKIAIOSFODNN7EXAMPLE"
|
||||
SECRET_KEY = "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY"
|
||||
SESSION_TOKEN = "fresh-session-token"
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clean_aws_env(self, monkeypatch):
|
||||
for env_var in ("AWS_BEARER_TOKEN_BEDROCK", "AWS_SESSION_TOKEN", "AWS_PROFILE"):
|
||||
monkeypatch.delenv(env_var, raising=False)
|
||||
|
||||
def _prepare(self, headers: Dict[str, Any], data: str, extra_headers: Optional[Dict[str, str]] = None):
|
||||
return BaseAWSLLM().get_request_headers(
|
||||
credentials=Credentials(self.ACCESS_KEY, self.SECRET_KEY, self.SESSION_TOKEN),
|
||||
aws_region_name="us-east-1",
|
||||
extra_headers=extra_headers,
|
||||
endpoint_url=self.URL,
|
||||
data=data,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
def test_stale_sigv4_headers_in_input_replaced_by_fresh_signature(self):
|
||||
first_prepped = self._prepare(
|
||||
headers={"Content-Type": "application/json"},
|
||||
data=json.dumps({"messages": [{"role": "user", "content": "original"}]}),
|
||||
)
|
||||
stale_authorization = first_prepped.headers["Authorization"]
|
||||
assert stale_authorization.startswith("AWS4-HMAC-SHA256")
|
||||
|
||||
stale_headers = {
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": stale_authorization,
|
||||
"X-Amz-Date": "20200101T000000Z",
|
||||
"X-Amz-Security-Token": "stale-session-token",
|
||||
}
|
||||
retry_data = json.dumps({"messages": [{"role": "user", "content": "retry"}]})
|
||||
second_prepped = self._prepare(headers=stale_headers, data=retry_data)
|
||||
|
||||
assert second_prepped.headers["X-Amz-Date"] != "20200101T000000Z"
|
||||
assert second_prepped.headers["X-Amz-Security-Token"] == self.SESSION_TOKEN
|
||||
assert second_prepped.headers["Authorization"] != stale_authorization
|
||||
assert second_prepped.headers["Authorization"].split("Signature=")[1] == _recomputed_sigv4_signature(
|
||||
url=self.URL,
|
||||
secret_key=self.SECRET_KEY,
|
||||
authorization=second_prepped.headers["Authorization"],
|
||||
headers=dict(second_prepped.headers),
|
||||
body=retry_data,
|
||||
)
|
||||
|
||||
def test_forwarded_headers_still_added_back_after_signing(self):
|
||||
prepped = self._prepare(
|
||||
headers={
|
||||
"Content-Type": "application/json",
|
||||
"anthropic-version": "bedrock-2023-05-31",
|
||||
"user-agent": "litellm-test-client",
|
||||
},
|
||||
data=json.dumps({"messages": []}),
|
||||
)
|
||||
assert prepped.headers["anthropic-version"] == "bedrock-2023-05-31"
|
||||
assert prepped.headers["user-agent"] == "litellm-test-client"
|
||||
assert prepped.headers["Content-Type"] == "application/json"
|
||||
|
||||
def test_extra_headers_bearer_authorization_still_overrides_signature(self):
|
||||
prepped = self._prepare(
|
||||
headers={"Content-Type": "application/json"},
|
||||
data=json.dumps({"messages": []}),
|
||||
extra_headers={"Authorization": "Bearer foo"},
|
||||
)
|
||||
assert prepped.headers["Authorization"] == "Bearer foo"
|
||||
|
|
|
|||
|
|
@ -1837,3 +1837,94 @@ async def test_alist_input_items_surfaces_upstream_error_status():
|
|||
)
|
||||
|
||||
assert excinfo.value.status_code == 404
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_anthropic_invalid_thinking_signature_retry_resigns_bedrock_request(monkeypatch):
|
||||
"""Regression: after Bedrock rejects a replayed thinking block (400 invalid signature),
|
||||
the strip-and-retry re-sign must not inherit attempt 1's SigV4 Authorization/X-Amz-Date;
|
||||
reusing them over the new stripped body makes AWS return 403 SignatureDoesNotMatch."""
|
||||
from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import (
|
||||
AmazonAnthropicClaudeMessagesConfig,
|
||||
)
|
||||
|
||||
for env_var in ("AWS_BEARER_TOKEN_BEDROCK", "AWS_SESSION_TOKEN", "AWS_PROFILE"):
|
||||
monkeypatch.delenv(env_var, raising=False)
|
||||
|
||||
handler = BaseLLMHTTPHandler()
|
||||
provider_config = AmazonAnthropicClaudeMessagesConfig()
|
||||
litellm_params = GenericLiteLLMParams(
|
||||
aws_access_key_id="AKIAIOSFODNN7EXAMPLE",
|
||||
aws_secret_access_key="wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY",
|
||||
aws_region_name="us-east-1",
|
||||
)
|
||||
request_url = "https://bedrock-runtime.us-east-1.amazonaws.com/model/test-model/invoke"
|
||||
request_body = {
|
||||
"anthropic_version": "bedrock-2023-05-31",
|
||||
"max_tokens": 100,
|
||||
"messages": [
|
||||
{"role": "user", "content": "hi"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{"type": "thinking", "thinking": "x", "signature": ""},
|
||||
{"type": "text", "text": "ok"},
|
||||
],
|
||||
},
|
||||
{"role": "user", "content": "continue"},
|
||||
],
|
||||
}
|
||||
first_attempt_headers, signed_json_body = provider_config.sign_request(
|
||||
headers={"Content-Type": "application/json"},
|
||||
optional_params=dict(litellm_params),
|
||||
request_data=request_body,
|
||||
api_base=request_url,
|
||||
api_key=None,
|
||||
stream=False,
|
||||
fake_stream=False,
|
||||
model="test-model",
|
||||
)
|
||||
|
||||
posts: list = []
|
||||
invalid_signature_response = httpx.Response(
|
||||
400,
|
||||
text='{"message": "messages.1.content.0: Invalid `signature` in `thinking` block"}',
|
||||
request=httpx.Request("POST", request_url),
|
||||
)
|
||||
ok_response = httpx.Response(200, json={"id": "msg_1"}, request=httpx.Request("POST", request_url))
|
||||
|
||||
class FakeAsyncClient:
|
||||
async def post(self, url, headers, data, stream=False, logging_obj=None):
|
||||
posts.append({"headers": dict(headers), "data": data})
|
||||
return invalid_signature_response if len(posts) == 1 else ok_response
|
||||
|
||||
logging_obj = Mock()
|
||||
logging_obj.model_call_details = {}
|
||||
|
||||
response = await handler._async_post_anthropic_messages_with_http_error_retry(
|
||||
async_httpx_client=FakeAsyncClient(),
|
||||
request_url=request_url,
|
||||
headers=dict(first_attempt_headers),
|
||||
signed_json_body=signed_json_body,
|
||||
request_body=request_body,
|
||||
stream=False,
|
||||
logging_obj=logging_obj,
|
||||
provider_config=provider_config,
|
||||
litellm_params=litellm_params,
|
||||
api_key=None,
|
||||
model="test-model",
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert len(posts) == 2
|
||||
retry_payload = json.loads(posts[1]["data"])
|
||||
retry_blocks = [
|
||||
block
|
||||
for message in retry_payload["messages"]
|
||||
if isinstance(message.get("content"), list)
|
||||
for block in message["content"]
|
||||
]
|
||||
assert retry_blocks and all(block["type"] != "thinking" for block in retry_blocks)
|
||||
retry_authorization = posts[1]["headers"]["Authorization"]
|
||||
assert retry_authorization.startswith("AWS4-HMAC-SHA256")
|
||||
assert retry_authorization != first_attempt_headers["Authorization"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue