Merge pull request #5143 from BerriAI/litellm_use_max_retries

fix(utils.py): set max_retries = num_retries, if given
This commit is contained in:
Krish Dholakia 2024-08-09 20:29:49 -07:00 committed by GitHub
commit a42164fa79
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 114 additions and 64 deletions

View file

@ -686,7 +686,9 @@ def completion(
proxy_server_request = kwargs.get("proxy_server_request", None)
fallbacks = kwargs.get("fallbacks", None)
headers = kwargs.get("headers", None) or extra_headers
num_retries = kwargs.get("num_retries", None) ## deprecated
num_retries = kwargs.get(
"num_retries", None
) ## alt. param for 'max_retries'. Use this to pass retries w/ instructor.
max_retries = kwargs.get("max_retries", None)
cooldown_time = kwargs.get("cooldown_time", None)
context_window_fallback_dict = kwargs.get("context_window_fallback_dict", None)
@ -762,8 +764,8 @@ def completion(
try:
if base_url is not None:
api_base = base_url
if max_retries is not None: # openai allows openai.OpenAI(max_retries=3)
num_retries = max_retries
if num_retries is not None:
max_retries = num_retries
logging = litellm_logging_obj
fallbacks = fallbacks or litellm.model_fallbacks
if fallbacks is not None:

View file

@ -1,5 +1,5 @@
# #### What this tests ####
# # This tests the LiteLLM Class
# # #### What this tests ####
# # # This tests the LiteLLM Class
# import sys, os
# import traceback
@ -11,83 +11,114 @@
# import litellm
# import asyncio
# litellm.set_verbose = True
# from litellm import Router
# # litellm.set_verbose = True
# # from litellm import Router
# import instructor
# from litellm import completion
# from pydantic import BaseModel
# # This enables response_model keyword
# # from client.chat.completions.create
# client = instructor.patch(
# Router(
# model_list=[
# {
# "model_name": "gpt-3.5-turbo", # openai model name
# "litellm_params": { # params for litellm completion/embedding call
# "model": "azure/chatgpt-v-2",
# "api_key": os.getenv("AZURE_API_KEY"),
# "api_version": os.getenv("AZURE_API_VERSION"),
# "api_base": os.getenv("AZURE_API_BASE"),
# },
# }
# ]
# )
# )
# class UserDetail(BaseModel):
# class User(BaseModel):
# name: str
# age: int
# user = client.chat.completions.create(
# client = instructor.from_litellm(completion)
# litellm.set_verbose = True
# resp = client.chat.completions.create(
# model="gpt-3.5-turbo",
# response_model=UserDetail,
# max_tokens=1024,
# messages=[
# {"role": "user", "content": "Extract Jason is 25 years old"},
# {
# "role": "user",
# "content": "Extract Jason is 25 years old.",
# }
# ],
# response_model=User,
# num_retries=10,
# )
# assert isinstance(user, UserDetail)
# assert user.name == "Jason"
# assert user.age == 25
# assert isinstance(resp, User)
# assert resp.name == "Jason"
# assert resp.age == 25
# print(f"user: {user}")
# # import instructor
# # from openai import AsyncOpenAI
# # from pydantic import BaseModel
# aclient = instructor.apatch(
# Router(
# model_list=[
# {
# "model_name": "gpt-3.5-turbo", # openai model name
# "litellm_params": { # params for litellm completion/embedding call
# "model": "azure/chatgpt-v-2",
# "api_key": os.getenv("AZURE_API_KEY"),
# "api_version": os.getenv("AZURE_API_VERSION"),
# "api_base": os.getenv("AZURE_API_BASE"),
# },
# }
# ],
# default_litellm_params={"acompletion": True},
# )
# )
# # # This enables response_model keyword
# # # from client.chat.completions.create
# # client = instructor.patch(
# # Router(
# # model_list=[
# # {
# # "model_name": "gpt-3.5-turbo", # openai model name
# # "litellm_params": { # params for litellm completion/embedding call
# # "model": "azure/chatgpt-v-2",
# # "api_key": os.getenv("AZURE_API_KEY"),
# # "api_version": os.getenv("AZURE_API_VERSION"),
# # "api_base": os.getenv("AZURE_API_BASE"),
# # },
# # }
# # ]
# # )
# # )
# class UserExtract(BaseModel):
# name: str
# age: int
# # class UserDetail(BaseModel):
# # name: str
# # age: int
# async def main():
# model = await aclient.chat.completions.create(
# model="gpt-3.5-turbo",
# response_model=UserExtract,
# messages=[
# {"role": "user", "content": "Extract jason is 25 years old"},
# ],
# )
# print(f"model: {model}")
# # user = client.chat.completions.create(
# # model="gpt-3.5-turbo",
# # response_model=UserDetail,
# # messages=[
# # {"role": "user", "content": "Extract Jason is 25 years old"},
# # ],
# # )
# # assert isinstance(user, UserDetail)
# # assert user.name == "Jason"
# # assert user.age == 25
# # print(f"user: {user}")
# # # import instructor
# # # from openai import AsyncOpenAI
# # aclient = instructor.apatch(
# # Router(
# # model_list=[
# # {
# # "model_name": "gpt-3.5-turbo", # openai model name
# # "litellm_params": { # params for litellm completion/embedding call
# # "model": "azure/chatgpt-v-2",
# # "api_key": os.getenv("AZURE_API_KEY"),
# # "api_version": os.getenv("AZURE_API_VERSION"),
# # "api_base": os.getenv("AZURE_API_BASE"),
# # },
# # }
# # ],
# # default_litellm_params={"acompletion": True},
# # )
# # )
# asyncio.run(main())
# # class UserExtract(BaseModel):
# # name: str
# # age: int
# # async def main():
# # model = await aclient.chat.completions.create(
# # model="gpt-3.5-turbo",
# # response_model=UserExtract,
# # messages=[
# # {"role": "user", "content": "Extract jason is 25 years old"},
# # ],
# # )
# # print(f"model: {model}")
# # asyncio.run(main())

View file

@ -446,3 +446,20 @@ def test_bedrock_optional_params_embeddings_provider_specific_params():
wait_for_model=True,
)
assert len(optional_params) == 1
def test_get_optional_params_num_retries():
"""
Relevant issue - https://github.com/BerriAI/litellm/issues/5124
"""
with patch("litellm.main.get_optional_params", new=MagicMock()) as mock_client:
_ = litellm.completion(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "Hello world"}],
num_retries=10,
)
mock_client.assert_called()
print(f"mock_client.call_args: {mock_client.call_args}")
assert mock_client.call_args.kwargs["max_retries"] == 10