From 6488de36c4310c311e7f07e071dc0e2b005c64a0 Mon Sep 17 00:00:00 2001 From: ishaan-jaff Date: Tue, 30 Jan 2024 08:12:43 -0800 Subject: [PATCH 1/3] (test) bedrock input validation - exceptions --- litellm/tests/test_embedding.py | 19 +++++++++++++++++++ 1 file changed, 19 insertions(+) diff --git a/litellm/tests/test_embedding.py b/litellm/tests/test_embedding.py index a005a6ad168..681471e3dbe 100644 --- a/litellm/tests/test_embedding.py +++ b/litellm/tests/test_embedding.py @@ -302,6 +302,25 @@ def test_bedrock_embedding_cohere(): # test_bedrock_embedding_cohere() +def test_demo_tokens_as_input_to_embeddings_fails_for_titan(): + litellm.set_verbose = True + + with pytest.raises( + litellm.BadRequestError, + match="BedrockException - Bedrock Embedding API input must be type str | List[str]", + ): + litellm.embedding(model="amazon.titan-embed-text-v1", input=[[1]]) + + with pytest.raises( + litellm.BadRequestError, + match="BedrockException - Bedrock Embedding API input must be type str | List[str]", + ): + litellm.embedding( + model="amazon.titan-embed-text-v1", + input=[1], + ) + + # comment out hf tests - since hf endpoints are unstable def test_hf_embedding(): try: From b9a117fa1d9eeae65fa55ae0fb584abc8f99f12e Mon Sep 17 00:00:00 2001 From: ishaan-jaff Date: Tue, 30 Jan 2024 08:14:35 -0800 Subject: [PATCH 2/3] (feat) Bedrock embedding - raise correct exception BadRequest --- litellm/llms/bedrock.py | 13 ++++++++++++- 1 file changed, 12 insertions(+), 1 deletion(-) diff --git a/litellm/llms/bedrock.py b/litellm/llms/bedrock.py index bcf35c3d1f4..16a0abbed7b 100644 --- a/litellm/llms/bedrock.py +++ b/litellm/llms/bedrock.py @@ -702,6 +702,11 @@ def _embedding_func_single( encoding=None, logging_obj=None, ): + if type(input) != str: + raise BedrockError( + message="Bedrock Embedding API input must be type str | List[str]", + status_code=400, + ) # logic for parsing in - calling - parsing out model embedding calls ## FORMAT EMBEDDING INPUT ## provider = model.split(".")[0] @@ -805,7 +810,7 @@ def embedding( logging_obj=logging_obj, ) ] - else: + elif type(input) == list: ## Embedding Call embeddings = [ _embedding_func_single( @@ -817,6 +822,12 @@ def embedding( ) for i in input ] # [TODO]: make these parallel calls + else: + # enters this branch if input = int, ex. input=2 + raise BedrockError( + message="Bedrock Embedding API input must be type str | List[str]", + status_code=400, + ) ## Populate OpenAI compliant dictionary embedding_response = [] From f941c57688c74d09a49719d8298c47df49f3f602 Mon Sep 17 00:00:00 2001 From: ishaan-jaff Date: Tue, 30 Jan 2024 08:31:21 -0800 Subject: [PATCH 3/3] (fix) use isinstance to check types --- litellm/llms/bedrock.py | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/litellm/llms/bedrock.py b/litellm/llms/bedrock.py index 16a0abbed7b..b67061c76b2 100644 --- a/litellm/llms/bedrock.py +++ b/litellm/llms/bedrock.py @@ -702,7 +702,7 @@ def _embedding_func_single( encoding=None, logging_obj=None, ): - if type(input) != str: + if isinstance(input, str) is False: raise BedrockError( message="Bedrock Embedding API input must be type str | List[str]", status_code=400, @@ -800,7 +800,8 @@ def embedding( aws_role_name=aws_role_name, aws_session_name=aws_session_name, ) - if type(input) == str: + if isinstance(input, str): + ## Embedding Call embeddings = [ _embedding_func_single( model, @@ -810,8 +811,8 @@ def embedding( logging_obj=logging_obj, ) ] - elif type(input) == list: - ## Embedding Call + elif isinstance(input, list): + ## Embedding Call - assuming this is a List[str] embeddings = [ _embedding_func_single( model,