diff --git a/litellm/main.py b/litellm/main.py index 271c54e514d..69c9121893f 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -348,6 +348,13 @@ def mock_completion( prompt_tokens=10, completion_tokens=20, total_tokens=30 ) + try: + _, custom_llm_provider, _, _ = litellm.utils.get_llm_provider(model=model) + model_response._hidden_params["custom_llm_provider"] = custom_llm_provider + except: + # dont let setting a hidden param block a mock_respose + pass + return model_response except: diff --git a/litellm/router.py b/litellm/router.py index b15687f677e..38ebcc1c940 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -997,6 +997,9 @@ class Router: """ try: kwargs["model"] = mg + kwargs.setdefault("metadata", {}).update( + {"model_group": mg} + ) # update model_group used, if fallbacks are done response = await self.async_function_with_retries( *args, **kwargs ) @@ -1025,8 +1028,10 @@ class Router: f"Falling back to model_group = {mg}" ) kwargs["model"] = mg - kwargs["metadata"]["model_group"] = mg - response = await self.async_function_with_retries( + kwargs.setdefault("metadata", {}).update( + {"model_group": mg} + ) # update model_group used, if fallbacks are done + response = await self.async_function_with_fallbacks( *args, **kwargs ) return response @@ -1191,6 +1196,9 @@ class Router: ## LOGGING kwargs = self.log_retry(kwargs=kwargs, e=original_exception) kwargs["model"] = mg + kwargs.setdefault("metadata", {}).update( + {"model_group": mg} + ) # update model_group used, if fallbacks are done response = self.function_with_fallbacks(*args, **kwargs) return response except Exception as e: @@ -1214,6 +1222,9 @@ class Router: ## LOGGING kwargs = self.log_retry(kwargs=kwargs, e=original_exception) kwargs["model"] = mg + kwargs.setdefault("metadata", {}).update( + {"model_group": mg} + ) # update model_group used, if fallbacks are done response = self.function_with_fallbacks(*args, **kwargs) return response except Exception as e: diff --git a/litellm/tests/test_router_fallbacks.py b/litellm/tests/test_router_fallbacks.py index 29bc0d7bf1e..5d17d36c9f7 100644 --- a/litellm/tests/test_router_fallbacks.py +++ b/litellm/tests/test_router_fallbacks.py @@ -716,7 +716,7 @@ def test_usage_based_routing_fallbacks(): # Constants for TPM and RPM allocation AZURE_FAST_TPM = 3 AZURE_BASIC_TPM = 4 - OPENAI_TPM = 2000 + OPENAI_TPM = 400 ANTHROPIC_TPM = 100000 def get_azure_params(deployment_name: str): @@ -775,6 +775,7 @@ def test_usage_based_routing_fallbacks(): model_list=model_list, fallbacks=fallbacks_list, set_verbose=True, + debug_level="DEBUG", routing_strategy="usage-based-routing", redis_host=os.environ["REDIS_HOST"], redis_port=os.environ["REDIS_PORT"], @@ -783,17 +784,32 @@ def test_usage_based_routing_fallbacks(): messages = [ {"content": "Tell me a joke.", "role": "user"}, ] - response = router.completion( - model="azure/gpt-4-fast", messages=messages, timeout=5 + model="azure/gpt-4-fast", + messages=messages, + timeout=5, + mock_response="very nice to meet you", ) print("response: ", response) print("response._hidden_params: ", response._hidden_params) - # in this test, we expect azure/gpt-4 fast to fail, then azure-gpt-4 basic to fail and then openai-gpt-4 to pass # the token count of this message is > AZURE_FAST_TPM, > AZURE_BASIC_TPM assert response._hidden_params["custom_llm_provider"] == "openai" + # now make 100 mock requests to OpenAI - expect it to fallback to anthropic-claude-instant-1.2 + for i in range(20): + response = router.completion( + model="azure/gpt-4-fast", + messages=messages, + timeout=5, + mock_response="very nice to meet you", + ) + print("response: ", response) + print("response._hidden_params: ", response._hidden_params) + if i == 19: + # by the 19th call we should have hit TPM LIMIT for OpenAI, it should fallback to anthropic-claude-instant-1.2 + assert response._hidden_params["custom_llm_provider"] == "anthropic" + except Exception as e: pytest.fail(f"An exception occurred {e}")