From 8ad1ae73e5076dd00128ef0814bfe89a75c85668 Mon Sep 17 00:00:00 2001 From: Brian Schultheiss Date: Sun, 23 Jun 2024 12:51:25 -0700 Subject: [PATCH 1/9] Support aws_session_token for bedrock client. https://github.com/BerriAI/litellm/issues/4346 --- litellm/llms/bedrock_httpx.py | 34 +++++++++++++ litellm/main.py | 94 +++++++++++++---------------------- 2 files changed, 68 insertions(+), 60 deletions(-) diff --git a/litellm/llms/bedrock_httpx.py b/litellm/llms/bedrock_httpx.py index 84ab10907c..3eb8acd390 100644 --- a/litellm/llms/bedrock_httpx.py +++ b/litellm/llms/bedrock_httpx.py @@ -301,6 +301,7 @@ class BedrockLLM(BaseLLM): self, aws_access_key_id: Optional[str] = None, aws_secret_access_key: Optional[str] = None, + aws_session_token: Optional[str] = None, aws_region_name: Optional[str] = None, aws_session_name: Optional[str] = None, aws_profile_name: Optional[str] = None, @@ -316,6 +317,7 @@ class BedrockLLM(BaseLLM): params_to_check: List[Optional[str]] = [ aws_access_key_id, aws_secret_access_key, + aws_session_token, aws_region_name, aws_session_name, aws_profile_name, @@ -333,6 +335,7 @@ class BedrockLLM(BaseLLM): ( aws_access_key_id, aws_secret_access_key, + aws_session_token, aws_region_name, aws_session_name, aws_profile_name, @@ -426,6 +429,18 @@ class BedrockLLM(BaseLLM): client = boto3.Session(profile_name=aws_profile_name) return client.get_credentials() + elif ( + aws_access_key_id is not None + and aws_secret_access_key is not None + and aws_session_token is not None + ): ### CHECK FOR AWS SESSION TOKEN ### + from botocore.credentials import Credentials + credentials = Credentials( + access_key=aws_access_key_id, + secret_key=aws_secret_access_key, + token=aws_session_token, + ) + return credentials else: session = boto3.Session( aws_access_key_id=aws_access_key_id, @@ -733,6 +748,7 @@ class BedrockLLM(BaseLLM): # pop aws_secret_access_key, aws_access_key_id, aws_region_name from kwargs, since completion calls fail with them aws_secret_access_key = optional_params.pop("aws_secret_access_key", None) aws_access_key_id = optional_params.pop("aws_access_key_id", None) + aws_session_token = optional_params.pop("aws_session_token", None) aws_region_name = optional_params.pop("aws_region_name", None) aws_role_name = optional_params.pop("aws_role_name", None) aws_session_name = optional_params.pop("aws_session_name", None) @@ -764,6 +780,7 @@ class BedrockLLM(BaseLLM): credentials: 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, @@ -1418,6 +1435,7 @@ class BedrockConverseLLM(BaseLLM): self, aws_access_key_id: Optional[str] = None, aws_secret_access_key: Optional[str] = None, + aws_session_token: Optional[str] = None, aws_region_name: Optional[str] = None, aws_session_name: Optional[str] = None, aws_profile_name: Optional[str] = None, @@ -1433,6 +1451,7 @@ class BedrockConverseLLM(BaseLLM): params_to_check: List[Optional[str]] = [ aws_access_key_id, aws_secret_access_key, + aws_session_token, aws_region_name, aws_session_name, aws_profile_name, @@ -1450,6 +1469,7 @@ class BedrockConverseLLM(BaseLLM): ( aws_access_key_id, aws_secret_access_key, + aws_session_token, aws_region_name, aws_session_name, aws_profile_name, @@ -1543,6 +1563,18 @@ class BedrockConverseLLM(BaseLLM): client = boto3.Session(profile_name=aws_profile_name) return client.get_credentials() + elif ( + aws_access_key_id is not None + and aws_secret_access_key is not None + and aws_session_token is not None + ): ### CHECK FOR AWS SESSION TOKEN ### + from botocore.credentials import Credentials + credentials = Credentials( + access_key=aws_access_key_id, + secret_key=aws_secret_access_key, + token=aws_session_token, + ) + return credentials else: session = boto3.Session( aws_access_key_id=aws_access_key_id, @@ -1678,6 +1710,7 @@ class BedrockConverseLLM(BaseLLM): # pop aws_secret_access_key, aws_access_key_id, aws_region_name from kwargs, since completion calls fail with them aws_secret_access_key = optional_params.pop("aws_secret_access_key", None) aws_access_key_id = optional_params.pop("aws_access_key_id", None) + aws_session_token = optional_params.pop("aws_session_token", None) aws_region_name = optional_params.pop("aws_region_name", None) aws_role_name = optional_params.pop("aws_role_name", None) aws_session_name = optional_params.pop("aws_session_name", None) @@ -1709,6 +1742,7 @@ class BedrockConverseLLM(BaseLLM): credentials: 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, diff --git a/litellm/main.py b/litellm/main.py index 307659c8a2..92cd2fee11 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -2176,13 +2176,23 @@ def completion( # boto3 reads keys from .env custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict - if ( - "aws_bedrock_client" in optional_params - ): # use old bedrock flow for aws_bedrock_client users. - response = bedrock.completion( + + if ("aws_bedrock_client" in optional_params): + # Extract credentials for legacy boto3 client and pass thru to httpx + aws_bedrock_client = optional_params.pop("aws_bedrock_client") + creds = aws_bedrock_client._get_credentials().get_frozen_credentials() + if creds.access_key: + optional_params["aws_access_key_id"] = creds.access_key + if creds.secret_key: + optional_params["aws_secret_access_key"] = creds.secret_key + if creds.token: + optional_params["aws_session_token"] = creds.token + + if model.startswith("anthropic"): + response = bedrock_converse_chat_completion.completion( model=model, messages=messages, - custom_prompt_dict=litellm.custom_prompt_dict, + custom_prompt_dict=custom_prompt_dict, model_response=model_response, print_verbose=print_verbose, optional_params=optional_params, @@ -2192,63 +2202,27 @@ def completion( logging_obj=logging, extra_headers=extra_headers, timeout=timeout, + acompletion=acompletion, + client=client, + ) + else: + response = bedrock_chat_completion.completion( + model=model, + messages=messages, + custom_prompt_dict=custom_prompt_dict, + model_response=model_response, + print_verbose=print_verbose, + optional_params=optional_params, + litellm_params=litellm_params, + logger_fn=logger_fn, + encoding=encoding, + logging_obj=logging, + extra_headers=extra_headers, + timeout=timeout, + acompletion=acompletion, + client=client, ) - if ( - "stream" in optional_params - and optional_params["stream"] == True - and not isinstance(response, CustomStreamWrapper) - ): - # don't try to access stream object, - if "ai21" in model: - response = CustomStreamWrapper( - response, - model, - custom_llm_provider="bedrock", - logging_obj=logging, - ) - else: - response = CustomStreamWrapper( - iter(response), - model, - custom_llm_provider="bedrock", - logging_obj=logging, - ) - else: - if model.startswith("anthropic"): - response = bedrock_converse_chat_completion.completion( - model=model, - messages=messages, - custom_prompt_dict=custom_prompt_dict, - model_response=model_response, - print_verbose=print_verbose, - optional_params=optional_params, - litellm_params=litellm_params, - logger_fn=logger_fn, - encoding=encoding, - logging_obj=logging, - extra_headers=extra_headers, - timeout=timeout, - acompletion=acompletion, - client=client, - ) - else: - response = bedrock_chat_completion.completion( - model=model, - messages=messages, - custom_prompt_dict=custom_prompt_dict, - model_response=model_response, - print_verbose=print_verbose, - optional_params=optional_params, - litellm_params=litellm_params, - logger_fn=logger_fn, - encoding=encoding, - logging_obj=logging, - extra_headers=extra_headers, - timeout=timeout, - acompletion=acompletion, - client=client, - ) if optional_params.get("stream", False): ## LOGGING logging.post_call( From 7f91e5354886fa1a445d953e2a368d507c3c0c3f Mon Sep 17 00:00:00 2001 From: Brian Schultheiss Date: Sun, 23 Jun 2024 13:15:04 -0700 Subject: [PATCH 2/9] updated documentation to reference boto3.client credential extraction, and update boto3.client creation to support session_token. --- docs/my-website/docs/providers/bedrock.md | 5 +++++ litellm/llms/bedrock.py | 25 ++++++++++++++++++++++- 2 files changed, 29 insertions(+), 1 deletion(-) diff --git a/docs/my-website/docs/providers/bedrock.md b/docs/my-website/docs/providers/bedrock.md index f380a6a50e..adbc54caa7 100644 --- a/docs/my-website/docs/providers/bedrock.md +++ b/docs/my-website/docs/providers/bedrock.md @@ -476,6 +476,7 @@ response = completion( messages=[{ "content": "Hello, how are you?","role": "user"}], aws_access_key_id="", aws_secret_access_key="", + aws_session_token="", aws_region_name="", ) ``` @@ -549,6 +550,10 @@ response = completion( This is a deprecated flow. Boto3 is not async. And boto3.client does not let us make the http call through httpx. Pass in your aws params through the method above 👆. [See Auth Code](https://github.com/BerriAI/litellm/blob/55a20c7cce99a93d36a82bf3ae90ba3baf9a7f89/litellm/llms/bedrock_httpx.py#L284) [Add new auth flow](https://github.com/BerriAI/litellm/issues) + +Experimental - 2024-Jun-23: + aws_access_key_id, aws_secret_access_key=, and aws_session_token will be extracted from boto3.client and be passed onto the httpx client + ::: Pass an external BedrockRuntime.Client object as a parameter to litellm.completion. Useful when using an AWS credentials profile, SSO session, assumed role session, or if environment variables are not available for auth. diff --git a/litellm/llms/bedrock.py b/litellm/llms/bedrock.py index d0d3bef6da..6c941bb558 100644 --- a/litellm/llms/bedrock.py +++ b/litellm/llms/bedrock.py @@ -558,6 +558,7 @@ def init_bedrock_client( region_name=None, aws_access_key_id: Optional[str] = None, aws_secret_access_key: Optional[str] = None, + aws_session_token: Optional[str] = None, aws_region_name: Optional[str] = None, aws_bedrock_runtime_endpoint: Optional[str] = None, aws_session_name: Optional[str] = None, @@ -591,6 +592,7 @@ def init_bedrock_client( ( aws_access_key_id, aws_secret_access_key, + aws_session_token, aws_region_name, aws_bedrock_runtime_endpoint, aws_session_name, @@ -668,6 +670,21 @@ def init_bedrock_client( endpoint_url=endpoint_url, config=config, ) + elif ( + aws_access_key_id is not None + and aws_secret_access_key is not None + and aws_session_token is not None + ): ### CHECK FOR AWS SESSION TOKEN ### + client = boto3.client( + service_name="bedrock-runtime", + aws_access_key_id=aws_access_key_id, + aws_secret_access_key=aws_secret_access_key, + aws_session_token=aws_session_token, + region_name=region_name, + endpoint_url=endpoint_url, + config=config, + ) + elif aws_role_name is not None and aws_session_name is not None: # use sts if role name passed in sts_client = boto3.client( @@ -786,9 +803,10 @@ def completion( _is_function_call = False json_schemas: dict = {} try: - # pop aws_secret_access_key, aws_access_key_id, aws_region_name from kwargs, since completion calls fail with them + # 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 = optional_params.pop("aws_secret_access_key", None) aws_access_key_id = optional_params.pop("aws_access_key_id", None) + aws_session_token = optional_params.pop("aws_session_token", None) aws_region_name = optional_params.pop("aws_region_name", None) aws_role_name = optional_params.pop("aws_role_name", None) aws_session_name = optional_params.pop("aws_session_name", None) @@ -806,6 +824,7 @@ def completion( client = init_bedrock_client( 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_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint, aws_role_name=aws_role_name, @@ -1328,6 +1347,7 @@ def embedding( # pop aws_secret_access_key, aws_access_key_id, aws_region_name from kwargs, since completion calls fail with them aws_secret_access_key = optional_params.pop("aws_secret_access_key", None) aws_access_key_id = optional_params.pop("aws_access_key_id", None) + aws_session_token = optional_params.pop("aws_session_token", None) aws_region_name = optional_params.pop("aws_region_name", None) aws_role_name = optional_params.pop("aws_role_name", None) aws_session_name = optional_params.pop("aws_session_name", None) @@ -1340,6 +1360,7 @@ def embedding( client = init_bedrock_client( 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_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint, aws_web_identity_token=aws_web_identity_token, @@ -1419,6 +1440,7 @@ def image_generation( # pop aws_secret_access_key, aws_access_key_id, aws_region_name from kwargs, since completion calls fail with them aws_secret_access_key = optional_params.pop("aws_secret_access_key", None) aws_access_key_id = optional_params.pop("aws_access_key_id", None) + aws_session_token = optional_params.pop("aws_session_token", None) aws_region_name = optional_params.pop("aws_region_name", None) aws_role_name = optional_params.pop("aws_role_name", None) aws_session_name = optional_params.pop("aws_session_name", None) @@ -1431,6 +1453,7 @@ def image_generation( client = init_bedrock_client( 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_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint, aws_web_identity_token=aws_web_identity_token, From 3fbb25f903808ebae31924d0c72a5da000019e1c Mon Sep 17 00:00:00 2001 From: Brian Schultheiss Date: Sun, 23 Jun 2024 13:37:38 -0700 Subject: [PATCH 3/9] Updated more references to AWS session token --- docs/my-website/docs/providers/aws_sagemaker.md | 1 + docs/my-website/docs/providers/bedrock.md | 3 ++- litellm/llms/bedrock.py | 5 +++-- litellm/llms/bedrock_httpx.py | 2 +- litellm/llms/sagemaker.py | 14 ++++++++++++-- litellm/proxy/proxy_server.py | 2 ++ litellm/types/router.py | 4 ++++ 7 files changed, 25 insertions(+), 6 deletions(-) diff --git a/docs/my-website/docs/providers/aws_sagemaker.md b/docs/my-website/docs/providers/aws_sagemaker.md index 2b65709e8e..5793fb05ae 100644 --- a/docs/my-website/docs/providers/aws_sagemaker.md +++ b/docs/my-website/docs/providers/aws_sagemaker.md @@ -59,6 +59,7 @@ response = completion( messages=[{ "content": "Hello, how are you?","role": "user"}], aws_access_key_id="", aws_secret_access_key="", + aws_session_token="", aws_region_name="", ) ``` diff --git a/docs/my-website/docs/providers/bedrock.md b/docs/my-website/docs/providers/bedrock.md index adbc54caa7..7f9b21b96b 100644 --- a/docs/my-website/docs/providers/bedrock.md +++ b/docs/my-website/docs/providers/bedrock.md @@ -538,6 +538,7 @@ response = completion( aws_region_name=aws_region_name, aws_access_key_id=aws_access_key_id, aws_secret_access_key=aws_secret_access_key, + aws_session_token=aws_session_token, aws_role_name=aws_role_name, aws_session_name="my-test-session", ) @@ -553,7 +554,7 @@ This is a deprecated flow. Boto3 is not async. And boto3.client does not let us Experimental - 2024-Jun-23: aws_access_key_id, aws_secret_access_key=, and aws_session_token will be extracted from boto3.client and be passed onto the httpx client - + ::: Pass an external BedrockRuntime.Client object as a parameter to litellm.completion. Useful when using an AWS credentials profile, SSO session, assumed role session, or if environment variables are not available for auth. diff --git a/litellm/llms/bedrock.py b/litellm/llms/bedrock.py index 6c941bb558..2403edf814 100644 --- a/litellm/llms/bedrock.py +++ b/litellm/llms/bedrock.py @@ -576,6 +576,7 @@ def init_bedrock_client( params_to_check = [ aws_access_key_id, aws_secret_access_key, + aws_session_token, aws_region_name, aws_bedrock_runtime_endpoint, aws_session_name, @@ -1344,7 +1345,7 @@ def embedding( encoding=None, ): ### BOTO3 INIT ### - # pop aws_secret_access_key, aws_access_key_id, aws_region_name from kwargs, since completion calls fail with them + # 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 = optional_params.pop("aws_secret_access_key", None) aws_access_key_id = optional_params.pop("aws_access_key_id", None) aws_session_token = optional_params.pop("aws_session_token", None) @@ -1437,7 +1438,7 @@ def image_generation( Bedrock Image Gen endpoint support """ ### BOTO3 INIT ### - # pop aws_secret_access_key, aws_access_key_id, aws_region_name from kwargs, since completion calls fail with them + # 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 = optional_params.pop("aws_secret_access_key", None) aws_access_key_id = optional_params.pop("aws_access_key_id", None) aws_session_token = optional_params.pop("aws_session_token", None) diff --git a/litellm/llms/bedrock_httpx.py b/litellm/llms/bedrock_httpx.py index 3eb8acd390..d00695f870 100644 --- a/litellm/llms/bedrock_httpx.py +++ b/litellm/llms/bedrock_httpx.py @@ -745,7 +745,7 @@ class BedrockLLM(BaseLLM): provider = model.split(".")[0] ## CREDENTIALS ## - # pop aws_secret_access_key, aws_access_key_id, aws_region_name from kwargs, since completion calls fail with them + # 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 = optional_params.pop("aws_secret_access_key", None) aws_access_key_id = optional_params.pop("aws_access_key_id", None) aws_session_token = optional_params.pop("aws_session_token", None) diff --git a/litellm/llms/sagemaker.py b/litellm/llms/sagemaker.py index 8e75428bb7..7d639b7bb2 100644 --- a/litellm/llms/sagemaker.py +++ b/litellm/llms/sagemaker.py @@ -162,9 +162,10 @@ def completion( ): import boto3 - # pop aws_secret_access_key, aws_access_key_id, aws_region_name from kwargs, since completion calls fail with them + # 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 = optional_params.pop("aws_secret_access_key", None) aws_access_key_id = optional_params.pop("aws_access_key_id", None) + aws_session_token = optional_params.pop("aws_session_token", None) aws_region_name = optional_params.pop("aws_region_name", None) model_id = optional_params.pop("model_id", None) @@ -175,6 +176,7 @@ def completion( service_name="sagemaker-runtime", aws_access_key_id=aws_access_key_id, aws_secret_access_key=aws_secret_access_key, + aws_session_token=aws_session_token, region_name=aws_region_name, ) else: @@ -249,6 +251,7 @@ def completion( model_id=model_id, aws_secret_access_key=aws_secret_access_key, aws_access_key_id=aws_access_key_id, + aws_session_token=aws_session_token, aws_region_name=aws_region_name, ) return response @@ -281,6 +284,7 @@ def completion( model_id=model_id, aws_secret_access_key=aws_secret_access_key, aws_access_key_id=aws_access_key_id, + aws_session_token=aws_session_token, aws_region_name=aws_region_name, ) data = json.dumps({"inputs": prompt, "parameters": inference_params}).encode( @@ -414,6 +418,7 @@ async def async_streaming( aws_secret_access_key: Optional[str], aws_access_key_id: Optional[str], aws_region_name: Optional[str], + aws_session_token: Optional[str] = None, ): """ Use aioboto3 @@ -429,6 +434,7 @@ async def async_streaming( service_name="sagemaker-runtime", aws_access_key_id=aws_access_key_id, aws_secret_access_key=aws_secret_access_key, + aws_session_token=aws_session_token, region_name=aws_region_name, ) else: @@ -481,6 +487,7 @@ async def async_completion( aws_secret_access_key: Optional[str], aws_access_key_id: Optional[str], aws_region_name: Optional[str], + aws_session_token: Optional[str] = None, ): """ Use aioboto3 @@ -496,6 +503,7 @@ async def async_completion( service_name="sagemaker-runtime", aws_access_key_id=aws_access_key_id, aws_secret_access_key=aws_secret_access_key, + aws_session_token=aws_session_token, region_name=aws_region_name, ) else: @@ -639,9 +647,10 @@ def embedding( ### BOTO3 INIT import boto3 - # pop aws_secret_access_key, aws_access_key_id, aws_region_name from kwargs, since completion calls fail with them + # 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 = optional_params.pop("aws_secret_access_key", None) aws_access_key_id = optional_params.pop("aws_access_key_id", None) + aws_session_token = optional_params.pop("aws_session_token", None) aws_region_name = optional_params.pop("aws_region_name", None) if aws_access_key_id is not None: @@ -651,6 +660,7 @@ def embedding( service_name="sagemaker-runtime", aws_access_key_id=aws_access_key_id, aws_secret_access_key=aws_secret_access_key, + aws_session_token=aws_session_token, region_name=aws_region_name, ) else: diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 30b90abe64..9c1039f51e 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -6343,6 +6343,7 @@ async def model_info_v2( _model["litellm_params"].pop("vertex_credentials", None) _model["litellm_params"].pop("aws_access_key_id", None) _model["litellm_params"].pop("aws_secret_access_key", None) + _model["litellm_params"].pop("aws_session_token", None) verbose_proxy_logger.debug("all_models: %s", all_models) return {"data": all_models} @@ -6859,6 +6860,7 @@ async def model_info_v1( model["litellm_params"].pop("vertex_credentials", None) model["litellm_params"].pop("aws_access_key_id", None) model["litellm_params"].pop("aws_secret_access_key", None) + model["litellm_params"].pop("aws_session_token", None) verbose_proxy_logger.debug("all_models: %s", all_models) return {"data": all_models} diff --git a/litellm/types/router.py b/litellm/types/router.py index e6864ffe2e..059a8620e5 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -145,6 +145,7 @@ class GenericLiteLLMParams(BaseModel): ## AWS BEDROCK / SAGEMAKER ## aws_access_key_id: Optional[str] = None aws_secret_access_key: Optional[str] = None + aws_session_token: Optional[str] = None aws_region_name: Optional[str] = None ## IBM WATSONX ## watsonx_region_name: Optional[str] = None @@ -178,6 +179,7 @@ class GenericLiteLLMParams(BaseModel): ## AWS BEDROCK / SAGEMAKER ## aws_access_key_id: Optional[str] = None, aws_secret_access_key: Optional[str] = None, + aws_session_token: Optional[str] = None, aws_region_name: Optional[str] = None, ## IBM WATSONX ## watsonx_region_name: Optional[str] = None, @@ -242,6 +244,7 @@ class LiteLLM_Params(GenericLiteLLMParams): ## AWS BEDROCK / SAGEMAKER ## aws_access_key_id: Optional[str] = None, aws_secret_access_key: Optional[str] = None, + aws_session_token: Optional[str] = None, aws_region_name: Optional[str] = None, **params, ): @@ -307,6 +310,7 @@ class LiteLLMParamsTypedDict(TypedDict, total=False): ## AWS BEDROCK / SAGEMAKER ## aws_access_key_id: Optional[str] aws_secret_access_key: Optional[str] + aws_session_token: Optional[str] aws_region_name: Optional[str] ## IBM WATSONX ## watsonx_region_name: Optional[str] From 5a6588342cab071322737d69043a3731e34d0468 Mon Sep 17 00:00:00 2001 From: Brian Schultheiss Date: Sun, 23 Jun 2024 15:19:54 -0700 Subject: [PATCH 4/9] added test for change --- docs/my-website/docs/providers/bedrock.md | 8 +- litellm/tests/test_bedrock_completion.py | 181 ++++++++++++++++++++++ 2 files changed, 185 insertions(+), 4 deletions(-) diff --git a/docs/my-website/docs/providers/bedrock.md b/docs/my-website/docs/providers/bedrock.md index 7f9b21b96b..1b073ad25e 100644 --- a/docs/my-website/docs/providers/bedrock.md +++ b/docs/my-website/docs/providers/bedrock.md @@ -474,10 +474,10 @@ from litellm import completion response = completion( model="bedrock/anthropic.claude-instant-v1", messages=[{ "content": "Hello, how are you?","role": "user"}], + aws_region_name="", aws_access_key_id="", aws_secret_access_key="", - aws_session_token="", - aws_region_name="", + aws_session_token=None, ) ``` @@ -553,7 +553,7 @@ This is a deprecated flow. Boto3 is not async. And boto3.client does not let us Experimental - 2024-Jun-23: - aws_access_key_id, aws_secret_access_key=, and aws_session_token will be extracted from boto3.client and be passed onto the httpx client + `aws_access_key_id`, `aws_secret_access_key`, and `aws_session_token` will be extracted from boto3.client and be passed onto the httpx client ::: @@ -569,7 +569,7 @@ bedrock = boto3.client( region_name="us-east-1", aws_access_key_id="", aws_secret_access_key="", - aws_session_token="", + aws_session_token=None, ) response = completion( diff --git a/litellm/tests/test_bedrock_completion.py b/litellm/tests/test_bedrock_completion.py index b953ca2a3a..614d1b76fa 100644 --- a/litellm/tests/test_bedrock_completion.py +++ b/litellm/tests/test_bedrock_completion.py @@ -15,6 +15,7 @@ from litellm import embedding, completion, completion_cost, Timeout, ModelRespon from litellm import RateLimitError from litellm.llms.custom_httpx.http_handler import HTTPHandler, AsyncHTTPHandler from unittest.mock import patch, AsyncMock, Mock +from litellm.llms.bedrock_httpx import BedrockLLM # litellm.num_retries = 3 litellm.cache = None @@ -205,6 +206,186 @@ def test_completion_bedrock_claude_sts_client_auth(): except Exception as e: pytest.fail(f"Error occurred: {e}") +@pytest.fixture() +def bedrock_session_token_creds(): + print("\ncalling oidc auto to get aws_session_token credentials") + import os + + aws_region_name = os.environ["AWS_REGION_NAME"] + aws_session_token = os.environ.get("AWS_SESSION_TOKEN") + + bllm = BedrockLLM() + if aws_session_token is not None: + # For local testing + creds = bllm.get_credentials( + aws_region_name=aws_region_name, + aws_access_key_id=os.environ['AWS_ACCESS_KEY_ID'], + aws_secret_access_key=os.environ['AWS_SECRET_ACCESS_KEY'], + aws_session_token=aws_session_token + ) + else: + # For circle-ci testing + # aws_role_name = os.environ["AWS_TEMP_ROLE_NAME"] + # TODO: This is using ai.moda's IAM role, we should use LiteLLM's IAM role eventually + aws_role_name = "arn:aws:iam::335785316107:role/litellm-github-unit-tests-circleci" + aws_web_identity_token = "oidc/circleci_v2/" + + creds = bllm.get_credentials( + aws_region_name=aws_region_name, + aws_web_identity_token=aws_web_identity_token, + aws_role_name=aws_role_name, + aws_session_name="my-test-session", + ) + return creds + +@pytest.mark.skipif( + os.environ.get("CIRCLE_OIDC_TOKEN_V2") is None, + reason="Cannot run without being in CircleCI Runner", +) +def test_completion_bedrock_claude_aws_session_token(bedrock_session_token_creds): + print("\ncalling bedrock claude with aws_session_token auth") + + import os + aws_region_name = os.environ["AWS_REGION_NAME"] + aws_access_key_id = bedrock_session_token_creds.access_key + aws_secret_access_key = bedrock_session_token_creds.secret_key + aws_session_token = bedrock_session_token_creds.token + + try: + litellm.set_verbose = True + + response_1 = completion( + model="bedrock/anthropic.claude-3-haiku-20240307-v1:0", + messages=messages, + max_tokens=10, + temperature=0.1, + aws_region_name=aws_region_name, + aws_access_key_id=aws_access_key_id, + aws_secret_access_key=aws_secret_access_key, + aws_session_token=aws_session_token, + ) + print(response_1) + assert len(response_1.choices) > 0 + assert len(response_1.choices[0].message.content) > 0 + + # This second call is to verify that the cache isn't breaking anything + response_2 = completion( + model="bedrock/anthropic.claude-3-haiku-20240307-v1:0", + messages=messages, + max_tokens=5, + temperature=0.2, + aws_region_name=aws_region_name, + aws_access_key_id=aws_access_key_id, + aws_secret_access_key=aws_secret_access_key, + aws_session_token=aws_session_token, + ) + print(response_2) + assert len(response_2.choices) > 0 + assert len(response_2.choices[0].message.content) > 0 + + # This third call is to verify that the cache isn't used for a different region + response_3 = completion( + model="bedrock/anthropic.claude-3-haiku-20240307-v1:0", + messages=messages, + max_tokens=6, + temperature=0.3, + aws_region_name="us-east-1", + aws_access_key_id=aws_access_key_id, + aws_secret_access_key=aws_secret_access_key, + aws_session_token=aws_session_token, + ) + print(response_3) + assert len(response_3.choices) > 0 + assert len(response_3.choices[0].message.content) > 0 + + except RateLimitError: + pass + except Exception as e: + pytest.fail(f"Error occurred: {e}") + +@pytest.mark.skipif( + os.environ.get("CIRCLE_OIDC_TOKEN_V2") is None, + reason="Cannot run without being in CircleCI Runner", +) +def test_completion_bedrock_claude_aws_bedrock_client(bedrock_session_token_creds): + print("\ncalling bedrock claude with aws_session_token auth") + + import os + import boto3 + from botocore.client import Config + + aws_region_name = os.environ["AWS_REGION_NAME"] + aws_access_key_id = bedrock_session_token_creds.access_key + aws_secret_access_key = bedrock_session_token_creds.secret_key + aws_session_token = bedrock_session_token_creds.token + + aws_bedrock_client_west = boto3.client( + service_name="bedrock-runtime", + region_name=aws_region_name, + aws_access_key_id=aws_access_key_id, + aws_secret_access_key=aws_secret_access_key, + aws_session_token=aws_session_token, + config= Config( + read_timeout=600 + ) + ) + + + try: + litellm.set_verbose = True + + response_1 = completion( + model="bedrock/anthropic.claude-3-haiku-20240307-v1:0", + messages=messages, + max_tokens=10, + temperature=0.1, + aws_bedrock_client=aws_bedrock_client_west, + ) + print(response_1) + assert len(response_1.choices) > 0 + assert len(response_1.choices[0].message.content) > 0 + + # This second call is to verify that the cache isn't breaking anything + response_2 = completion( + model="bedrock/anthropic.claude-3-haiku-20240307-v1:0", + messages=messages, + max_tokens=5, + temperature=0.2, + aws_bedrock_client=aws_bedrock_client_west, + ) + print(response_2) + assert len(response_2.choices) > 0 + assert len(response_2.choices[0].message.content) > 0 + + # This third call is to verify that the cache isn't used for a different region + aws_bedrock_client_east = boto3.client( + service_name="bedrock-runtime", + region_name="us-east-1", + aws_access_key_id=aws_access_key_id, + aws_secret_access_key=aws_secret_access_key, + aws_session_token=aws_session_token, + config= Config( + read_timeout=600 + ) + ) + + response_3 = completion( + model="bedrock/anthropic.claude-3-haiku-20240307-v1:0", + messages=messages, + max_tokens=6, + temperature=0.3, + aws_bedrock_client=aws_bedrock_client_east, + ) + print(response_3) + assert len(response_3.choices) > 0 + assert len(response_3.choices[0].message.content) > 0 + + except RateLimitError: + pass + except Exception as e: + pytest.fail(f"Error occurred: {e}") + + # test_completion_bedrock_claude_sts_client_auth() From 80b4af7abec2393288aaa4cc01584b34a5d0e6db Mon Sep 17 00:00:00 2001 From: Brian Schultheiss Date: Tue, 25 Jun 2024 13:29:33 -0700 Subject: [PATCH 5/9] Revert some non-essential changes --- docs/my-website/docs/providers/aws_sagemaker.md | 1 - litellm/llms/sagemaker.py | 14 ++------------ litellm/proxy/proxy_server.py | 2 -- litellm/types/router.py | 4 ---- 4 files changed, 2 insertions(+), 19 deletions(-) diff --git a/docs/my-website/docs/providers/aws_sagemaker.md b/docs/my-website/docs/providers/aws_sagemaker.md index 5793fb05ae..2b65709e8e 100644 --- a/docs/my-website/docs/providers/aws_sagemaker.md +++ b/docs/my-website/docs/providers/aws_sagemaker.md @@ -59,7 +59,6 @@ response = completion( messages=[{ "content": "Hello, how are you?","role": "user"}], aws_access_key_id="", aws_secret_access_key="", - aws_session_token="", aws_region_name="", ) ``` diff --git a/litellm/llms/sagemaker.py b/litellm/llms/sagemaker.py index 7d639b7bb2..8e75428bb7 100644 --- a/litellm/llms/sagemaker.py +++ b/litellm/llms/sagemaker.py @@ -162,10 +162,9 @@ def completion( ): import boto3 - # pop aws_secret_access_key, aws_access_key_id, aws_session_token, aws_region_name from kwargs, since completion calls fail with them + # pop aws_secret_access_key, aws_access_key_id, aws_region_name from kwargs, since completion calls fail with them aws_secret_access_key = optional_params.pop("aws_secret_access_key", None) aws_access_key_id = optional_params.pop("aws_access_key_id", None) - aws_session_token = optional_params.pop("aws_session_token", None) aws_region_name = optional_params.pop("aws_region_name", None) model_id = optional_params.pop("model_id", None) @@ -176,7 +175,6 @@ def completion( service_name="sagemaker-runtime", aws_access_key_id=aws_access_key_id, aws_secret_access_key=aws_secret_access_key, - aws_session_token=aws_session_token, region_name=aws_region_name, ) else: @@ -251,7 +249,6 @@ def completion( model_id=model_id, aws_secret_access_key=aws_secret_access_key, aws_access_key_id=aws_access_key_id, - aws_session_token=aws_session_token, aws_region_name=aws_region_name, ) return response @@ -284,7 +281,6 @@ def completion( model_id=model_id, aws_secret_access_key=aws_secret_access_key, aws_access_key_id=aws_access_key_id, - aws_session_token=aws_session_token, aws_region_name=aws_region_name, ) data = json.dumps({"inputs": prompt, "parameters": inference_params}).encode( @@ -418,7 +414,6 @@ async def async_streaming( aws_secret_access_key: Optional[str], aws_access_key_id: Optional[str], aws_region_name: Optional[str], - aws_session_token: Optional[str] = None, ): """ Use aioboto3 @@ -434,7 +429,6 @@ async def async_streaming( service_name="sagemaker-runtime", aws_access_key_id=aws_access_key_id, aws_secret_access_key=aws_secret_access_key, - aws_session_token=aws_session_token, region_name=aws_region_name, ) else: @@ -487,7 +481,6 @@ async def async_completion( aws_secret_access_key: Optional[str], aws_access_key_id: Optional[str], aws_region_name: Optional[str], - aws_session_token: Optional[str] = None, ): """ Use aioboto3 @@ -503,7 +496,6 @@ async def async_completion( service_name="sagemaker-runtime", aws_access_key_id=aws_access_key_id, aws_secret_access_key=aws_secret_access_key, - aws_session_token=aws_session_token, region_name=aws_region_name, ) else: @@ -647,10 +639,9 @@ def embedding( ### BOTO3 INIT import boto3 - # pop aws_secret_access_key, aws_access_key_id, aws_session_token, aws_region_name from kwargs, since completion calls fail with them + # pop aws_secret_access_key, aws_access_key_id, aws_region_name from kwargs, since completion calls fail with them aws_secret_access_key = optional_params.pop("aws_secret_access_key", None) aws_access_key_id = optional_params.pop("aws_access_key_id", None) - aws_session_token = optional_params.pop("aws_session_token", None) aws_region_name = optional_params.pop("aws_region_name", None) if aws_access_key_id is not None: @@ -660,7 +651,6 @@ def embedding( service_name="sagemaker-runtime", aws_access_key_id=aws_access_key_id, aws_secret_access_key=aws_secret_access_key, - aws_session_token=aws_session_token, region_name=aws_region_name, ) else: diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 9c1039f51e..30b90abe64 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -6343,7 +6343,6 @@ async def model_info_v2( _model["litellm_params"].pop("vertex_credentials", None) _model["litellm_params"].pop("aws_access_key_id", None) _model["litellm_params"].pop("aws_secret_access_key", None) - _model["litellm_params"].pop("aws_session_token", None) verbose_proxy_logger.debug("all_models: %s", all_models) return {"data": all_models} @@ -6860,7 +6859,6 @@ async def model_info_v1( model["litellm_params"].pop("vertex_credentials", None) model["litellm_params"].pop("aws_access_key_id", None) model["litellm_params"].pop("aws_secret_access_key", None) - model["litellm_params"].pop("aws_session_token", None) verbose_proxy_logger.debug("all_models: %s", all_models) return {"data": all_models} diff --git a/litellm/types/router.py b/litellm/types/router.py index 059a8620e5..e6864ffe2e 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -145,7 +145,6 @@ class GenericLiteLLMParams(BaseModel): ## AWS BEDROCK / SAGEMAKER ## aws_access_key_id: Optional[str] = None aws_secret_access_key: Optional[str] = None - aws_session_token: Optional[str] = None aws_region_name: Optional[str] = None ## IBM WATSONX ## watsonx_region_name: Optional[str] = None @@ -179,7 +178,6 @@ class GenericLiteLLMParams(BaseModel): ## AWS BEDROCK / SAGEMAKER ## aws_access_key_id: Optional[str] = None, aws_secret_access_key: Optional[str] = None, - aws_session_token: Optional[str] = None, aws_region_name: Optional[str] = None, ## IBM WATSONX ## watsonx_region_name: Optional[str] = None, @@ -244,7 +242,6 @@ class LiteLLM_Params(GenericLiteLLMParams): ## AWS BEDROCK / SAGEMAKER ## aws_access_key_id: Optional[str] = None, aws_secret_access_key: Optional[str] = None, - aws_session_token: Optional[str] = None, aws_region_name: Optional[str] = None, **params, ): @@ -310,7 +307,6 @@ class LiteLLMParamsTypedDict(TypedDict, total=False): ## AWS BEDROCK / SAGEMAKER ## aws_access_key_id: Optional[str] aws_secret_access_key: Optional[str] - aws_session_token: Optional[str] aws_region_name: Optional[str] ## IBM WATSONX ## watsonx_region_name: Optional[str] From ac7bc0025e1e925dea0d4f2a714b2482a88503ea Mon Sep 17 00:00:00 2001 From: Brian Schultheiss Date: Tue, 25 Jun 2024 13:33:40 -0700 Subject: [PATCH 6/9] Revert some non-essential changes --- docs/my-website/docs/providers/bedrock.md | 8 +++----- 1 file changed, 3 insertions(+), 5 deletions(-) diff --git a/docs/my-website/docs/providers/bedrock.md b/docs/my-website/docs/providers/bedrock.md index 1b073ad25e..b72dac10bc 100644 --- a/docs/my-website/docs/providers/bedrock.md +++ b/docs/my-website/docs/providers/bedrock.md @@ -474,10 +474,9 @@ from litellm import completion response = completion( model="bedrock/anthropic.claude-instant-v1", messages=[{ "content": "Hello, how are you?","role": "user"}], - aws_region_name="", aws_access_key_id="", aws_secret_access_key="", - aws_session_token=None, + aws_region_name="", ) ``` @@ -538,7 +537,6 @@ response = completion( aws_region_name=aws_region_name, aws_access_key_id=aws_access_key_id, aws_secret_access_key=aws_secret_access_key, - aws_session_token=aws_session_token, aws_role_name=aws_role_name, aws_session_name="my-test-session", ) @@ -553,7 +551,7 @@ This is a deprecated flow. Boto3 is not async. And boto3.client does not let us Experimental - 2024-Jun-23: - `aws_access_key_id`, `aws_secret_access_key`, and `aws_session_token` will be extracted from boto3.client and be passed onto the httpx client + `aws_access_key_id`, `aws_secret_access_key`, and `aws_session_token` will be extracted from boto3.client and be passed into the httpx client ::: @@ -569,7 +567,7 @@ bedrock = boto3.client( region_name="us-east-1", aws_access_key_id="", aws_secret_access_key="", - aws_session_token=None, + aws_session_token="", ) response = completion( From 746f864fb2fbc15fc844932a89cf1d440f58c807 Mon Sep 17 00:00:00 2001 From: Brian Schultheiss Date: Tue, 25 Jun 2024 13:57:12 -0700 Subject: [PATCH 7/9] Revert some non-essential changes --- litellm/llms/bedrock.py | 7 +------ 1 file changed, 1 insertion(+), 6 deletions(-) diff --git a/litellm/llms/bedrock.py b/litellm/llms/bedrock.py index 2403edf814..07d9834bf4 100644 --- a/litellm/llms/bedrock.py +++ b/litellm/llms/bedrock.py @@ -558,7 +558,6 @@ def init_bedrock_client( region_name=None, aws_access_key_id: Optional[str] = None, aws_secret_access_key: Optional[str] = None, - aws_session_token: Optional[str] = None, aws_region_name: Optional[str] = None, aws_bedrock_runtime_endpoint: Optional[str] = None, aws_session_name: Optional[str] = None, @@ -576,7 +575,6 @@ def init_bedrock_client( params_to_check = [ aws_access_key_id, aws_secret_access_key, - aws_session_token, aws_region_name, aws_bedrock_runtime_endpoint, aws_session_name, @@ -593,7 +591,6 @@ def init_bedrock_client( ( aws_access_key_id, aws_secret_access_key, - aws_session_token, aws_region_name, aws_bedrock_runtime_endpoint, aws_session_name, @@ -1438,10 +1435,9 @@ def image_generation( Bedrock Image Gen endpoint support """ ### BOTO3 INIT ### - # pop aws_secret_access_key, aws_access_key_id, aws_session_token, aws_region_name from kwargs, since completion calls fail with them + # pop aws_secret_access_key, aws_access_key_id, aws_region_name from kwargs, since completion calls fail with them aws_secret_access_key = optional_params.pop("aws_secret_access_key", None) aws_access_key_id = optional_params.pop("aws_access_key_id", None) - aws_session_token = optional_params.pop("aws_session_token", None) aws_region_name = optional_params.pop("aws_region_name", None) aws_role_name = optional_params.pop("aws_role_name", None) aws_session_name = optional_params.pop("aws_session_name", None) @@ -1454,7 +1450,6 @@ def image_generation( client = init_bedrock_client( 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_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint, aws_web_identity_token=aws_web_identity_token, From 5dce53579ec413e0d59807af838f852beb99636f Mon Sep 17 00:00:00 2001 From: Brian Schultheiss Date: Tue, 25 Jun 2024 14:09:55 -0700 Subject: [PATCH 8/9] Revert some non-essential changes --- litellm/llms/bedrock.py | 23 ++--------------------- 1 file changed, 2 insertions(+), 21 deletions(-) diff --git a/litellm/llms/bedrock.py b/litellm/llms/bedrock.py index 07d9834bf4..d0d3bef6da 100644 --- a/litellm/llms/bedrock.py +++ b/litellm/llms/bedrock.py @@ -668,21 +668,6 @@ def init_bedrock_client( endpoint_url=endpoint_url, config=config, ) - elif ( - aws_access_key_id is not None - and aws_secret_access_key is not None - and aws_session_token is not None - ): ### CHECK FOR AWS SESSION TOKEN ### - client = boto3.client( - service_name="bedrock-runtime", - aws_access_key_id=aws_access_key_id, - aws_secret_access_key=aws_secret_access_key, - aws_session_token=aws_session_token, - region_name=region_name, - endpoint_url=endpoint_url, - config=config, - ) - elif aws_role_name is not None and aws_session_name is not None: # use sts if role name passed in sts_client = boto3.client( @@ -801,10 +786,9 @@ def completion( _is_function_call = False json_schemas: dict = {} try: - # pop aws_secret_access_key, aws_access_key_id, aws_session_token, aws_region_name from kwargs, since completion calls fail with them + # pop aws_secret_access_key, aws_access_key_id, aws_region_name from kwargs, since completion calls fail with them aws_secret_access_key = optional_params.pop("aws_secret_access_key", None) aws_access_key_id = optional_params.pop("aws_access_key_id", None) - aws_session_token = optional_params.pop("aws_session_token", None) aws_region_name = optional_params.pop("aws_region_name", None) aws_role_name = optional_params.pop("aws_role_name", None) aws_session_name = optional_params.pop("aws_session_name", None) @@ -822,7 +806,6 @@ def completion( client = init_bedrock_client( 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_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint, aws_role_name=aws_role_name, @@ -1342,10 +1325,9 @@ def embedding( encoding=None, ): ### BOTO3 INIT ### - # pop aws_secret_access_key, aws_access_key_id, aws_session_token, aws_region_name from kwargs, since completion calls fail with them + # pop aws_secret_access_key, aws_access_key_id, aws_region_name from kwargs, since completion calls fail with them aws_secret_access_key = optional_params.pop("aws_secret_access_key", None) aws_access_key_id = optional_params.pop("aws_access_key_id", None) - aws_session_token = optional_params.pop("aws_session_token", None) aws_region_name = optional_params.pop("aws_region_name", None) aws_role_name = optional_params.pop("aws_role_name", None) aws_session_name = optional_params.pop("aws_session_name", None) @@ -1358,7 +1340,6 @@ def embedding( client = init_bedrock_client( 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_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint, aws_web_identity_token=aws_web_identity_token, From 09492ccebace713fa94e5d368ebbff4c4e67fd30 Mon Sep 17 00:00:00 2001 From: Brian Schultheiss Date: Tue, 25 Jun 2024 14:33:40 -0700 Subject: [PATCH 9/9] Update tests to verify streaming works --- litellm/tests/test_bedrock_completion.py | 46 ++++++++++++++++++++++++ 1 file changed, 46 insertions(+) diff --git a/litellm/tests/test_bedrock_completion.py b/litellm/tests/test_bedrock_completion.py index 614d1b76fa..d21f30549d 100644 --- a/litellm/tests/test_bedrock_completion.py +++ b/litellm/tests/test_bedrock_completion.py @@ -238,6 +238,20 @@ def bedrock_session_token_creds(): ) return creds +def process_stream_response(res, messages): + import types + if isinstance(res, litellm.utils.CustomStreamWrapper): + chunks = [] + for part in res: + chunks.append(part) + text = part.choices[0].delta.content or "" + print(text, end="") + res = litellm.stream_chunk_builder(chunks, messages=messages) + else: + raise ValueError("Response object is not a streaming response") + + return res + @pytest.mark.skipif( os.environ.get("CIRCLE_OIDC_TOKEN_V2") is None, reason="Cannot run without being in CircleCI Runner", @@ -298,6 +312,23 @@ def test_completion_bedrock_claude_aws_session_token(bedrock_session_token_creds assert len(response_3.choices) > 0 assert len(response_3.choices[0].message.content) > 0 + # This fourth call is to verify streaming api works + response_4 = completion( + model="bedrock/anthropic.claude-3-haiku-20240307-v1:0", + messages=messages, + max_tokens=6, + temperature=0.3, + aws_region_name="us-east-1", + aws_access_key_id=aws_access_key_id, + aws_secret_access_key=aws_secret_access_key, + aws_session_token=aws_session_token, + stream=True + ) + response_4 = process_stream_response(response_4, messages) + print(response_4) + assert len(response_4.choices) > 0 + assert len(response_4.choices[0].message.content) > 0 + except RateLimitError: pass except Exception as e: @@ -380,6 +411,21 @@ def test_completion_bedrock_claude_aws_bedrock_client(bedrock_session_token_cred assert len(response_3.choices) > 0 assert len(response_3.choices[0].message.content) > 0 + # This fourth call is to verify streaming api works + response_4 = completion( + model="bedrock/anthropic.claude-3-haiku-20240307-v1:0", + messages=messages, + max_tokens=6, + temperature=0.3, + aws_bedrock_client=aws_bedrock_client_east, + stream=True + ) + response_4 = process_stream_response(response_4, messages) + print(response_4) + assert len(response_4.choices) > 0 + assert len(response_4.choices[0].message.content) > 0 + + except RateLimitError: pass except Exception as e: