From ad55f4dbb5dee960d341ca2470e1e615b9255ecc Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Tue, 5 Mar 2024 19:00:03 -0800 Subject: [PATCH 1/5] feat(proxy_server.py): retry if virtual key is rate limited currently for chat completions --- .../proxy/hooks/parallel_request_limiter.py | 8 ++++- litellm/proxy/proxy_server.py | 19 ++++++++++++ proxy_server_config.yaml | 2 ++ tests/test_keys.py | 29 +++++++++++++++++++ 4 files changed, 57 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index 4221b064eee..8982e4e2bf8 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -71,7 +71,9 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): ): self.print_verbose(f"Inside Max Parallel Request Pre-Call Hook") api_key = user_api_key_dict.api_key - max_parallel_requests = user_api_key_dict.max_parallel_requests or sys.maxsize + max_parallel_requests = user_api_key_dict.max_parallel_requests + if max_parallel_requests is None: + max_parallel_requests = sys.maxsize tpm_limit = getattr(user_api_key_dict, "tpm_limit", sys.maxsize) if tpm_limit is None: tpm_limit = sys.maxsize @@ -105,6 +107,10 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): and rpm_limit == sys.maxsize ): pass + elif max_parallel_requests == 0 or tpm_limit == 0 or rpm_limit == 0: + raise HTTPException( + status_code=429, detail="Max parallel request limit reached." + ) elif current is None: new_val = { "current_requests": 1, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 628f5585232..ef54f29bd39 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -8,6 +8,7 @@ import hashlib, uuid import warnings import importlib import warnings +import backoff def showwarning(message, category, filename, lineno, file=None, line=None): @@ -2298,6 +2299,11 @@ def parse_cache_control(cache_control): return cache_dict +def on_backoff(details): + # The 'tries' key in the details dictionary contains the number of completed tries + verbose_proxy_logger.debug(f"Backing off... this was attempt #{details['tries']}") + + @router.on_event("startup") async def startup_event(): global prisma_client, master_key, use_background_health_checks, llm_router, llm_model_list, general_settings, proxy_budget_rescheduler_min_time, proxy_budget_rescheduler_max_time, litellm_proxy_admin_name @@ -2613,6 +2619,19 @@ async def completion( dependencies=[Depends(user_api_key_auth)], tags=["chat/completions"], ) # azure compatible endpoint +@backoff.on_exception( + backoff.expo, + Exception, # base exception to catch for the backoff + max_tries=litellm.num_retries or 3, # maximum number of retries + max_time=litellm.request_timeout or 60, # maximum total time to retry for + on_backoff=on_backoff, # specifying the function to call on backoff + giveup=lambda e: not ( + isinstance(e, ProxyException) + and getattr(e, "message", None) is not None + and isinstance(e.message, str) + and "Max parallel request limit reached" in e.message + ), # the result of the logical expression is on the second position +) async def chat_completion( request: Request, fastapi_response: Response, diff --git a/proxy_server_config.yaml b/proxy_server_config.yaml index 64183f2165d..4b454f5bd92 100644 --- a/proxy_server_config.yaml +++ b/proxy_server_config.yaml @@ -38,6 +38,8 @@ litellm_settings: drop_params: True max_budget: 100 budget_duration: 30d + num_retries: 5 + request_timeout: 600 general_settings: master_key: sk-1234 # [OPTIONAL] Only use this if you to require all calls to contain this key (Authorization: Bearer sk-1234) proxy_budget_rescheduler_min_time: 10 diff --git a/tests/test_keys.py b/tests/test_keys.py index 5a7b79e1cb3..413c24bc159 100644 --- a/tests/test_keys.py +++ b/tests/test_keys.py @@ -6,6 +6,7 @@ import asyncio, time import aiohttp from openai import AsyncOpenAI import sys, os +from typing import Optional sys.path.insert( 0, os.path.abspath("../") @@ -19,6 +20,7 @@ async def generate_key( budget=None, budget_duration=None, models=["azure-models", "gpt-4", "dall-e-3"], + max_parallel_requests: Optional[int] = None, ): url = "http://0.0.0.0:4000/key/generate" headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"} @@ -28,6 +30,7 @@ async def generate_key( "duration": None, "max_budget": budget, "budget_duration": budget_duration, + "max_parallel_requests": max_parallel_requests, } print(f"data: {data}") @@ -524,3 +527,29 @@ async def test_key_info_spend_values_sagemaker(): rounded_key_info_spend = round(key_info["info"]["spend"], 8) assert rounded_key_info_spend > 0 # assert rounded_response_cost == rounded_key_info_spend + + +@pytest.mark.asyncio +async def test_key_rate_limit(): + """ + Tests backoff/retry logic on parallel request error. + - Create key with max parallel requests 0 + - run 2 requests -> both fail + - Create key with max parallel request 1 + - run 2 requests + - both should succeed + """ + async with aiohttp.ClientSession() as session: + key_gen = await generate_key(session=session, i=0, max_parallel_requests=0) + new_key = key_gen["key"] + try: + await chat_completion(session=session, key=new_key) + pytest.fail(f"Expected this call to fail") + except Exception as e: + pass + key_gen = await generate_key(session=session, i=0, max_parallel_requests=1) + new_key = key_gen["key"] + try: + await chat_completion(session=session, key=new_key) + except Exception as e: + pytest.fail(f"Expected this call to work - {str(e)}") From a3a751ce6213bf2907fc3d4c45a9f8903e566662 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Tue, 5 Mar 2024 20:45:16 -0800 Subject: [PATCH 2/5] fix(utils.py): fix mistral api exception mapping --- litellm/utils.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/utils.py b/litellm/utils.py index 68dc137afbc..ba40cdcf5af 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -6658,7 +6658,7 @@ def exception_type( message=f"{exception_provider} - {message}", llm_provider=custom_llm_provider, model=model, - response=httpx.Response(status_code=500, request=_request), + request=_request, ) elif hasattr(original_exception, "status_code"): exception_mapping_worked = True From 8a4a14cc95df7cb99745fb17110fe69ae1fb8ee8 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Tue, 5 Mar 2024 21:12:50 -0800 Subject: [PATCH 3/5] test(test_caching.py): fix test to check on id --- litellm/tests/test_caching.py | 15 ++++++--------- 1 file changed, 6 insertions(+), 9 deletions(-) diff --git a/litellm/tests/test_caching.py b/litellm/tests/test_caching.py index f649bff0278..0e5a7ab5f6a 100644 --- a/litellm/tests/test_caching.py +++ b/litellm/tests/test_caching.py @@ -438,11 +438,10 @@ def test_redis_cache_completion_stream(): temperature=0.2, stream=True, ) - response_1_content = "" + response_1_id = "" for chunk in response1: print(chunk) - response_1_content += chunk.choices[0].delta.content or "" - print(response_1_content) + response_1_id = chunk.id time.sleep(0.5) response2 = completion( model="gpt-3.5-turbo", @@ -451,15 +450,13 @@ def test_redis_cache_completion_stream(): temperature=0.2, stream=True, ) - response_2_content = "" + response_2_id = "" for chunk in response2: print(chunk) - response_2_content += chunk.choices[0].delta.content or "" - print("\nresponse 1", response_1_content) - print("\nresponse 2", response_2_content) + response_2_id += chunk.id assert ( - response_1_content == response_2_content - ), f"Response 1 != Response 2. Same params, Response 1{response_1_content} != Response 2{response_2_content}" + response_1_id == response_2_id + ), f"Response 1 != Response 2. Same params, Response 1{response_1_id} != Response 2{response_2_id}" litellm.success_callback = [] litellm.cache = None litellm.success_callback = [] From 7d824225a5fcccbc5887c081f2ef2566d8f4c61c Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Tue, 5 Mar 2024 21:37:59 -0800 Subject: [PATCH 4/5] fix(utils.py): set status code for api error --- litellm/utils.py | 1 + 1 file changed, 1 insertion(+) diff --git a/litellm/utils.py b/litellm/utils.py index ba40cdcf5af..c2b1a730f1e 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -6655,6 +6655,7 @@ def exception_type( method="POST", url="https://api.openai.com/v1" ) raise APIError( + status_code=500, message=f"{exception_provider} - {message}", llm_provider=custom_llm_provider, model=model, From 7e3734d037dbde3538d207082d5e03a8211dbea7 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 6 Mar 2024 19:05:39 -0800 Subject: [PATCH 5/5] test(test_completion.py): handle gemini timeout error --- litellm/tests/test_completion.py | 2 ++ litellm/utils.py | 5 ++++- 2 files changed, 6 insertions(+), 1 deletion(-) diff --git a/litellm/tests/test_completion.py b/litellm/tests/test_completion.py index f5e145769bd..8e695f3f713 100644 --- a/litellm/tests/test_completion.py +++ b/litellm/tests/test_completion.py @@ -2139,6 +2139,8 @@ async def test_acompletion_gemini(): response = await litellm.acompletion(model=model_name, messages=messages) # Add any assertions here to check the response print(f"response: {response}") + except litellm.Timeout as e: + pass except litellm.APIError as e: pass except Exception as e: diff --git a/litellm/utils.py b/litellm/utils.py index c2b1a730f1e..d09f330a1d5 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -7105,7 +7105,10 @@ def exception_type( llm_provider="palm", response=original_exception.response, ) - if "504 Deadline expired before operation could complete." in error_str: + if ( + "504 Deadline expired before operation could complete." in error_str + or "504 Deadline Exceeded" in error_str + ): exception_mapping_worked = True raise Timeout( message=f"PalmException - {original_exception.message}",