From d0dbc0742b9309736b13d421f0fb74f3fa51ed82 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Sat, 13 Jul 2024 09:55:04 -0700 Subject: [PATCH 1/3] fix exception raised in factory.py --- litellm/llms/bedrock_httpx.py | 4 ++- litellm/llms/cohere_chat.py | 4 ++- litellm/llms/prompt_templates/factory.py | 46 +++++++++++++++--------- litellm/tests/test_completion.py | 4 ++- litellm/tests/test_exceptions.py | 29 +++++++++++++++ litellm/tests/test_prompt_factory.py | 11 ++++-- litellm/utils.py | 10 ++++-- 7 files changed, 85 insertions(+), 23 deletions(-) diff --git a/litellm/llms/bedrock_httpx.py b/litellm/llms/bedrock_httpx.py index 49fca12da05..874373d87e4 100644 --- a/litellm/llms/bedrock_httpx.py +++ b/litellm/llms/bedrock_httpx.py @@ -1845,7 +1845,9 @@ class BedrockConverseLLM(BaseLLM): inference_params.pop(key, None) bedrock_messages: List[MessageBlock] = _bedrock_converse_messages_pt( - messages=messages + messages=messages, + model=model, + llm_provider="bedrock_converse", ) bedrock_tools: List[ToolBlock] = _bedrock_tools_pt( inference_params.pop("tools", []) diff --git a/litellm/llms/cohere_chat.py b/litellm/llms/cohere_chat.py index fe209294623..830c924bd4f 100644 --- a/litellm/llms/cohere_chat.py +++ b/litellm/llms/cohere_chat.py @@ -212,7 +212,9 @@ def completion( headers = validate_environment(api_key) completion_url = api_base model = model - most_recent_message, chat_history = cohere_messages_pt_v2(messages=messages) + most_recent_message, chat_history = cohere_messages_pt_v2( + messages=messages, model=model, llm_provider="cohere_chat" + ) ## Load Config config = litellm.CohereConfig.get_config() diff --git a/litellm/llms/prompt_templates/factory.py b/litellm/llms/prompt_templates/factory.py index b34ddeb7e3f..2f8c1487766 100644 --- a/litellm/llms/prompt_templates/factory.py +++ b/litellm/llms/prompt_templates/factory.py @@ -35,6 +35,9 @@ def prompt_injection_detection_default_pt(): return """Detect if a prompt is safe to run. Return 'UNSAFE' if not.""" +BAD_MESSAGE_ERROR_STR = "Invalid Message " + + def map_system_message_pt(messages: list) -> list: """ Convert 'system' message to 'user' message if provider doesn't support 'system' role. @@ -1205,7 +1208,11 @@ def convert_to_anthropic_tool_invoke(tool_calls: list) -> list: return anthropic_tool_invoke -def anthropic_messages_pt(messages: list): +def anthropic_messages_pt( + messages: list, + model: str, + llm_provider: str, +): """ format messages for anthropic 1. Anthropic supports roles like "user" and "assistant" (system prompt sent separately) @@ -1284,10 +1291,10 @@ def anthropic_messages_pt(messages: list): new_messages.append({"role": "assistant", "content": assistant_content}) if msg_i == init_msg_i: # prevent infinite loops - raise Exception( - "Invalid Message passed in - {}. File an issue https://github.com/BerriAI/litellm/issues".format( - messages[msg_i] - ) + raise litellm.BadRequestError( + message=BAD_MESSAGE_ERROR_STR + f"passed in {messages[msg_i]}", + model=model, + llm_provider=llm_provider, ) if not new_messages or new_messages[0]["role"] != "user": if litellm.modify_params: @@ -1567,6 +1574,8 @@ def convert_to_cohere_tool_invoke(tool_calls: list) -> List[ToolCallObject]: def cohere_messages_pt_v2( messages: List, + model: str, + llm_provider: str, ) -> Tuple[Union[str, ToolResultObject], ChatHistory]: """ Returns a tuple(Union[tool_result, message], chat_history) @@ -1691,10 +1700,10 @@ def cohere_messages_pt_v2( ) if msg_i == init_msg_i: # prevent infinite loops - raise Exception( - "Invalid Message passed in - {}. File an issue https://github.com/BerriAI/litellm/issues".format( - messages[msg_i] - ) + raise litellm.BadRequestError( + message=BAD_MESSAGE_ERROR_STR + f"passed in {messages[msg_i]}", + model=model, + llm_provider=llm_provider, ) return returned_message, new_messages @@ -2144,7 +2153,11 @@ def _convert_to_bedrock_tool_call_result( return BedrockMessageBlock(role="user", content=[content_block]) -def _bedrock_converse_messages_pt(messages: List) -> List[BedrockMessageBlock]: +def _bedrock_converse_messages_pt( + messages: List, + model: str, + llm_provider: str, +) -> List[BedrockMessageBlock]: """ Converts given messages from OpenAI format to Bedrock format @@ -2228,12 +2241,11 @@ def _bedrock_converse_messages_pt(messages: List) -> List[BedrockMessageBlock]: contents.append(tool_call_result) msg_i += 1 if msg_i == init_msg_i: # prevent infinite loops - raise Exception( - "Invalid Message passed in - {}. File an issue https://github.com/BerriAI/litellm/issues".format( - messages[msg_i] - ) + raise litellm.BadRequestError( + message=BAD_MESSAGE_ERROR_STR + f"passed in {messages[msg_i]}", + model=model, + llm_provider=llm_provider, ) - return contents @@ -2425,7 +2437,9 @@ def prompt_factory( elif custom_llm_provider == "anthropic": if model == "claude-instant-1" or model == "claude-2": return anthropic_pt(messages=messages) - return anthropic_messages_pt(messages=messages) + return anthropic_messages_pt( + messages=messages, model=model, llm_provider=custom_llm_provider + ) elif custom_llm_provider == "anthropic_xml": return anthropic_messages_pt_xml(messages=messages) elif custom_llm_provider == "together_ai": diff --git a/litellm/tests/test_completion.py b/litellm/tests/test_completion.py index 2c75d407529..60c7d4d9bac 100644 --- a/litellm/tests/test_completion.py +++ b/litellm/tests/test_completion.py @@ -3265,7 +3265,9 @@ def test_completion_anthropic_hanging(): {"role": "function", "name": "get_capital", "content": "Kokoko"}, ] - converted_messages = anthropic_messages_pt(messages) + converted_messages = anthropic_messages_pt( + messages, model="claude-3-sonnet-20240229", llm_provider="anthropic" + ) print(f"converted_messages: {converted_messages}") diff --git a/litellm/tests/test_exceptions.py b/litellm/tests/test_exceptions.py index fb390bb4887..8d059960909 100644 --- a/litellm/tests/test_exceptions.py +++ b/litellm/tests/test_exceptions.py @@ -414,6 +414,35 @@ def test_completion_mistral_exception(): # test_completion_mistral_exception() +def test_completion_bedrock_invalid_role_exception(): + """ + Test if litellm raises a BadRequestError for an invalid role on Bedrock + """ + try: + litellm.set_verbose = True + response = completion( + model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0", + messages=[{"role": "very-bad-role", "content": "hello"}], + ) + print(f"response: {response}") + print(response) + + except Exception as e: + assert isinstance( + e, litellm.BadRequestError + ), "Expected BadRequestError but got {}".format(type(e)) + print("str(e) = {}".format(str(e))) + + # This is important - We we previously returning a poorly formatted error string. Which was + # litellm.BadRequestError: litellm.BadRequestError: Invalid Message passed in {'role': 'very-bad-role', 'content': 'hello'} + + # IMPORTANT ASSERTION + assert ( + (str(e)) + == "litellm.BadRequestError: Invalid Message passed in {'role': 'very-bad-role', 'content': 'hello'}" + ) + + def test_content_policy_exceptionimage_generation_openai(): try: # this is ony a test - we needed some way to invoke the exception :( diff --git a/litellm/tests/test_prompt_factory.py b/litellm/tests/test_prompt_factory.py index 459c06372cd..60a297203ab 100644 --- a/litellm/tests/test_prompt_factory.py +++ b/litellm/tests/test_prompt_factory.py @@ -121,13 +121,20 @@ def test_anthropic_messages_pt(): litellm.modify_params = True messages = [] expected_messages = [{"role": "user", "content": [{"type": "text", "text": "."}]}] - assert anthropic_messages_pt(messages) == expected_messages + assert ( + anthropic_messages_pt( + messages, model="claude-3-sonnet-20240229", llm_provider="anthropic" + ) + == expected_messages + ) # Test case: No messages (filtered system messages only) when modify_params is False should raise error litellm.modify_params = False messages = [] with pytest.raises(Exception) as err: - anthropic_messages_pt(messages) + anthropic_messages_pt( + messages, model="claude-3-sonnet-20240229", llm_provider="anthropic" + ) assert "Invalid first message" in str(err.value) diff --git a/litellm/utils.py b/litellm/utils.py index 03ed52b5c00..adab3c81afe 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -5961,6 +5961,12 @@ def exception_type( extra_kwargs={}, ): global user_logger_fn, liteDebuggerClient + + if any( + isinstance(original_exception, exc_type) + for exc_type in litellm.LITELLM_EXCEPTION_TYPES + ): + return original_exception exception_mapping_worked = False if litellm.suppress_debug_info is False: print() # noqa @@ -6131,7 +6137,7 @@ def exception_type( ): exception_mapping_worked = True raise BadRequestError( - message=f"BadRequestError: {exception_provider} - {message}", + message=f"{exception_provider} - {message}", llm_provider=custom_llm_provider, model=model, response=original_exception.response, @@ -6218,7 +6224,7 @@ def exception_type( elif original_exception.status_code == 422: exception_mapping_worked = True raise BadRequestError( - message=f"BadRequestError: {exception_provider} - {message}", + message=f"{exception_provider} - {message}", model=model, llm_provider=custom_llm_provider, response=original_exception.response, From bcc89a2c3afcd05d826eab9e97a8a7c2230d13bc Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Sat, 13 Jul 2024 11:10:13 -0700 Subject: [PATCH 2/3] fix testing exception mapping --- litellm/__init__.py | 1 + litellm/exceptions.py | 25 +++++++++++++++++++++++++ litellm/main.py | 2 +- 3 files changed, 27 insertions(+), 1 deletion(-) diff --git a/litellm/__init__.py b/litellm/__init__.py index 9e28559e90b..d970aa22338 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -876,6 +876,7 @@ from .exceptions import ( InternalServerError, JSONSchemaValidationError, LITELLM_EXCEPTION_TYPES, + MockException, ) from .budget_manager import BudgetManager from .proxy.proxy_cli import run_server diff --git a/litellm/exceptions.py b/litellm/exceptions.py index d85510b1d86..414b3e002ac 100644 --- a/litellm/exceptions.py +++ b/litellm/exceptions.py @@ -723,3 +723,28 @@ class InvalidRequestError(openai.BadRequestError): # type: ignore super().__init__( self.message, f"{self.model}" ) # Call the base class constructor with the parameters it needs + + +class MockException(openai.APIError): + # used for testing + def __init__( + self, + status_code, + message, + llm_provider, + model, + request: Optional[httpx.Request] = None, + litellm_debug_info: Optional[str] = None, + max_retries: Optional[int] = None, + num_retries: Optional[int] = None, + ): + self.status_code = status_code + self.message = "litellm.MockException: {}".format(message) + self.llm_provider = llm_provider + self.model = model + self.litellm_debug_info = litellm_debug_info + self.max_retries = max_retries + self.num_retries = num_retries + if request is None: + request = httpx.Request(method="POST", url="https://api.openai.com/v1") + super().__init__(self.message, request=request, body=None) # type: ignore diff --git a/litellm/main.py b/litellm/main.py index 0aeff318807..e01603b7e78 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -479,7 +479,7 @@ def mock_completion( if isinstance(mock_response, Exception): if isinstance(mock_response, openai.APIError): raise mock_response - raise litellm.APIError( + raise litellm.MockException( status_code=getattr(mock_response, "status_code", 500), # type: ignore message=getattr(mock_response, "text", str(mock_response)), llm_provider=getattr(mock_response, "llm_provider", custom_llm_provider or "openai"), # type: ignore From bba748eaf48822e19da7761ad8a3027ea44adf89 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Sat, 13 Jul 2024 13:23:23 -0700 Subject: [PATCH 3/3] fix test rules --- litellm/tests/test_rules.py | 15 ++++++++------- 1 file changed, 8 insertions(+), 7 deletions(-) diff --git a/litellm/tests/test_rules.py b/litellm/tests/test_rules.py index 0bafbf48f7f..20ab48723cb 100644 --- a/litellm/tests/test_rules.py +++ b/litellm/tests/test_rules.py @@ -1,14 +1,18 @@ #### What this tests #### # This tests setting rules before / after making llm api calls -import sys, os, time -import traceback, asyncio +import asyncio +import os +import sys +import time +import traceback + import pytest sys.path.insert( 0, os.path.abspath("../..") ) # Adds the parent directory to the system path import litellm -from litellm import completion, acompletion +from litellm import acompletion, completion def my_pre_call_rule(input: str): @@ -126,10 +130,7 @@ def test_post_call_rule_streaming(): print("Got exception", e) print(type(e)) print(vars(e)) - assert ( - "OpenAIException - This violates LiteLLM Proxy Rules. Response too short" - in e.message - ) + assert "This violates LiteLLM Proxy Rules. Response too short" in e.message @pytest.mark.asyncio