diff --git a/litellm/main.py b/litellm/main.py index 3f236523995..f2e1f7e2998 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -254,9 +254,10 @@ def completion( metadata = kwargs.get('metadata', None) fallbacks = kwargs.get('fallbacks', None) headers = kwargs.get("headers", None) + num_retries = kwargs.get("num_retries", None) ######## end of unpacking kwargs ########### openai_params = ["functions", "function_call", "temperature", "temperature", "top_p", "n", "stream", "stop", "max_tokens", "presence_penalty", "frequency_penalty", "logit_bias", "user", "request_timeout", "api_base", "api_version", "api_key"] - litellm_params = ["metadata", "acompletion", "caching", "return_async", "mock_response", "api_key", "api_version", "api_base", "force_timeout", "logger_fn", "verbose", "custom_llm_provider", "litellm_logging_obj", "litellm_call_id", "use_client", "id", "fallbacks", "azure", "headers", "model_list"] + litellm_params = ["metadata", "acompletion", "caching", "return_async", "mock_response", "api_key", "api_version", "api_base", "force_timeout", "logger_fn", "verbose", "custom_llm_provider", "litellm_logging_obj", "litellm_call_id", "use_client", "id", "fallbacks", "azure", "headers", "model_list", "num_retries"] default_params = openai_params + litellm_params non_default_params = {k: v for k,v in kwargs.items() if k not in default_params} # model-specific params - pass them straight to the model/provider if mock_response: @@ -1325,9 +1326,19 @@ def completion( return response except Exception as e: ## Map to OpenAI Exception - raise exception_type( - model=model, custom_llm_provider=custom_llm_provider, original_exception=e, completion_kwargs=args, - ) + try: + raise exception_type( + model=model, custom_llm_provider=custom_llm_provider, original_exception=e, completion_kwargs=args, + ) + except Exception as e: + if num_retries: + if (isinstance(e, openai.APIError) + or isinstance(e, openai.Timeout) + or isinstance(e, openai.Timeout) + or isinstance(e, openai.ServiceUnavailableError)): + return completion_with_retries(num_retries=num_retries, **args) + else: + raise e def completion_with_retries(*args, **kwargs): @@ -1338,8 +1349,9 @@ def completion_with_retries(*args, **kwargs): import tenacity except: raise Exception("tenacity import failed please run `pip install tenacity`") - - retryer = tenacity.Retrying(stop=tenacity.stop_after_attempt(3), reraise=True) + + num_retries = kwargs.pop("num_retries", 3) + retryer = tenacity.Retrying(stop=tenacity.stop_after_attempt(num_retries), reraise=True) return retryer(completion, *args, **kwargs) diff --git a/litellm/tests/test_completion_with_retries.py b/litellm/tests/test_completion_with_retries.py index bfc077b1d2f..9b54f94f8f2 100644 --- a/litellm/tests/test_completion_with_retries.py +++ b/litellm/tests/test_completion_with_retries.py @@ -1,37 +1,61 @@ -# import sys, os -# import traceback -# from dotenv import load_dotenv +import sys, os +import traceback +from dotenv import load_dotenv -# load_dotenv() -# import os +load_dotenv() +import os -# sys.path.insert( -# 0, os.path.abspath("../..") -# ) # Adds the parent directory to the system path -# import pytest -# import litellm -# from litellm import completion_with_retries -# from litellm import ( -# AuthenticationError, -# InvalidRequestError, -# RateLimitError, -# ServiceUnavailableError, -# OpenAIError, -# ) +sys.path.insert( + 0, os.path.abspath("../..") +) # Adds the parent directory to the system path +import pytest +import openai +import litellm +from litellm import completion_with_retries, completion +from litellm import ( + AuthenticationError, + InvalidRequestError, + RateLimitError, + ServiceUnavailableError, + OpenAIError, +) -# user_message = "Hello, whats the weather in San Francisco??" -# messages = [{"content": user_message, "role": "user"}] +user_message = "Hello, whats the weather in San Francisco??" +messages = [{"content": user_message, "role": "user"}] -# def logger_fn(user_model_dict): -# # print(f"user_model_dict: {user_model_dict}") -# pass +def logger_fn(user_model_dict): + # print(f"user_model_dict: {user_model_dict}") + pass -# # normal call +# normal call +def test_completion_custom_provider_model_name(): + try: + response = completion_with_retries( + model="together_ai/togethercomputer/llama-2-70b-chat", + messages=messages, + logger_fn=logger_fn, + ) + # Add any assertions here to check the response + print(response) + except Exception as e: + pytest.fail(f"Error occurred: {e}") + +# completion with num retries +def test_completion_with_num_retries(): + try: + response = completion(model="j2-ultra", messages=[{"messages": "vibe", "bad": "message"}], num_retries=2) + except openai.APIError as e: + pass + except Exception as e: + pytest.fail(f"Unmapped exception occurred") + +test_completion_with_num_retries() +# bad call # def test_completion_custom_provider_model_name(): # try: # response = completion_with_retries( -# model="together_ai/togethercomputer/llama-2-70b-chat", +# model="bad-model", # messages=messages, # logger_fn=logger_fn, # ) @@ -40,45 +64,31 @@ # except Exception as e: # pytest.fail(f"Error occurred: {e}") - -# # bad call -# # def test_completion_custom_provider_model_name(): -# # try: -# # response = completion_with_retries( -# # model="bad-model", -# # messages=messages, -# # logger_fn=logger_fn, -# # ) -# # # Add any assertions here to check the response -# # print(response) -# # except Exception as e: -# # pytest.fail(f"Error occurred: {e}") - -# # impact on exception mapping -# def test_context_window(): -# sample_text = "how does a court case get to the Supreme Court?" * 5000 -# messages = [{"content": sample_text, "role": "user"}] -# try: -# model = "chatgpt-test" -# response = completion_with_retries( -# model=model, -# messages=messages, -# custom_llm_provider="azure", -# logger_fn=logger_fn, -# ) -# print(f"response: {response}") -# except InvalidRequestError as e: -# print(f"InvalidRequestError: {e.llm_provider}") -# return -# except OpenAIError as e: -# print(f"OpenAIError: {e.llm_provider}") -# return -# except Exception as e: -# print("Uncaught Error in test_context_window") -# print(f"Error Type: {type(e).__name__}") -# print(f"Uncaught Exception - {e}") -# pytest.fail(f"Error occurred: {e}") -# return +# impact on exception mapping +def test_context_window(): + sample_text = "how does a court case get to the Supreme Court?" * 5000 + messages = [{"content": sample_text, "role": "user"}] + try: + model = "chatgpt-test" + response = completion_with_retries( + model=model, + messages=messages, + custom_llm_provider="azure", + logger_fn=logger_fn, + ) + print(f"response: {response}") + except InvalidRequestError as e: + print(f"InvalidRequestError: {e.llm_provider}") + return + except OpenAIError as e: + print(f"OpenAIError: {e.llm_provider}") + return + except Exception as e: + print("Uncaught Error in test_context_window") + print(f"Error Type: {type(e).__name__}") + print(f"Uncaught Exception - {e}") + pytest.fail(f"Error occurred: {e}") + return # test_context_window()