mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-21 00:21:49 +00:00
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:
commit
a2626726a2
17 changed files with 396 additions and 286 deletions
|
|
@ -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 {}
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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 ###
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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=(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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 = (
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(";")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
46
tests/test_litellm/types/llms/test_types_llms_bedrock.py
Normal file
46
tests/test_litellm/types/llms/test_types_llms_bedrock.py
Normal 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"
|
||||
Loading…
Add table
Reference in a new issue