mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
82e6b84f5a
commit
89c3a8216b
16 changed files with 363 additions and 79 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
52
tests/test_litellm/llms/bedrock/event_loop_probe.py
Normal file
52
tests/test_litellm/llms/bedrock/event_loop_probe.py
Normal 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()
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue