refactor(bedrock): resolve AWS credentials from one typed auth struct

Every Bedrock and SageMaker call site hand-copied the same nine aws_* kwargs
into BaseAWSLLM.get_credentials, so each new auth param has to be threaded
into a dozen places and any site that misses one silently assumes the role
with the wrong parameters.

Introduce AwsAuthParams, a frozen pydantic model whose fields are exactly the
credential-shaped params get_credentials accepts, plus resolve_credentials on
BaseAWSLLM and pop_aws_auth_params for the call sites that must strip the keys
out of optional_params. Deriving AWS_AUTH_PARAM_KEYS from the model's fields
means the mirror list in common_utils can no longer drift from the struct.

Behavior is unchanged: the same values reach STS from the same call sites.
Dropping any one field from the resolver fails one of the new tests.

Claude-Session: https://claude.ai/code/session_01E6zsK1DBcXfbetkgX86fw2
This commit is contained in:
ryan-crabbe-berri 2026-09-09 17:51:31 -07:00
parent c1a83fc005
commit b40fc0ac22
13 changed files with 248 additions and 255 deletions

View file

@ -6,11 +6,12 @@ import json
import os
import re
import urllib.parse
from collections.abc import Callable, Mapping
from collections.abc import Callable, Mapping, MutableMapping
from concurrent.futures import ThreadPoolExecutor
from datetime import datetime
from functools import partial
from threading import Lock
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, ParamSpec, TypeVar, cast, get_args, overload
import httpx
@ -31,6 +32,7 @@ from litellm.constants import (
from litellm.litellm_core_utils.aws_partition import contains_bedrock_arn, get_aws_dns_suffix
from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.secret_managers.main import get_secret, get_secret_str
from litellm.types.llms.bedrock import AWS_AUTH_PARAM_KEYS, AwsAuthParams
if TYPE_CHECKING:
from botocore.awsrequest import AWSPreparedRequest
@ -53,6 +55,14 @@ _STS_REGION_FROM_ENDPOINT_PATTERN: Final = re.compile(
SIGV4_COMPUTED_HEADERS: Final = frozenset({"authorization", "x-amz-date", "x-amz-security-token", "date"})
def pop_aws_auth_params(
optional_params: MutableMapping[str, object], # mutable-ok: pops the aws_* keys out of the caller's mapping
) -> AwsAuthParams:
return AwsAuthParams.model_validate(
MappingProxyType({key: optional_params.pop(key, None) for key in AWS_AUTH_PARAM_KEYS})
)
class BedrockRequestTarget(BaseModel):
aws_region_name: str
aws_bedrock_runtime_endpoint: str | None
@ -379,6 +389,20 @@ class BaseAWSLLM(SignsRequestsWithAWS):
else:
return self._get_or_set_cached_credentials(args, self._auth_with_env_vars)
def resolve_credentials(self, auth_params: AwsAuthParams, aws_region_name: str | None) -> Credentials:
return self.get_credentials(
aws_access_key_id=auth_params.aws_access_key_id,
aws_secret_access_key=auth_params.aws_secret_access_key,
aws_session_token=auth_params.aws_session_token,
aws_region_name=aws_region_name,
aws_session_name=auth_params.aws_session_name,
aws_profile_name=auth_params.aws_profile_name,
aws_role_name=auth_params.aws_role_name,
aws_web_identity_token=auth_params.aws_web_identity_token,
aws_sts_endpoint=auth_params.aws_sts_endpoint,
aws_external_id=auth_params.aws_external_id,
)
def _get_aws_region_from_model_arn(self, model: str | None) -> str | None:
try:
# First check if the string contains the expected prefix
@ -1453,22 +1477,10 @@ class BaseAWSLLM(SignsRequestsWithAWS):
from botocore.credentials import Credentials
except ImportError:
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
## CREDENTIALS ##
# pop aws_secret_access_key, aws_access_key_id, aws_region_name from kwargs, since completion calls fail with them
aws_secret_access_key: Final = optional_params.pop("aws_secret_access_key", None)
aws_access_key_id: Final = optional_params.pop("aws_access_key_id", None)
aws_session_token: Final = optional_params.pop("aws_session_token", None)
aws_region_name: Final = self._get_aws_region_name(optional_params, model)
optional_params.pop("aws_region_name", None)
aws_role_name: Final = optional_params.pop("aws_role_name", None)
aws_session_name: Final = optional_params.pop("aws_session_name", None)
aws_profile_name: Final = optional_params.pop("aws_profile_name", None)
aws_web_identity_token: Final = optional_params.pop("aws_web_identity_token", None)
aws_sts_endpoint: Final = optional_params.pop("aws_sts_endpoint", None)
aws_bedrock_runtime_endpoint: Final = optional_params.pop(
"aws_bedrock_runtime_endpoint", None
) # https://bedrock-runtime.{region_name}.amazonaws.com
aws_external_id: Final = optional_params.pop("aws_external_id", None)
auth_params: Final = pop_aws_auth_params(optional_params)
aws_bedrock_runtime_endpoint: Final = optional_params.pop("aws_bedrock_runtime_endpoint", None)
if bearer_token is not None:
return BearerRequestTarget(
@ -1476,18 +1488,7 @@ class BaseAWSLLM(SignsRequestsWithAWS):
aws_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint,
)
credentials: Final[Credentials] = 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,
aws_region_name=aws_region_name,
aws_session_name=aws_session_name,
aws_profile_name=aws_profile_name,
aws_role_name=aws_role_name,
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
)
credentials: Final[Credentials] = self.resolve_credentials(auth_params, aws_region_name)
return Boto3CredentialsInfo(
credentials=credentials,
aws_region_name=aws_region_name,
@ -1621,31 +1622,9 @@ class BaseAWSLLM(SignsRequestsWithAWS):
except ImportError:
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
## CREDENTIALS ##
# pop aws_secret_access_key, aws_access_key_id, aws_session_token, aws_region_name from kwargs, since completion calls fail with them
aws_secret_access_key: Final = optional_params.get("aws_secret_access_key", None)
aws_access_key_id: Final = optional_params.get("aws_access_key_id", None)
aws_session_token: Final = optional_params.get("aws_session_token", None)
aws_role_name: Final = optional_params.get("aws_role_name", None)
aws_session_name: Final = optional_params.get("aws_session_name", None)
aws_profile_name: Final = optional_params.get("aws_profile_name", None)
aws_web_identity_token: Final = optional_params.get("aws_web_identity_token", None)
aws_sts_endpoint: Final = optional_params.get("aws_sts_endpoint", None)
aws_external_id: Final = optional_params.get("aws_external_id", None)
auth_params: Final = AwsAuthParams.model_validate(optional_params)
aws_region_name: Final = self._get_aws_region_name(optional_params=optional_params, model=model)
credentials: Final[Credentials] = 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,
aws_region_name=aws_region_name,
aws_session_name=aws_session_name,
aws_profile_name=aws_profile_name,
aws_role_name=aws_role_name,
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
)
credentials: Final[Credentials] = self.resolve_credentials(auth_params, aws_region_name)
sigv4: Final = SigV4Auth(credentials, service_name, aws_region_name)
headers = headers or {}

View file

@ -6,6 +6,7 @@ from openai.types.batch import BatchRequestCounts
from openai.types.batch import Metadata as OpenAIBatchMetadata
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
from litellm.types.llms.bedrock import AwsAuthParams
from litellm.types.utils import LiteLLMBatch
if TYPE_CHECKING:
@ -128,11 +129,10 @@ class BedrockBatchesHandler:
from litellm.llms.bedrock.batches.transformation import BedrockBatchesConfig
creds: Final = BedrockBatchesConfig().get_credentials(
auth_params: Final = AwsAuthParams(
aws_access_key_id=aws_access_key_id,
aws_secret_access_key=aws_secret_access_key,
aws_session_token=aws_session_token,
aws_region_name=region,
aws_session_name=aws_session_name,
aws_profile_name=aws_profile_name,
aws_role_name=aws_role_name,
@ -140,6 +140,7 @@ class BedrockBatchesHandler:
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
)
creds: Final = BedrockBatchesConfig().resolve_credentials(auth_params, region)
client: Final = boto3.client(
"bedrock",
@ -154,15 +155,7 @@ class BedrockBatchesHandler:
batch_id=batch_id,
aws_region_name=region,
logging_obj=logging_obj,
aws_access_key_id=aws_access_key_id,
aws_secret_access_key=aws_secret_access_key,
aws_session_token=aws_session_token,
aws_session_name=aws_session_name,
aws_profile_name=aws_profile_name,
aws_role_name=aws_role_name,
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
**auth_params.model_dump(),
)
try:
@ -306,18 +299,7 @@ class BedrockBatchesHandler:
# BaseAWSLLM) lazily to avoid a circular import at module load.
from litellm.llms.bedrock.batches.transformation import BedrockBatchesConfig
creds: Final = BedrockBatchesConfig().get_credentials(
aws_access_key_id=kwargs.get("aws_access_key_id"),
aws_secret_access_key=kwargs.get("aws_secret_access_key"),
aws_session_token=kwargs.get("aws_session_token"),
aws_region_name=region,
aws_session_name=kwargs.get("aws_session_name"),
aws_profile_name=kwargs.get("aws_profile_name"),
aws_role_name=kwargs.get("aws_role_name"),
aws_web_identity_token=kwargs.get("aws_web_identity_token"),
aws_sts_endpoint=kwargs.get("aws_sts_endpoint"),
aws_external_id=kwargs.get("aws_external_id"),
)
creds: Final = BedrockBatchesConfig().resolve_credentials(AwsAuthParams.model_validate(kwargs), region)
client: Final = boto3.client(
"bedrock",

View file

@ -21,7 +21,7 @@ from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts
from litellm.types.utils import ModelResponse
from litellm.utils import CustomStreamWrapper
from ..base_aws_llm import BaseAWSLLM, Credentials, bedrock_bearer_token, run_aws_signing
from ..base_aws_llm import BaseAWSLLM, Credentials, bedrock_bearer_token, pop_aws_auth_params, run_aws_signing
from ..common_utils import BedrockError, _get_all_bedrock_regions, error_response_text
from .invoke_handler import AWSEventStreamDecoder, MockResponseIterator, make_call
@ -343,20 +343,8 @@ class BedrockConverseLLM(BaseAWSLLM):
model_id=unencoded_model_id,
)
## CREDENTIALS ##
# pop aws_secret_access_key, aws_access_key_id, aws_region_name from kwargs, since completion calls fail with them
aws_secret_access_key: Final = optional_params.pop("aws_secret_access_key", None)
aws_access_key_id: Final = optional_params.pop("aws_access_key_id", None)
aws_session_token: Final = optional_params.pop("aws_session_token", None)
aws_role_name: Final = optional_params.pop("aws_role_name", None)
aws_session_name: Final = optional_params.pop("aws_session_name", None)
aws_profile_name: Final = optional_params.pop("aws_profile_name", None)
aws_bedrock_runtime_endpoint: Final = optional_params.pop(
"aws_bedrock_runtime_endpoint", None
) # https://bedrock-runtime.{region_name}.amazonaws.com
aws_web_identity_token: Final = optional_params.pop("aws_web_identity_token", None)
aws_sts_endpoint: Final = optional_params.pop("aws_sts_endpoint", None)
aws_external_id: Final = optional_params.pop("aws_external_id", None)
auth_params: Final = pop_aws_auth_params(optional_params)
aws_bedrock_runtime_endpoint: Final = optional_params.pop("aws_bedrock_runtime_endpoint", None)
optional_params.pop("aws_region_name", None)
litellm_params["aws_region_name"] = aws_region_name # [DO NOT DELETE] important for async calls
@ -364,18 +352,7 @@ class BedrockConverseLLM(BaseAWSLLM):
credentials: Final[Credentials | None] = (
None
if bedrock_bearer_token(api_key) is not None
else 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,
aws_region_name=aws_region_name,
aws_session_name=aws_session_name,
aws_profile_name=aws_profile_name,
aws_role_name=aws_role_name,
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
)
else self.resolve_credentials(auth_params, aws_region_name)
)
### SET RUNTIME ENDPOINT ###

View file

@ -28,6 +28,7 @@ from litellm.llms.base_llm.anthropic_messages.transformation import (
from litellm.llms.base_llm.base_utils import BaseLLMModelInfo, BaseTokenCounter
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.secret_managers.main import get_secret, get_secret_str
from litellm.types.llms.bedrock import AWS_AUTH_PARAM_KEYS, AwsAuthParams
if TYPE_CHECKING:
from litellm.types.llms.openai import AllMessageValues
@ -82,18 +83,7 @@ class BedrockError(BaseLLMException):
)
_BEDROCK_AWS_AUTH_PARAMETER_KEYS: Final[tuple[str, ...]] = (
"aws_access_key_id",
"aws_secret_access_key",
"aws_session_token",
"aws_region_name",
"aws_session_name",
"aws_profile_name",
"aws_role_name",
"aws_web_identity_token",
"aws_sts_endpoint",
"aws_external_id",
)
_BEDROCK_AWS_AUTH_PARAMETER_KEYS: Final[tuple[str, ...]] = (*AWS_AUTH_PARAM_KEYS, "aws_region_name")
def merge_bedrock_aws_request_params(
@ -1650,19 +1640,9 @@ class CommonBatchFilesUtils:
except ImportError:
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
# Get AWS credentials using existing methods
aws_region_name: Final = self._base_aws._get_aws_region_name(optional_params=optional_params, model="")
credentials: Final = self._base_aws.get_credentials(
aws_access_key_id=optional_params.get("aws_access_key_id"),
aws_secret_access_key=optional_params.get("aws_secret_access_key"),
aws_session_token=optional_params.get("aws_session_token"),
aws_region_name=aws_region_name,
aws_session_name=optional_params.get("aws_session_name"),
aws_profile_name=optional_params.get("aws_profile_name"),
aws_role_name=optional_params.get("aws_role_name"),
aws_web_identity_token=optional_params.get("aws_web_identity_token"),
aws_sts_endpoint=optional_params.get("aws_sts_endpoint"),
aws_external_id=optional_params.get("aws_external_id"),
credentials: Final = self._base_aws.resolve_credentials(
AwsAuthParams.model_validate(optional_params), aws_region_name
)
# Prepare the request data

View file

@ -26,7 +26,14 @@ from litellm.types.llms.bedrock import (
)
from litellm.types.utils import EmbeddingResponse, LlmProviders
from ..base_aws_llm import AWSPreparedRequest, BaseAWSLLM, Credentials, bedrock_bearer_token, run_aws_signing
from ..base_aws_llm import (
AWSPreparedRequest,
BaseAWSLLM,
Credentials,
bedrock_bearer_token,
pop_aws_auth_params,
run_aws_signing,
)
from ..common_utils import BedrockError
from .amazon_nova_transformation import AmazonNovaEmbeddingConfig
from .amazon_titan_g1_transformation import AmazonTitanG1Config
@ -75,18 +82,8 @@ class BedrockEmbedding(BaseAWSLLM):
optional_params: dict,
bearer_token: str | None = None,
) -> tuple[Credentials | None, str]:
## CREDENTIALS ##
# pop aws_secret_access_key, aws_access_key_id, aws_session_token, aws_region_name from kwargs, since completion calls fail with them
aws_secret_access_key: Final = optional_params.pop("aws_secret_access_key", None)
aws_access_key_id: Final = optional_params.pop("aws_access_key_id", None)
aws_session_token: Final = optional_params.pop("aws_session_token", None)
auth_params: Final = pop_aws_auth_params(optional_params)
aws_region_name = optional_params.pop("aws_region_name", None)
aws_role_name: Final = optional_params.pop("aws_role_name", None)
aws_session_name: Final = optional_params.pop("aws_session_name", None)
aws_profile_name: Final = optional_params.pop("aws_profile_name", None)
aws_web_identity_token: Final = optional_params.pop("aws_web_identity_token", None)
aws_sts_endpoint: Final = optional_params.pop("aws_sts_endpoint", None)
aws_external_id: Final = optional_params.pop("aws_external_id", None)
### SET REGION NAME ###
if aws_region_name is None:
@ -104,20 +101,7 @@ class BedrockEmbedding(BaseAWSLLM):
aws_region_name = "us-west-2"
credentials: Final[Credentials | None] = (
None
if bearer_token is not None
else 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,
aws_region_name=aws_region_name,
aws_session_name=aws_session_name,
aws_profile_name=aws_profile_name,
aws_role_name=aws_role_name,
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
)
None if bearer_token is not None else self.resolve_credentials(auth_params, aws_region_name)
)
return credentials, aws_region_name

View file

@ -11,6 +11,7 @@ from litellm.litellm_core_utils.cloud_storage_security import (
validate_managed_cloud_file_id,
)
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.bedrock import AwsAuthParams
from litellm.types.llms.openai import (
FileContentRequest,
HttpxBinaryResponseContent,
@ -101,19 +102,9 @@ class BedrockFilesHandler(BaseAWSLLM):
allow_legacy_cloud_file_ids=should_allow_legacy_cloud_file_ids(optional_params),
)
# Get AWS credentials
aws_region_name: Final = self._get_aws_region_name(optional_params=optional_params, model="")
credentials: Final[Credentials] = self.get_credentials(
aws_access_key_id=optional_params.get("aws_access_key_id"),
aws_secret_access_key=optional_params.get("aws_secret_access_key"),
aws_session_token=optional_params.get("aws_session_token"),
aws_region_name=aws_region_name,
aws_session_name=optional_params.get("aws_session_name"),
aws_profile_name=optional_params.get("aws_profile_name"),
aws_role_name=optional_params.get("aws_role_name"),
aws_web_identity_token=optional_params.get("aws_web_identity_token"),
aws_sts_endpoint=optional_params.get("aws_sts_endpoint"),
aws_external_id=optional_params.get("aws_external_id"),
credentials: Final[Credentials] = self.resolve_credentials(
AwsAuthParams.model_validate(optional_params), aws_region_name
)
# Create S3 client

View file

@ -41,7 +41,7 @@ from litellm.llms.base_llm.files.transformation import (
BaseFilesConfig,
LiteLLMLoggingObj,
)
from litellm.types.llms.bedrock import BedrockBatchRecordKind
from litellm.types.llms.bedrock import AwsAuthParams, BedrockBatchRecordKind
from litellm.types.llms.openai import (
AllMessageValues,
CreateFileRequest,
@ -133,21 +133,10 @@ def _responses_request_adapter() -> TypeAdapter[ResponsesAPIOptionalRequestParam
return TypeAdapter(ResponsesAPIOptionalRequestParams)
class _BedrockS3RequestParams(BaseModel):
class _BedrockS3RequestParams(AwsAuthParams):
"""Typed view of the credential/region params the S3 GetObject path reads."""
model_config = ConfigDict(extra="ignore")
aws_access_key_id: str | None = None
aws_secret_access_key: str | None = None
aws_session_token: str | None = None
aws_region_name: str | None = None
aws_session_name: str | None = None
aws_profile_name: str | None = None
aws_role_name: str | None = None
aws_web_identity_token: str | None = None
aws_sts_endpoint: str | None = None
aws_external_id: str | None = None
s3_region_name: str | None = None
s3_endpoint_url: str | None = None
@ -1019,20 +1008,8 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
except ImportError:
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
# Get AWS credentials using existing methods
aws_region_name: Final = self._get_aws_region_name(optional_params=optional_params, model="")
credentials: Final = self.get_credentials(
aws_access_key_id=optional_params.get("aws_access_key_id"),
aws_secret_access_key=optional_params.get("aws_secret_access_key"),
aws_session_token=optional_params.get("aws_session_token"),
aws_region_name=aws_region_name,
aws_session_name=optional_params.get("aws_session_name"),
aws_profile_name=optional_params.get("aws_profile_name"),
aws_role_name=optional_params.get("aws_role_name"),
aws_web_identity_token=optional_params.get("aws_web_identity_token"),
aws_sts_endpoint=optional_params.get("aws_sts_endpoint"),
aws_external_id=optional_params.get("aws_external_id"),
)
credentials: Final = self.resolve_credentials(AwsAuthParams.model_validate(optional_params), aws_region_name)
# Calculate SHA256 hash of the content (REQUIRED for S3)
content_hash: Final = hashlib.sha256(content.encode("utf-8")).hexdigest()
@ -1296,18 +1273,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
except ImportError:
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
credentials: Final = self.get_credentials( # any-ok: boto3 Credentials is untyped
aws_access_key_id=request_params.aws_access_key_id,
aws_secret_access_key=request_params.aws_secret_access_key,
aws_session_token=request_params.aws_session_token,
aws_region_name=aws_region_name,
aws_session_name=request_params.aws_session_name,
aws_profile_name=request_params.aws_profile_name,
aws_role_name=request_params.aws_role_name,
aws_web_identity_token=request_params.aws_web_identity_token,
aws_sts_endpoint=request_params.aws_sts_endpoint,
aws_external_id=request_params.aws_external_id,
)
credentials: Final = self.resolve_credentials(request_params, aws_region_name)
empty_body_hash: Final = hashlib.sha256(b"").hexdigest()
aws_request: Final = AWSRequest( # any-ok: botocore AWSRequest is untyped

View file

@ -18,6 +18,7 @@ from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.litellm_core_utils.realtime_streaming import DefaultLoggedRealTimeEventTypes
from litellm.types.llms.bedrock import AwsAuthParams
from litellm.types.llms.openai import OpenAIRealtimeEvents
from litellm.types.realtime import RealtimeResponseTransformInput
@ -149,12 +150,10 @@ class BedrockRealtime(BaseAWSLLM):
verbose_proxy_logger.debug("Bedrock Realtime: Connecting to %s with model %s", endpoint_uri, model)
credentials: Final = await run_aws_signing(
self.get_credentials,
auth_params: Final = AwsAuthParams(
aws_access_key_id=aws_access_key_id,
aws_secret_access_key=aws_secret_access_key,
aws_session_token=aws_session_token,
aws_region_name=aws_region_name,
aws_session_name=aws_session_name,
aws_profile_name=aws_profile_name,
aws_role_name=aws_role_name,
@ -162,7 +161,8 @@ class BedrockRealtime(BaseAWSLLM):
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
)
if credentials is None:
credentials: Final = await run_aws_signing(self.resolve_credentials, auth_params, aws_region_name)
if credentials is None: # pyright: ignore[reportUnnecessaryComparison] # boto3.Session() env fallback yields None
raise BedrockError(
status_code=401,
message=(

View file

@ -6,7 +6,7 @@ from typing import Final
import httpx
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, pop_aws_auth_params
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.utils import ModelResponse, get_secret
@ -23,19 +23,9 @@ class SagemakerChatHandler(BaseAWSLLM):
from botocore.credentials import Credentials
except ImportError:
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
## CREDENTIALS ##
# pop aws_secret_access_key, aws_access_key_id, aws_session_token, aws_region_name from kwargs, since completion calls fail with them
aws_secret_access_key: Final = optional_params.pop("aws_secret_access_key", None)
aws_access_key_id: Final = optional_params.pop("aws_access_key_id", None)
aws_session_token: Final = optional_params.pop("aws_session_token", None)
auth_params: Final = pop_aws_auth_params(optional_params)
aws_region_name = optional_params.pop("aws_region_name", None)
aws_role_name: Final = optional_params.pop("aws_role_name", None)
aws_session_name: Final = optional_params.pop("aws_session_name", None)
aws_profile_name: Final = optional_params.pop("aws_profile_name", None)
optional_params.pop("aws_bedrock_runtime_endpoint", None) # https://bedrock-runtime.{region_name}.amazonaws.com
aws_web_identity_token: Final = optional_params.pop("aws_web_identity_token", None)
aws_sts_endpoint: Final = optional_params.pop("aws_sts_endpoint", None)
aws_external_id: Final = optional_params.pop("aws_external_id", None)
optional_params.pop("aws_bedrock_runtime_endpoint", None)
### SET REGION NAME ###
if aws_region_name is None:
@ -52,18 +42,7 @@ class SagemakerChatHandler(BaseAWSLLM):
if aws_region_name is None:
aws_region_name = "us-west-2"
credentials: Final[Credentials] = 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,
aws_region_name=aws_region_name,
aws_session_name=aws_session_name,
aws_profile_name=aws_profile_name,
aws_role_name=aws_role_name,
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
)
credentials: Final[Credentials] = self.resolve_credentials(auth_params, aws_region_name)
return credentials, aws_region_name
def _prepare_request(

View file

@ -10,7 +10,7 @@ from litellm._logging import verbose_logger
from litellm.litellm_core_utils.asyncify import asyncify
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, pop_aws_auth_params
from litellm.llms.custom_httpx.http_handler import (
_get_httpx_client,
get_async_httpx_client,
@ -46,19 +46,9 @@ class SagemakerLLM(BaseAWSLLM):
from botocore.credentials import Credentials
except ImportError:
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
## CREDENTIALS ##
# pop aws_secret_access_key, aws_access_key_id, aws_session_token, aws_region_name from kwargs, since completion calls fail with them
aws_secret_access_key: Final = optional_params.pop("aws_secret_access_key", None)
aws_access_key_id: Final = optional_params.pop("aws_access_key_id", None)
aws_session_token: Final = optional_params.pop("aws_session_token", None)
auth_params: Final = pop_aws_auth_params(optional_params)
aws_region_name = optional_params.pop("aws_region_name", None)
aws_role_name: Final = optional_params.pop("aws_role_name", None)
aws_session_name: Final = optional_params.pop("aws_session_name", None)
aws_profile_name: Final = optional_params.pop("aws_profile_name", None)
optional_params.pop("aws_bedrock_runtime_endpoint", None) # https://bedrock-runtime.{region_name}.amazonaws.com
aws_web_identity_token: Final = optional_params.pop("aws_web_identity_token", None)
aws_sts_endpoint: Final = optional_params.pop("aws_sts_endpoint", None)
aws_external_id: Final = optional_params.pop("aws_external_id", None)
optional_params.pop("aws_bedrock_runtime_endpoint", None)
### SET REGION NAME ###
if aws_region_name is None:
@ -75,18 +65,7 @@ class SagemakerLLM(BaseAWSLLM):
if aws_region_name is None:
aws_region_name = "us-west-2"
credentials: Final[Credentials] = 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,
aws_region_name=aws_region_name,
aws_session_name=aws_session_name,
aws_profile_name=aws_profile_name,
aws_role_name=aws_role_name,
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
)
credentials: Final[Credentials] = self.resolve_credentials(auth_params, aws_region_name)
return credentials, aws_region_name
def _prepare_request(

View file

@ -3,6 +3,7 @@ from collections.abc import Sequence
from enum import Enum
from typing import TYPE_CHECKING, Any, Final, Literal, TypeAlias
from pydantic import BaseModel, ConfigDict
from typing_extensions import ReadOnly, Required, TypedDict, override
from .openai import ChatCompletionToolCallChunk
@ -1107,6 +1108,25 @@ class BedrockTag(TypedDict):
value: str
class AwsAuthParams(BaseModel):
"""Every credential-shaped aws_* param BaseAWSLLM.get_credentials accepts; region is resolved separately."""
model_config = ConfigDict(frozen=True, extra="ignore")
aws_access_key_id: str | None = None
aws_secret_access_key: str | None = None
aws_session_token: str | None = None
aws_session_name: str | None = None
aws_profile_name: str | None = None
aws_role_name: str | None = None
aws_web_identity_token: str | None = None
aws_sts_endpoint: str | None = None
aws_external_id: str | None = None
AWS_AUTH_PARAM_KEYS: Final[tuple[str, ...]] = tuple(AwsAuthParams.model_fields)
class BedrockCreateBatchRequest(TypedDict, total=False):
"""
Request structure for creating a Bedrock batch inference job.

View file

@ -3278,3 +3278,113 @@ def test_run_aws_signing_leaves_the_default_executor_free_for_other_providers():
other_provider, signing_thread = asyncio.run(scenario())
assert other_provider != signing_thread
assert signing_thread.startswith("aws-signing")
def _recording_boto3_client(recorded: Dict[str, Any]):
"""boto3.client replacement that records the STS client kwargs and the assume-role params."""
def _client(service_name, **client_kwargs):
recorded["client_kwargs"] = client_kwargs
sts = MagicMock()
def _assume(**params):
recorded["assume_role"] = params
return {
"Credentials": {
"AccessKeyId": "ASIAASSUMED",
"SecretAccessKey": "assumed-secret",
"SessionToken": "assumed-token",
"Expiration": datetime.now(timezone.utc) + timedelta(minutes=30),
}
}
def _assume_web_identity(**params):
recorded["assume_role_with_web_identity"] = params
return {
"Credentials": {
"AccessKeyId": "ASIAWEBIDENTITY",
"SecretAccessKey": "assumed-secret",
"SessionToken": "assumed-token",
"Expiration": datetime.now(timezone.utc) + timedelta(minutes=30),
},
"PackedPolicySize": 10,
}
sts.assume_role.side_effect = _assume
sts.assume_role_with_web_identity.side_effect = _assume_web_identity
return sts
return _client
def test_resolve_credentials_forwards_static_keys_role_session_and_external_id():
"""Every field the role-assumption route reads must reach STS, so a dropped struct field fails here."""
from litellm.types.llms.bedrock import AwsAuthParams
auth_params = AwsAuthParams(
aws_access_key_id="AKIACALLER",
aws_secret_access_key="caller-secret",
aws_session_token="caller-token",
aws_role_name="arn:aws:iam::123456789012:role/litellm-target",
aws_session_name="litellm-session",
aws_external_id="litellm-external-id",
aws_sts_endpoint="https://custom-sts.example",
)
recorded: Dict[str, Any] = {}
with (
patch.dict(os.environ, _os_environ_without_aws_keys(), clear=True),
patch("boto3.client", side_effect=_recording_boto3_client(recorded)),
):
credentials = BaseAWSLLM().resolve_credentials(auth_params, "us-east-1")
assert recorded["client_kwargs"]["aws_access_key_id"] == "AKIACALLER"
assert recorded["client_kwargs"]["aws_secret_access_key"] == "caller-secret"
assert recorded["client_kwargs"]["aws_session_token"] == "caller-token"
assert recorded["client_kwargs"]["endpoint_url"] == "https://custom-sts.example"
assert recorded["assume_role"]["RoleArn"] == "arn:aws:iam::123456789012:role/litellm-target"
assert recorded["assume_role"]["RoleSessionName"] == "litellm-session"
assert recorded["assume_role"]["ExternalId"] == "litellm-external-id"
assert credentials.access_key == "ASIAASSUMED"
def test_resolve_credentials_forwards_web_identity_token():
"""A struct carrying a web-identity token must take the web-identity route, not plain role assumption."""
from litellm.types.llms.bedrock import AwsAuthParams
auth_params = AwsAuthParams(
aws_web_identity_token="unresolvable-oidc-token",
aws_role_name="arn:aws:iam::123456789012:role/litellm-wif",
aws_session_name="litellm-wif-session",
)
recorded: Dict[str, Any] = {}
with (
patch.dict(os.environ, _os_environ_without_aws_keys(), clear=True),
patch("boto3.client", side_effect=_recording_boto3_client(recorded)),
):
with pytest.raises(AwsAuthError) as exc:
BaseAWSLLM().resolve_credentials(auth_params, "us-east-1")
assert exc.value.status_code == 401
assert "assume_role" not in recorded
def test_resolve_credentials_forwards_profile_name():
"""The profile route must receive the struct's profile name rather than the ambient session."""
from litellm.types.llms.bedrock import AwsAuthParams
auth_params = AwsAuthParams(aws_profile_name="litellm-qa-profile")
session_instance = MagicMock()
session_instance.get_credentials.return_value = Credentials(
access_key="AKIAPROFILE", secret_key="profile-secret", token=None
)
with (
patch.dict(os.environ, _os_environ_without_aws_keys(), clear=True),
patch("boto3.Session", return_value=session_instance) as mock_session_cls,
):
credentials = BaseAWSLLM().resolve_credentials(auth_params, "us-east-1")
assert mock_session_cls.call_args.kwargs["profile_name"] == "litellm-qa-profile"
assert credentials.access_key == "AKIAPROFILE"

View file

@ -0,0 +1,46 @@
import pytest
from pydantic import ValidationError
from litellm.types.llms.bedrock import AWS_AUTH_PARAM_KEYS, AwsAuthParams
def test_model_validate_keeps_auth_params_and_ignores_request_params():
auth_params = AwsAuthParams.model_validate(
{
"aws_role_name": "arn:aws:iam::999999999999:role/litellm-role",
"aws_session_name": "litellm-session",
"aws_external_id": "litellm-external-id",
"aws_region_name": "us-west-2",
"aws_bedrock_runtime_endpoint": "https://bedrock.example.com",
"model": "anthropic.claude-haiku-4-5-20251001-v1:0",
"temperature": 0.1,
"messages": [{"role": "user", "content": "hi"}],
}
)
assert auth_params.aws_role_name == "arn:aws:iam::999999999999:role/litellm-role"
assert auth_params.aws_session_name == "litellm-session"
assert auth_params.aws_external_id == "litellm-external-id"
assert auth_params.aws_access_key_id is None
assert set(auth_params.model_dump()) == set(AWS_AUTH_PARAM_KEYS)
assert not set(AWS_AUTH_PARAM_KEYS) & {"aws_region_name", "aws_bedrock_runtime_endpoint", "model", "temperature"}
@pytest.mark.parametrize(
("field", "value"),
[
("aws_role_name", 1234),
("aws_session_name", ["litellm-session"]),
("aws_external_id", {"id": "x"}),
],
)
def test_model_validate_rejects_non_string_credentials(field, value):
with pytest.raises(ValidationError):
AwsAuthParams.model_validate({field: value})
def test_frozen_struct_rejects_field_assignment():
auth_params = AwsAuthParams(aws_role_name="arn:aws:iam::999999999999:role/litellm-role")
with pytest.raises(ValidationError):
auth_params.aws_role_name = "arn:aws:iam::999999999999:role/other-role"