Merge pull request #40500 from BerriAI/litellm_bedrock_aws_auth_params

fix(bedrock): send aws_session_tags on every STS call via one typed auth struct
This commit is contained in:
Mateo Wang 2026-09-18 12:01:52 -07:00 committed by GitHub
commit a2626726a2
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
17 changed files with 396 additions and 286 deletions

View file

@ -6,7 +6,7 @@ import json
import os
import re
import urllib.parse
from collections.abc import Callable, Mapping, Sequence
from collections.abc import Callable, Mapping, MutableMapping, Sequence
from concurrent.futures import ThreadPoolExecutor
from datetime import datetime
from functools import partial
@ -33,7 +33,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 AwsSessionTag
from litellm.types.llms.bedrock import AWS_AUTH_PARAM_KEYS, AwsAuthParams, AwsSessionTag
if TYPE_CHECKING:
from botocore.awsrequest import AWSPreparedRequest
@ -168,6 +168,14 @@ def build_web_identity_session_policy() -> WebIdentitySessionPolicy:
)
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
@ -501,6 +509,21 @@ 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,
aws_session_tags=_canonical_aws_session_tags(auth_params.aws_session_tags),
)
def _get_aws_region_from_model_arn(self, model: str | None) -> str | None:
try:
# First check if the string contains the expected prefix
@ -1515,23 +1538,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)
aws_session_tags: Final = optional_params.pop("aws_session_tags", 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(
@ -1539,19 +1549,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,
aws_session_tags=aws_session_tags,
)
credentials: Final[Credentials] = self.resolve_credentials(auth_params, aws_region_name)
return Boto3CredentialsInfo(
credentials=credentials,
aws_region_name=aws_region_name,
@ -1685,33 +1683,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)
aws_session_tags: Final = optional_params.get("aws_session_tags", 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,
aws_session_tags=aws_session_tags,
)
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,7 +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 AwsSessionTag
from litellm.types.llms.bedrock import AwsAuthParams, AwsSessionTag
from litellm.types.utils import LiteLLMBatch
if TYPE_CHECKING:
@ -130,11 +130,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,
@ -143,6 +142,7 @@ class BedrockBatchesHandler:
aws_external_id=aws_external_id,
aws_session_tags=aws_session_tags,
)
creds: Final = BedrockBatchesConfig().resolve_credentials(auth_params, region)
client: Final = boto3.client(
"bedrock",
@ -157,16 +157,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,
aws_session_tags=aws_session_tags,
**auth_params.model_dump(),
)
try:
@ -310,19 +301,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"),
aws_session_tags=kwargs.get("aws_session_tags"),
)
creds: Final = BedrockBatchesConfig().resolve_credentials(AwsAuthParams.model_validate(kwargs), region)
client: Final = boto3.client(
"bedrock",

View file

@ -17,7 +17,7 @@ from litellm.llms.custom_httpx.http_handler import (
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
@ -323,21 +323,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)
aws_session_tags: Final = optional_params.pop("aws_session_tags", 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
@ -345,19 +332,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,
aws_session_tags=aws_session_tags,
)
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
@ -83,19 +84,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",
"aws_session_tags",
)
_BEDROCK_AWS_AUTH_PARAMETER_KEYS: Final[tuple[str, ...]] = (*AWS_AUTH_PARAM_KEYS, "aws_region_name")
def merge_bedrock_aws_request_params(
@ -1669,20 +1658,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"),
aws_session_tags=optional_params.get("aws_session_tags"),
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,19 +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)
aws_session_tags: Final = optional_params.pop("aws_session_tags", None)
### SET REGION NAME ###
if aws_region_name is None:
@ -105,21 +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,
aws_session_tags=aws_session_tags,
)
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

@ -46,7 +46,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,
@ -142,21 +142,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
@ -1157,20 +1146,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()
@ -1517,18 +1494,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

@ -29,6 +29,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
@ -257,6 +258,7 @@ class BedrockRealtime(BaseAWSLLM):
aws_sts_endpoint: str | None = None,
aws_bedrock_runtime_endpoint: str | None = None,
aws_external_id: str | None = None,
aws_session_tags: object = None,
**kwargs: object,
):
"""
@ -297,20 +299,20 @@ 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,
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
aws_session_tags=aws_session_tags,
)
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,20 +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)
aws_session_tags: Final = optional_params.pop("aws_session_tags", None)
optional_params.pop("aws_bedrock_runtime_endpoint", None)
### SET REGION NAME ###
if aws_region_name is None:
@ -53,19 +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,
aws_session_tags=aws_session_tags,
)
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,20 +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)
aws_session_tags: Final = optional_params.pop("aws_session_tags", None)
optional_params.pop("aws_bedrock_runtime_endpoint", None)
### SET REGION NAME ###
if aws_region_name is None:
@ -76,19 +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,
aws_session_tags=aws_session_tags,
)
credentials: Final[Credentials] = self.resolve_credentials(auth_params, aws_region_name)
return credentials, aws_region_name
def _prepare_request(

View file

@ -481,6 +481,7 @@ async def _arealtime(
aws_sts_endpoint: Final = kwargs.get("aws_sts_endpoint")
aws_bedrock_runtime_endpoint: Final = kwargs.get("aws_bedrock_runtime_endpoint")
aws_external_id: Final = kwargs.get("aws_external_id")
aws_session_tags: Final = kwargs.get("aws_session_tags")
await bedrock_realtime.async_realtime(
model=model,
@ -500,6 +501,7 @@ async def _arealtime(
aws_sts_endpoint=aws_sts_endpoint,
aws_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint,
aws_external_id=aws_external_id,
aws_session_tags=aws_session_tags,
)
elif _custom_llm_provider == "xai":
api_base = (

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
@ -1112,6 +1113,26 @@ class AwsSessionTag(TypedDict):
Value: str # writable-ok: boto3's STS stubs type assume_role Tags as writable TagTypeDef, which rejects ReadOnly
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_session_tags: object = 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

@ -185,19 +185,19 @@ class DummyCredentials:
],
)
@pytest.mark.parametrize(
"param_name, param_value",
"param_name, param_value, expected_credentials_value",
[
("aws_session_token", "dummy_session_token"),
("aws_session_name", "dummy_session_name"),
("aws_profile_name", "dummy_profile_name"),
("aws_role_name", "dummy_role_name"),
("aws_web_identity_token", "dummy_web_identity_token"),
("aws_sts_endpoint", "dummy_sts_endpoint"),
("aws_external_id", "dummy_external_id"),
("aws_session_tags", [{"Key": "team", "Value": "genai"}]),
("aws_session_token", "dummy_session_token", "dummy_session_token"),
("aws_session_name", "dummy_session_name", "dummy_session_name"),
("aws_profile_name", "dummy_profile_name", "dummy_profile_name"),
("aws_role_name", "dummy_role_name", "dummy_role_name"),
("aws_web_identity_token", "dummy_web_identity_token", "dummy_web_identity_token"),
("aws_sts_endpoint", "dummy_sts_endpoint", "dummy_sts_endpoint"),
("aws_external_id", "dummy_external_id", "dummy_external_id"),
("aws_session_tags", [{"Key": "team", "Value": "genai"}], ({"Key": "team", "Value": "genai"},)),
],
)
def test_dynamic_aws_params_propagation(model, param_name, param_value):
def test_dynamic_aws_params_propagation(model, param_name, param_value, expected_credentials_value):
"""
When passed to litellm.completion, each dynamic AWS authentication parameter
should propagate down to the get_credentials() call in BaseAWSLLM.
@ -282,6 +282,4 @@ def test_dynamic_aws_params_propagation(model, param_name, param_value):
)
# We now assert that get_credentials() was called with the dynamic param.
assert (
dummy_get_credentials.called_kwargs.get(param_name) == param_value
)
assert dummy_get_credentials.called_kwargs.get(param_name) == expected_credentials_value

View file

@ -2629,6 +2629,100 @@ def test_sign_s3_request_without_body_assumes_role_with_external_id(monkeypatch)
assert "ASIAFILESGETROLE" in authorization
class _SessionTagGatedSTSClient:
"""Mimics a trust policy with an aws:RequestTag condition: assume_role only succeeds with the expected tags."""
def __init__(self, expected_tags, access_key_id):
self.expected_tags = expected_tags
self.access_key_id = access_key_id
def get_caller_identity(self):
return {"Arn": "arn:aws:iam::111111111111:user/litellm-proxy-pod"}
def assume_role(self, **params):
import datetime
from botocore.exceptions import ClientError
if list(params.get("Tags") or ()) != self.expected_tags:
raise ClientError(
{"Error": {"Code": "AccessDenied", "Message": "is not authorized to perform: sts:TagSession"}},
"AssumeRole",
)
return {
"Credentials": {
"AccessKeyId": self.access_key_id,
"SecretAccessKey": "assumed-secret",
"SessionToken": "assumed-session-token",
"Expiration": datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(minutes=30),
}
}
def test_sign_s3_request_assumes_role_with_session_tags():
"""The deployment's aws_session_tags must reach STS when signing the S3 upload, not only on chat calls."""
from unittest.mock import patch
import boto3
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
expected_tags = [{"Key": "team", "Value": "genai"}]
optional_params = {
"aws_region_name": "us-east-1",
"aws_access_key_id": "AKIAFILESPUTCALLER",
"aws_secret_access_key": "pod-caller-secret",
"aws_role_name": "arn:aws:iam::999999999999:role/litellm-files-put-role",
"aws_session_name": "litellm-files-put-session",
"aws_session_tags": [{"Key": "team", "Value": "genai"}],
}
with patch.object(boto3, "client", return_value=_SessionTagGatedSTSClient(expected_tags, "ASIAFILESPUTTAGGED")):
signed_headers, _signed_body = BedrockFilesConfig()._sign_s3_request(
content='{"custom_id": "req-1"}',
api_base="https://s3.us-east-1.amazonaws.com/safe-bucket/litellm-bedrock-files-model-id-abc.jsonl",
optional_params=optional_params,
)
authorization = {key.lower(): value for key, value in signed_headers.items()}["authorization"]
assert "ASIAFILESPUTTAGGED" in authorization
def test_sign_s3_request_without_body_assumes_role_with_session_tags():
"""The deployment's aws_session_tags must reach STS when signing the S3 download too."""
from unittest.mock import patch
import boto3
from litellm.llms.bedrock.files.transformation import (
BedrockFilesConfig,
_BedrockS3RequestParams,
)
expected_tags = [{"Key": "team", "Value": "genai"}]
request_params = _BedrockS3RequestParams.model_validate(
{
"aws_region_name": "us-east-1",
"aws_access_key_id": "AKIAFILESGETCALLER",
"aws_secret_access_key": "pod-caller-secret",
"aws_role_name": "arn:aws:iam::999999999999:role/litellm-files-get-role",
"aws_session_name": "litellm-files-get-session",
"aws_session_tags": [{"Key": "team", "Value": "genai"}],
}
)
with patch.object(boto3, "client", return_value=_SessionTagGatedSTSClient(expected_tags, "ASIAFILESGETTAGGED")):
signed_headers = BedrockFilesConfig()._sign_s3_request_without_body(
method="GET",
api_base="https://s3.us-east-1.amazonaws.com/safe-bucket/litellm-bedrock-files-model-id-abc.jsonl",
aws_region_name="us-east-1",
request_params=request_params,
)
authorization = {key.lower(): value for key, value in signed_headers.items()}["authorization"]
assert "ASIAFILESGETTAGGED" in authorization
def _s3_signature_for(method: str, url: str, headers: Mapping[str, str]) -> str:
sent = {name.lower(): value for name, value in headers.items()}
signed_names = sent["authorization"].split("SignedHeaders=")[1].split(",")[0].split(";")

View file

@ -855,6 +855,7 @@ class TestBedrockRealtimeAwsAuth:
aws_role_name="arn:aws:iam::123456789012:role/nova-sonic",
aws_session_name="realtime-session",
aws_external_id="realtime-external-id",
aws_session_tags=[{"Key": "team", "Value": "realtime"}],
)
assert handler.get_credentials_kwargs == {
@ -868,6 +869,7 @@ class TestBedrockRealtimeAwsAuth:
"aws_web_identity_token": None,
"aws_sts_endpoint": None,
"aws_external_id": "realtime-external-id",
"aws_session_tags": ({"Key": "team", "Value": "realtime"},),
}
resolver = stub_aws_sdk_client["config_kwargs"]["aws_credentials_identity_resolver"]
assert isinstance(resolver, FakeStaticCredentialsResolver)

View file

@ -10,6 +10,7 @@ from fastapi.testclient import TestClient
from collections.abc import Callable
from datetime import datetime, timedelta, timezone
from typing import Any, Dict, Optional
from unittest.mock import MagicMock, patch
@ -3555,3 +3556,148 @@ 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, dict[str, object]]) -> Callable[..., MagicMock]:
"""boto3.client replacement that records the STS client kwargs and the assume-role params."""
def _client(service_name: str, **client_kwargs: object) -> MagicMock:
recorded["client_kwargs"] = client_kwargs
sts = MagicMock()
def _assume(**params: object) -> dict[str, object]:
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: object) -> dict[str, object]:
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",
aws_session_tags=[{"Key": "team", "Value": "genai"}, {"Key": "cost-center", "Value": "42"}],
)
recorded: dict[str, dict[str, object]] = {}
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 recorded["assume_role"]["Tags"] == (
{"Key": "cost-center", "Value": "42"},
{"Key": "team", "Value": "genai"},
)
assert credentials.access_key == "ASIAASSUMED"
@pytest.mark.parametrize(
"malformed_tags",
[
"team=genai",
{"team": "genai"},
[{"key": "team", "value": "genai"}],
[{"Key": "team"}],
],
)
def test_resolve_credentials_rejects_malformed_session_tags(malformed_tags):
"""A struct built from raw config must surface the friendly session-tag error before STS is called."""
from litellm.types.llms.bedrock import AwsAuthParams
auth_params = AwsAuthParams(
aws_role_name="arn:aws:iam::123456789012:role/litellm-target",
aws_session_name="litellm-session",
aws_session_tags=malformed_tags,
)
recorded: dict[str, dict[str, object]] = {}
with (
patch.dict(os.environ, _os_environ_without_aws_keys(), clear=True),
patch("boto3.client", side_effect=_recording_boto3_client(recorded)),
):
with pytest.raises(ValueError, match="Invalid 'aws_session_tags' value"):
BaseAWSLLM().resolve_credentials(auth_params, "us-east-1")
assert "assume_role" not in recorded
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, dict[str, object]] = {}
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"