fix(bedrock): sign requests off the event loop on every async path

SigV4 signing resolves AWS credentials, and botocore refreshes expiring
credentials inside that signing with a blocking HTTP call. Every async
Bedrock path that still signed on the event loop (/v1/messages, Converse,
count tokens, the agent-runtime and Comprehend Medical pass-throughs,
async-invoke status polling, realtime, AgentCore, SQS, S3) now signs on a
worker thread, so one Bedrock request no longer stalls the whole worker.

Fixes #40165
This commit is contained in:
mateo-berri 2026-09-08 11:54:43 -07:00
parent 82e6b84f5a
commit 89c3a8216b
16 changed files with 363 additions and 79 deletions

View file

@ -5,6 +5,7 @@ Sends JSON-RPC envelopes directly to AgentCore endpoints, bypassing the
completion bridge that would otherwise strip the envelope.
"""
import asyncio
import json
from collections.abc import AsyncIterator, Mapping
from typing import Any, Final
@ -45,7 +46,8 @@ class BedrockAgentCoreA2AHandler:
Returns:
A2A JSON-RPC response dict from the AgentCore agent
"""
url, headers, body = BedrockAgentCoreA2ATransformation.get_url_and_signed_request(
url, headers, body = await asyncio.to_thread(
BedrockAgentCoreA2ATransformation.get_url_and_signed_request,
request_id=request_id,
params=params,
litellm_params=litellm_params,
@ -91,7 +93,8 @@ class BedrockAgentCoreA2AHandler:
Yields:
A2A streaming response events from the AgentCore agent
"""
url, headers, body = BedrockAgentCoreA2ATransformation.get_url_and_signed_request(
url, headers, body = await asyncio.to_thread(
BedrockAgentCoreA2ATransformation.get_url_and_signed_request,
request_id=request_id,
params=params,
litellm_params=litellm_params,

View file

@ -366,7 +366,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
# Sign the request
aws_request: Final = AWSRequest(method="PUT", url=url, data=json_string, headers=headers)
aws_region_name: Final = self.get_aws_region_name_for_non_llm_api_calls(aws_region_name=self.s3_region_name)
S3SigV4Auth(credentials, "s3", aws_region_name).add_auth(aws_request)
await asyncio.to_thread(S3SigV4Auth(credentials, "s3", aws_region_name).add_auth, aws_request)
# Prepare the signed headers
signed_headers: Final = dict(aws_request.headers.items())
@ -597,7 +597,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
# Sign the request
aws_request: Final = AWSRequest(method="GET", url=url, headers=headers)
S3SigV4Auth(credentials, "s3", self.s3_region_name).add_auth(aws_request)
await asyncio.to_thread(S3SigV4Auth(credentials, "s3", self.s3_region_name).add_auth, aws_request)
# Prepare the signed headers
signed_headers: Final = dict(aws_request.headers.items())

View file

@ -295,7 +295,7 @@ class SQSLogger(CustomBatchLogger, BaseAWSLLM):
data=prepped.body,
headers=prepped.headers,
)
SigV4Auth(credentials, "sqs", self.sqs_region_name).add_auth(aws_request)
await asyncio.to_thread(SigV4Auth(credentials, "sqs", self.sqs_region_name).add_auth, aws_request)
signed_headers: Final = dict(aws_request.headers.items())

View file

@ -1,3 +1,4 @@
import asyncio
import base64
import hashlib
import json
@ -7,7 +8,7 @@ import urllib.parse
from collections.abc import Callable, Mapping
from datetime import datetime
from threading import Lock
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, cast, get_args, overload
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, ParamSpec, TypeVar, cast, get_args, overload
import httpx
from pydantic import BaseModel, ValidationError
@ -1668,3 +1669,19 @@ class BaseAWSLLM:
request_headers_dict["Authorization"] = incoming_authorization
return request_headers_dict, request.body
_SignParams = ParamSpec("_SignParams")
_SignedRequest = TypeVar("_SignedRequest")
async def sign_request_off_loop_if_aws(
provider_config: object,
sign_request: Callable[_SignParams, _SignedRequest],
/,
*args: _SignParams.args,
**kwargs: _SignParams.kwargs, # kwargs-ok: ParamSpec forwarding keeps the wrapped sign_request signature
) -> _SignedRequest:
if isinstance(provider_config, BaseAWSLLM):
return await asyncio.to_thread(sign_request, *args, **kwargs)
return sign_request(*args, **kwargs)

View file

@ -1,3 +1,4 @@
import asyncio
import json
from collections.abc import Mapping
from types import MappingProxyType
@ -136,7 +137,8 @@ class BedrockConverseLLM(BaseAWSLLM):
)
data: Final = json.dumps(request_data)
prepped: Final = self.get_request_headers(
prepped: Final = await asyncio.to_thread(
self.get_request_headers,
credentials=credentials,
aws_region_name=litellm_params.get("aws_region_name") or "us-west-2",
extra_headers=headers,
@ -206,7 +208,8 @@ class BedrockConverseLLM(BaseAWSLLM):
)
data: Final = json.dumps(request_data)
prepped: Final = self.get_request_headers(
prepped: Final = await asyncio.to_thread(
self.get_request_headers,
credentials=credentials,
aws_region_name=litellm_params.get("aws_region_name") or "us-west-2",
extra_headers=headers,

View file

@ -4,6 +4,7 @@ AWS Bedrock CountTokens API handler.
Simplified handler leveraging existing LiteLLM Bedrock infrastructure.
"""
import asyncio
from typing import Any, Final
import httpx
@ -12,7 +13,7 @@ import litellm
from litellm._logging import verbose_logger
from litellm.llms.bedrock.common_utils import BedrockError
from litellm.llms.bedrock.count_tokens.transformation import BedrockCountTokensConfig
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, get_async_httpx_client
class BedrockCountTokensHandler(BedrockCountTokensConfig):
@ -27,6 +28,7 @@ class BedrockCountTokensHandler(BedrockCountTokensConfig):
request_data: dict[str, Any],
litellm_params: dict[str, Any],
resolved_model: str,
client: AsyncHTTPHandler | None = None,
) -> dict[str, Any]:
"""
Handle a CountTokens request using existing LiteLLM patterns.
@ -75,7 +77,8 @@ class BedrockCountTokensHandler(BedrockCountTokensConfig):
# Extract api_key for bearer token auth if provided
api_key: Final = litellm_params.get("api_key", None)
headers: Final = {"Content-Type": "application/json"}
signed_headers, signed_body = self._sign_request(
signed_headers, signed_body = await asyncio.to_thread(
self._sign_request,
service_name="bedrock",
headers=headers,
optional_params=litellm_params,
@ -85,7 +88,7 @@ class BedrockCountTokensHandler(BedrockCountTokensConfig):
api_key=api_key,
)
async_client: Final = get_async_httpx_client(llm_provider=litellm.LlmProviders.BEDROCK)
async_client: Final = client or get_async_httpx_client(llm_provider=litellm.LlmProviders.BEDROCK)
response: Final = await async_client.post(
endpoint_url,

View file

@ -2,10 +2,11 @@
Handles embedding calls to Bedrock's `/invoke` endpoint
"""
import asyncio
import copy
import json
import urllib.parse
from collections.abc import Callable
from collections.abc import Callable, Mapping
from typing import TYPE_CHECKING, Final, get_args, overload
import httpx
@ -26,7 +27,7 @@ from litellm.types.llms.bedrock import (
)
from litellm.types.utils import EmbeddingResponse, LlmProviders
from ..base_aws_llm import BaseAWSLLM, Credentials, bedrock_bearer_token
from ..base_aws_llm import AWSPreparedRequest, BaseAWSLLM, Credentials, bedrock_bearer_token
from ..common_utils import BedrockError
from .amazon_nova_transformation import AmazonNovaEmbeddingConfig
from .amazon_titan_g1_transformation import AmazonTitanG1Config
@ -41,6 +42,20 @@ if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
def _sign_get_request(
credentials: Credentials, url: str, headers: Mapping[str, str], aws_region_name: str
) -> AWSPreparedRequest:
try:
from botocore.auth import SigV4Auth
from botocore.awsrequest import AWSRequest
except ImportError:
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
request: Final = AWSRequest(method="GET", url=url, data=None, headers=dict(headers))
SigV4Auth(credentials, "bedrock", aws_region_name).add_auth(request)
return request.prepare()
class BedrockEmbedding(BaseAWSLLM):
@overload
def _load_credentials(
@ -599,9 +614,6 @@ class BedrockEmbedding(BaseAWSLLM):
dict: Status response from AWS Bedrock
"""
# Get AWS credentials using the same method as other Bedrock methods
credentials, _ = self._load_credentials(kwargs)
# Get the runtime endpoint
endpoint_url, _ = self.get_runtime_endpoint(
api_base=None,
@ -618,27 +630,13 @@ class BedrockEmbedding(BaseAWSLLM):
# Prepare headers for GET request
headers: Final = {"Content-Type": "application/json"}
# Use AWSRequest directly for GET requests (get_request_headers hardcodes POST)
try:
from botocore.auth import SigV4Auth
from botocore.awsrequest import AWSRequest
except ImportError:
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
def sign_status_request() -> AWSPreparedRequest:
credentials, _ = self._load_credentials(kwargs)
return _sign_get_request(
credentials=credentials, url=status_url, headers=headers, aws_region_name=aws_region_name
)
# Create AWSRequest with GET method and encoded URL
request: Final = AWSRequest(
method="GET",
url=status_url,
data=None, # GET request, no body
headers=headers,
)
# Sign the request - SigV4Auth will create canonical string from request URL
sigv4: Final = SigV4Auth(credentials, "bedrock", aws_region_name)
sigv4.add_auth(request)
# Prepare the request
prepped: Final = request.prepare()
prepped: Final = await asyncio.to_thread(sign_status_request)
# LOGGING
if logging_obj is not None:

View file

@ -149,7 +149,8 @@ class BedrockRealtime(BaseAWSLLM):
verbose_proxy_logger.debug("Bedrock Realtime: Connecting to %s with model %s", endpoint_uri, model)
credentials: Final = self.get_credentials(
credentials: Final = await asyncio.to_thread(
self.get_credentials,
aws_access_key_id=aws_access_key_id,
aws_secret_access_key=aws_secret_access_key,
aws_session_token=aws_session_token,
@ -169,7 +170,7 @@ class BedrockRealtime(BaseAWSLLM):
"or configure credentials in the environment"
),
)
frozen_credentials: Final = credentials.get_frozen_credentials()
frozen_credentials: Final = await asyncio.to_thread(credentials.get_frozen_credentials)
# Initialize Bedrock client with aws_sdk_bedrock_runtime
config: Final = Config(

View file

@ -77,6 +77,7 @@ from litellm.llms.base_llm.vector_store_files.transformation import (
BaseVectorStoreFilesConfig,
)
from litellm.llms.base_llm.videos.transformation import BaseVideoConfig
from litellm.llms.bedrock.base_aws_llm import sign_request_off_loop_if_aws
from litellm.llms.custom_httpx.container_handler import raise_for_error_status
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
@ -1961,7 +1962,9 @@ class BaseLLMHTTPHandler:
api_key=api_key,
)
signed_headers, signed_json_body = provider_config.sign_request(
signed_headers, signed_json_body = await sign_request_off_loop_if_aws(
provider_config,
provider_config.sign_request,
headers=headers,
optional_params=optional_params,
request_data=data,
@ -2062,7 +2065,9 @@ class BaseLLMHTTPHandler:
max_attempts,
)
provider_config.transform_anthropic_messages_request_on_http_error(e=e, request_data=request_body)
headers, signed_json_body = provider_config.sign_request(
headers, signed_json_body = await sign_request_off_loop_if_aws(
provider_config,
provider_config.sign_request,
headers=headers,
optional_params=optional_params_dict,
request_data=request_body,
@ -2222,7 +2227,9 @@ class BaseLLMHTTPHandler:
stream=stream,
)
headers, signed_json_body = anthropic_messages_provider_config.sign_request(
headers, signed_json_body = await sign_request_off_loop_if_aws(
anthropic_messages_provider_config,
anthropic_messages_provider_config.sign_request,
headers=headers,
optional_params=dict(litellm_params), # dynamic aws_* params are passed under litellm_params
request_data=request_body,
@ -2898,7 +2905,9 @@ class BaseLLMHTTPHandler:
fake_stream=fake_stream,
)
headers, signed_body = responses_api_provider_config.sign_request(
headers, signed_body = await sign_request_off_loop_if_aws(
responses_api_provider_config,
responses_api_provider_config.sign_request,
headers=headers,
optional_params=dict(litellm_params),
request_data=data,
@ -4606,7 +4615,9 @@ class BaseLLMHTTPHandler:
)
data = BaseResponsesAPIConfig.normalize_responses_api_request_dict(data)
headers, signed_body = responses_api_provider_config.sign_request(
headers, signed_body = await sign_request_off_loop_if_aws(
responses_api_provider_config,
responses_api_provider_config.sign_request,
headers=headers,
optional_params=dict(litellm_params),
request_data=data,
@ -9833,7 +9844,9 @@ class BaseLLMHTTPHandler:
)
all_optional_params: Final[dict[str, object]] = dict(litellm_params)
all_optional_params.update(vector_store_search_optional_params or {})
headers, signed_json_body = vector_store_provider_config.sign_request(
headers, signed_json_body = await sign_request_off_loop_if_aws(
vector_store_provider_config,
vector_store_provider_config.sign_request,
headers=headers,
optional_params=all_optional_params,
request_data=request_body,

View file

@ -8,6 +8,7 @@ Use litellm with Anthropic SDK, Vertex AI SDK, Cohere SDK, etc.
from __future__ import annotations
import asyncio
import hmac
import inspect
import json
@ -15,6 +16,7 @@ import os
import re
from collections.abc import AsyncGenerator, Callable, Mapping, Sequence
from dataclasses import dataclass
from functools import partial
from types import MappingProxyType
from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, cast
@ -84,6 +86,9 @@ from litellm.utils import ProviderConfigManager
from .passthrough_endpoint_router import PassthroughEndpointRouter
if TYPE_CHECKING:
from botocore.awsrequest import AWSPreparedRequest
from botocore.credentials import Credentials
from litellm.proxy.proxy_server import ProxyConfig as _ProxyConfig
from litellm.router import Router
@ -1099,13 +1104,6 @@ async def bedrock_proxy_route(
"""
create_request_copy(request)
try:
from botocore.auth import SigV4Auth
from botocore.awsrequest import AWSRequest
from botocore.credentials import Credentials
except ImportError:
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
aws_region_name: Final = get_secret_str(secret_name="AWS_REGION_NAME")
if not _is_bedrock_agent_runtime_route(endpoint=endpoint):
return await bedrock_llm_proxy_route(
@ -1139,17 +1137,20 @@ async def bedrock_proxy_route(
from litellm.llms.bedrock.chat import BedrockConverseLLM
bedrock_llm: Final = BedrockConverseLLM()
credentials: Final[Credentials] = bedrock_llm.get_credentials()
sigv4: Final = SigV4Auth(credentials, "bedrock", aws_region_name)
headers: Final = {"Content-Type": "application/json"}
# Assuming the body contains JSON data, parse it
try:
data: Final = await _json_request_body(request)
except Exception as e:
raise HTTPException(status_code=400, detail={"error": e})
_request: Final = AWSRequest(method="POST", url=str(updated_url), data=json.dumps(data), headers=headers)
sigv4.add_auth(_request)
prepped: Final = _request.prepare()
prepped: Final = await asyncio.to_thread(
_sign_aws_json_post,
get_credentials=bedrock_llm.get_credentials,
service_name="bedrock",
aws_region_name=aws_region_name,
url=str(updated_url),
body=json.dumps(data),
headers=MappingProxyType({"Content-Type": "application/json"}),
)
## check for streaming
is_streaming_request = False
@ -1177,6 +1178,25 @@ async def bedrock_proxy_route(
return received_value
def _sign_aws_json_post(
get_credentials: Callable[[], Credentials],
service_name: str,
aws_region_name: str | None,
url: str,
body: str,
headers: Mapping[str, str],
) -> AWSPreparedRequest:
try:
from botocore.auth import SigV4Auth
from botocore.awsrequest import AWSRequest
except ImportError:
raise ImportError(f"Missing boto3 to call {service_name}. Run 'pip install boto3'.")
aws_request: Final = AWSRequest(method="POST", url=url, data=body, headers=dict(headers))
SigV4Auth(get_credentials(), service_name, aws_region_name).add_auth(aws_request)
return aws_request.prepare()
COMPREHEND_MEDICAL_TARGET_PREFIX: Final = "ComprehendMedical_20181030"
@ -1207,13 +1227,6 @@ async def comprehend_medical_proxy_route(
[Docs](https://docs.litellm.ai/docs/pass_through/comprehend_medical)
"""
try:
from botocore.auth import SigV4Auth
from botocore.awsrequest import AWSRequest
from botocore.credentials import Credentials
except ImportError:
raise ImportError("Missing boto3 to call comprehendmedical. Run 'pip install boto3'.")
from .llm_provider_handlers.comprehend_medical_passthrough_logging_handler import (
COMPREHEND_MEDICAL_SUPPORTED_OPERATIONS,
)
@ -1246,18 +1259,21 @@ async def comprehend_medical_proxy_route(
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
credentials: Final[Credentials] = BaseAWSLLM().get_credentials(aws_region_name=aws_region_name)
sigv4: Final = SigV4Auth(credentials, "comprehendmedical", aws_region_name)
headers: Final = MappingProxyType(
{
"Content-Type": "application/x-amz-json-1.1",
"X-Amz-Target": f"{COMPREHEND_MEDICAL_TARGET_PREFIX}.{operation}",
}
)
target_url: Final = f"https://comprehendmedical.{aws_region_name}.{get_aws_dns_suffix(aws_region_name)}/"
_request: Final = AWSRequest(method="POST", url=target_url, data=json.dumps(data), headers=headers)
sigv4.add_auth(_request)
prepped: Final = _request.prepare()
prepped: Final = await asyncio.to_thread(
_sign_aws_json_post,
get_credentials=partial(BaseAWSLLM().get_credentials, aws_region_name=aws_region_name),
service_name="comprehendmedical",
aws_region_name=aws_region_name,
url=target_url,
body=json.dumps(data),
headers=MappingProxyType(
{
"Content-Type": "application/x-amz-json-1.1",
"X-Amz-Target": f"{COMPREHEND_MEDICAL_TARGET_PREFIX}.{operation}",
}
),
)
endpoint_func: Final = create_pass_through_route(
endpoint=operation,

View file

@ -6,6 +6,7 @@ extension, and AWS credential resolution is stubbed so nothing reaches STS.
from __future__ import annotations
import asyncio
from unittest.mock import MagicMock, patch
import httpx
@ -16,6 +17,7 @@ from litellm.llms.bedrock.chat.converse_handler import BedrockConverseLLM
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.rust_bridge import chat_completions as bridge
from litellm.types.utils import ModelResponse
from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe
RUST_RESPONSE = {
"created": 1_700_000_000,
@ -308,7 +310,9 @@ CONVERSE_RESPONSE = {
}
async def _drive_async_completion(*, skip_pre_call_logging: bool, logging_obj):
async def _drive_async_completion(
*, skip_pre_call_logging: bool, logging_obj, credentials: Credentials = RESOLVED_CREDENTIALS
):
"""Run the real `async_completion` with a stubbed transport."""
import httpx as _httpx
@ -335,7 +339,7 @@ async def _drive_async_completion(*, skip_pre_call_logging: bool, logging_obj):
stream=None,
optional_params={"maxTokens": 16},
litellm_params={"aws_region_name": "us-west-2"},
credentials=RESOLVED_CREDENTIALS,
credentials=credentials,
headers={},
client=client,
skip_pre_call_logging=skip_pre_call_logging,
@ -357,6 +361,23 @@ async def test_async_completion_logs_pre_call_by_default():
assert logging_obj.pre_call.call_count == 1
@pytest.mark.asyncio
async def test_async_completion_signs_off_the_event_loop(monkeypatch):
"""Regression for issue #40165: botocore refreshes expiring credentials inside SigV4 signing with a
blocking HTTP call, so `async_completion` must sign on a worker thread to keep the loop serving."""
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
probe = EventLoopProbe()
release = asyncio.create_task(probe.release_refresh_from_the_loop())
response = await _drive_async_completion(
skip_pre_call_logging=False, logging_obj=MagicMock(), credentials=probe.credentials()
)
await release
assert response.choices[0].message.content == "hi"
assert probe.served_during_refresh is True
def _sync_client_returning_converse_response():
client = MagicMock()
client.post.side_effect = lambda **_kwargs: httpx.Response(

View file

@ -0,0 +1,50 @@
import asyncio
from unittest.mock import AsyncMock
import httpx
import pytest
from litellm.llms.bedrock.count_tokens.handler import BedrockCountTokensHandler
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe
class _ProbedCountTokensHandler(BedrockCountTokensHandler):
def __init__(self, probe: EventLoopProbe) -> None:
super().__init__()
self._probe = probe
def get_credentials(self, **kwargs):
return self._probe.credentials()
@pytest.mark.asyncio
async def test_handle_count_tokens_request_signs_off_the_event_loop(monkeypatch):
"""Regression for issue #40165: the count_tokens handler signed on the loop, so botocore's blocking
credential refresh inside SigV4 stalled every other request on the worker."""
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
probe = EventLoopProbe()
client = AsyncMock(spec=AsyncHTTPHandler)
client.post = AsyncMock(
return_value=httpx.Response(
200,
json={"inputTokens": 7},
request=httpx.Request("POST", "https://bedrock-runtime.us-west-2.amazonaws.com/"),
)
)
release = asyncio.create_task(probe.release_refresh_from_the_loop())
result = await _ProbedCountTokensHandler(probe).handle_count_tokens_request(
request_data={
"model": "us.anthropic.claude-haiku-4-5-20251001-v1:0",
"messages": [{"role": "user", "content": "hi"}],
},
litellm_params={"aws_region_name": "us-west-2"},
resolved_model="us.anthropic.claude-haiku-4-5-20251001-v1:0",
client=client,
)
await release
assert result == {"input_tokens": 7}
assert client.post.call_args.kwargs["headers"]["Authorization"].startswith("AWS4-HMAC-SHA256")
assert probe.served_during_refresh is True

View file

@ -0,0 +1,52 @@
"""Refreshable credentials whose refresh only completes while the event loop keeps serving."""
from __future__ import annotations
import asyncio
import threading
from datetime import datetime, timedelta, timezone
from typing import Final
from botocore.credentials import RefreshableCredentials
REFRESH_RELEASE_TIMEOUT_SECONDS: Final = 2.0
class EventLoopProbe:
"""Blocks inside botocore's credential refresh until a coroutine on the loop releases it.
Signing on the event loop thread can never be released, so `served_during_refresh` reads False there
and True only when the refresh ran on another thread while the loop stayed responsive.
"""
def __init__(self) -> None:
self.refresh_started: Final = threading.Event()
self.loop_served: Final = threading.Event()
self.served_during_refresh: bool | None = None
def refresh(self) -> dict[str, str | None]:
self.refresh_started.set()
served: Final = self.loop_served.wait(timeout=REFRESH_RELEASE_TIMEOUT_SECONDS)
if self.served_during_refresh is None:
self.served_during_refresh = served
return {
"access_key": "AKIAREFRESHED",
"secret_key": "refreshed-secret",
"token": None,
"expiry_time": (datetime.now(timezone.utc) + timedelta(hours=1)).isoformat(),
}
def credentials(self) -> RefreshableCredentials:
return RefreshableCredentials(
access_key="AKIASTALE",
secret_key="stale-secret",
token=None,
expiry_time=datetime.now(timezone.utc) + timedelta(seconds=60),
refresh_using=self.refresh,
method="event-loop-probe",
)
async def release_refresh_from_the_loop(self) -> None:
while not self.refresh_started.is_set():
await asyncio.sleep(0.005)
self.loop_served.set()

View file

@ -1,3 +1,4 @@
import asyncio
import json
import os
import threading
@ -22,7 +23,9 @@ from litellm.llms.bedrock.base_aws_llm import (
AwsAuthError,
BaseAWSLLM,
Boto3CredentialsInfo,
sign_request_off_loop_if_aws,
)
from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe
# Global variable for the base_aws_llm.py file path
@ -3215,3 +3218,24 @@ class TestGetRequestHeadersResign:
extra_headers={"Authorization": "Bearer foo"},
)
assert prepped.headers["Authorization"] == "Bearer foo"
@pytest.mark.asyncio
async def test_sign_request_off_loop_if_aws_keeps_the_loop_serving_while_credentials_refresh():
"""Regression for issue #40165: an AWS provider's signing (and the botocore credential refresh
inside it) must run off the event loop, so other requests keep being served meanwhile."""
probe = EventLoopProbe()
def sign(headers: dict[str, str]) -> dict[str, str]:
request = AWSRequest(
method="POST", url="https://bedrock-runtime.us-west-2.amazonaws.com/", data="{}", headers=headers
)
SigV4Auth(probe.credentials(), "bedrock", "us-west-2").add_auth(request)
return dict(request.headers)
release = asyncio.create_task(probe.release_refresh_from_the_loop())
signed = await sign_request_off_loop_if_aws(BaseAWSLLM(), sign, headers={"Content-Type": "application/json"})
await release
assert "Authorization" in signed
assert probe.served_during_refresh is True

View file

@ -29,10 +29,14 @@ from litellm.llms.custom_httpx.llm_http_handler import (
_rust_responses_websocket_enabled,
)
from litellm.llms.azure.videos.transformation import AzureVideoConfig
from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import (
AmazonAnthropicClaudeMessagesConfig,
)
from litellm.llms.openai.videos.transformation import OpenAIVideoConfig
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import ImageObject, ImageResponse, ModelResponse, TranscriptionResponse
from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe
_ACTIVE_KEY = "_code_interpreter_interception_active"
_SANDBOX_KEY = "_code_interpreter_interception_sandbox_key"
@ -749,6 +753,62 @@ async def test_anthropic_messages_streaming_response_aclose_closes_agentic_upstr
assert tracker.closed is True
class _ProbedBedrockMessagesConfig(AmazonAnthropicClaudeMessagesConfig):
def __init__(self, probe: EventLoopProbe) -> None:
super().__init__()
self._probe = probe
def get_credentials(self, **kwargs):
return self._probe.credentials()
@pytest.mark.asyncio
async def test_async_anthropic_messages_handler_signs_bedrock_off_the_event_loop(monkeypatch):
"""Regression for issue #40165: /v1/messages on Bedrock signed on the loop, so botocore's blocking
credential refresh inside SigV4 stalled every other request on the worker."""
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
probe = EventLoopProbe()
handler = BaseLLMHTTPHandler()
upstream_response = httpx.Response(
200,
json={
"id": "msg_123",
"type": "message",
"role": "assistant",
"content": [{"type": "text", "text": "hi"}],
"model": "claude-haiku-4-5-20251001",
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 1, "output_tokens": 1},
},
request=httpx.Request("POST", "https://bedrock-runtime.us-west-2.amazonaws.com/"),
)
mock_client = AsyncMock(spec=AsyncHTTPHandler)
mock_client.post = AsyncMock(return_value=upstream_response)
mock_logging_obj = Mock()
mock_logging_obj.model_call_details = {}
mock_logging_obj.dynamic_success_callbacks = None
release = asyncio.create_task(probe.release_refresh_from_the_loop())
await handler.async_anthropic_messages_handler(
model="us.anthropic.claude-haiku-4-5-20251001-v1:0",
messages=[{"role": "user", "content": "hi"}],
anthropic_messages_provider_config=_ProbedBedrockMessagesConfig(probe),
anthropic_messages_optional_request_params={"max_tokens": 16},
custom_llm_provider="bedrock",
litellm_params=GenericLiteLLMParams(aws_region_name="us-west-2"),
logging_obj=mock_logging_obj,
client=mock_client,
stream=False,
kwargs={},
)
await release
sent_headers = mock_client.post.call_args.kwargs["headers"]
assert sent_headers["Authorization"].startswith("AWS4-HMAC-SHA256")
assert probe.served_during_refresh is True
@pytest.mark.asyncio
async def test_async_anthropic_messages_handler_passes_litellm_metadata():
"""Ensure litellm_metadata from kwargs is forwarded via update_from_kwargs.

View file

@ -1,3 +1,4 @@
import asyncio
import base64
import contextlib
import json
@ -19,6 +20,7 @@ from starlette.datastructures import FormData
import litellm
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
BaseOpenAIPassThroughHandler,
@ -1852,11 +1854,11 @@ class TestBedrockAgentRuntimePassthroughToggle:
return request
@contextlib.contextmanager
def _patched_dispatch(self, general_settings: Mapping[str, object]):
def _patched_dispatch(self, general_settings: Mapping[str, object], credentials: object | None = None):
from botocore.credentials import Credentials
bedrock_llm: Final = Mock()
bedrock_llm.get_credentials = Mock(return_value=Credentials("ak", "sk"))
bedrock_llm.get_credentials = Mock(return_value=credentials or Credentials("ak", "sk"))
forwarder: Final = AsyncMock(return_value="forwarded")
with (
@ -1891,6 +1893,27 @@ class TestBedrockAgentRuntimePassthroughToggle:
forwarder.assert_awaited_once()
assert "bedrock-agent-runtime.us-east-1.amazonaws.com" in create_route.call_args.kwargs["target"]
@pytest.mark.asyncio
async def test_agent_runtime_dispatch_signs_off_the_event_loop(self, monkeypatch):
"""Regression for issue #40165: the agent-runtime pass-through signed on the loop, so botocore's
blocking credential refresh inside SigV4 stalled every other request on the worker."""
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
probe: Final = EventLoopProbe()
release: Final = asyncio.create_task(probe.release_refresh_from_the_loop())
with self._patched_dispatch(MappingProxyType({}), credentials=probe.credentials()) as (create_route, forwarder):
result: Final = await bedrock_proxy_route(
endpoint=self.AGENT_RUNTIME_ENDPOINT,
request=self._mock_request(),
fastapi_response=Mock(),
user_api_key_dict=UserAPIKeyAuth(),
)
await release
assert result == "forwarded"
assert create_route.call_args.kwargs["custom_headers"]["Authorization"].startswith("AWS4-HMAC-SHA256")
assert probe.served_during_refresh is True
@pytest.mark.asyncio
@pytest.mark.parametrize("value", (True, "true", "True"))
async def test_agent_runtime_dispatch_rejected_when_disabled(self, value: bool | str):