diff --git a/litellm/router_strategy/lowest_tpm_rpm_v2.py b/litellm/router_strategy/lowest_tpm_rpm_v2.py index bc7cb46b07e..e38f3b53365 100644 --- a/litellm/router_strategy/lowest_tpm_rpm_v2.py +++ b/litellm/router_strategy/lowest_tpm_rpm_v2.py @@ -65,10 +65,13 @@ class LowestTPMLoggingHandler_v2(CustomLogger): # ------------ # Setup values # ------------ + dt = get_utc_datetime() current_minute = dt.strftime("%H-%M") model_id = deployment.get("model_info", {}).get("id") - rpm_key = f"{model_id}:rpm:{current_minute}" + deployment_name = deployment.get("litellm_params", {}).get("model") + rpm_key = f"{model_id}:{deployment_name}:rpm:{current_minute}" + local_result = self.router_cache.get_cache( key=rpm_key, local_only=True ) # check local result first @@ -230,9 +233,9 @@ class LowestTPMLoggingHandler_v2(CustomLogger): if standard_logging_object is None: raise ValueError("standard_logging_object not passed in.") model_group = standard_logging_object.get("model_group") - model = standard_logging_object.get("model") + model = standard_logging_object["hidden_params"].get("litellm_model_name") id = standard_logging_object.get("model_id") - if model_group is None or id is None: + if model_group is None or id is None or model is None: return elif isinstance(id, int): id = str(id) diff --git a/litellm/types/utils.py b/litellm/types/utils.py index a2d41d8fb9d..5ecd490ee89 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -1625,13 +1625,16 @@ class StandardLoggingAdditionalHeaders(TypedDict, total=False): class StandardLoggingHiddenParams(TypedDict): - model_id: Optional[str] + model_id: Optional[ + str + ] # id of the model in the router, separates multiple models with the same name but different credentials cache_key: Optional[str] api_base: Optional[str] response_cost: Optional[str] litellm_overhead_time_ms: Optional[float] additional_headers: Optional[StandardLoggingAdditionalHeaders] batch_models: Optional[List[str]] + litellm_model_name: Optional[str] # the model name sent to the provider by litellm class StandardLoggingModelInformation(TypedDict): diff --git a/tests/local_testing/test_tpm_rpm_routing_v2.py b/tests/local_testing/test_tpm_rpm_routing_v2.py index a7073b4acd8..cb6e77de7dd 100644 --- a/tests/local_testing/test_tpm_rpm_routing_v2.py +++ b/tests/local_testing/test_tpm_rpm_routing_v2.py @@ -20,7 +20,7 @@ sys.path.insert( from unittest.mock import AsyncMock, MagicMock, patch from litellm.types.utils import StandardLoggingPayload import pytest - +from litellm.types.router import DeploymentTypedDict import litellm from litellm import Router from litellm.caching.caching import DualCache @@ -47,12 +47,14 @@ def test_tpm_rpm_updated(): deployment_id = "1234" deployment = "azure/chatgpt-v-2" total_tokens = 50 - standard_logging_payload = create_standard_logging_payload() + standard_logging_payload: StandardLoggingPayload = create_standard_logging_payload() standard_logging_payload["model_group"] = model_group standard_logging_payload["model_id"] = deployment_id standard_logging_payload["total_tokens"] = total_tokens + standard_logging_payload["hidden_params"]["litellm_model_name"] = deployment kwargs = { "litellm_params": { + "model": deployment, "metadata": { "model_group": model_group, "deployment": deployment, @@ -62,10 +64,16 @@ def test_tpm_rpm_updated(): "standard_logging_object": standard_logging_payload, } + litellm_deployment_dict: DeploymentTypedDict = { + "model_name": model_group, + "litellm_params": {"model": deployment}, + "model_info": {"id": deployment_id}, + } + start_time = time.time() response_obj = {"usage": {"total_tokens": total_tokens}} end_time = time.time() - lowest_tpm_logger.pre_call_check(deployment=kwargs["litellm_params"]) + lowest_tpm_logger.pre_call_check(deployment=litellm_deployment_dict) lowest_tpm_logger.log_success_event( response_obj=response_obj, kwargs=kwargs, @@ -74,8 +82,8 @@ def test_tpm_rpm_updated(): ) dt = get_utc_datetime() current_minute = dt.strftime("%H-%M") - tpm_count_api_key = f"{deployment_id}:tpm:{current_minute}" - rpm_count_api_key = f"{deployment_id}:rpm:{current_minute}" + tpm_count_api_key = f"{deployment_id}:{deployment}:tpm:{current_minute}" + rpm_count_api_key = f"{deployment_id}:{deployment}:rpm:{current_minute}" print(f"tpm_count_api_key={tpm_count_api_key}") assert response_obj["usage"]["total_tokens"] == test_cache.get_cache( @@ -113,6 +121,7 @@ def test_get_available_deployments(): standard_logging_payload["model_group"] = model_group standard_logging_payload["model_id"] = deployment_id standard_logging_payload["total_tokens"] = total_tokens + standard_logging_payload["hidden_params"]["litellm_model_name"] = deployment kwargs = { "litellm_params": { "metadata": { @@ -135,10 +144,11 @@ def test_get_available_deployments(): ## DEPLOYMENT 2 ## total_tokens = 20 deployment_id = "5678" - standard_logging_payload = create_standard_logging_payload() + standard_logging_payload: StandardLoggingPayload = create_standard_logging_payload() standard_logging_payload["model_group"] = model_group standard_logging_payload["model_id"] = deployment_id standard_logging_payload["total_tokens"] = total_tokens + standard_logging_payload["hidden_params"]["litellm_model_name"] = deployment kwargs = { "litellm_params": { "metadata": { @@ -209,11 +219,12 @@ def test_router_get_available_deployments(): print(f"router id's: {router.get_model_ids()}") ## DEPLOYMENT 1 ## deployment_id = 1 - standard_logging_payload = create_standard_logging_payload() + standard_logging_payload: StandardLoggingPayload = create_standard_logging_payload() standard_logging_payload["model_group"] = "azure-model" standard_logging_payload["model_id"] = str(deployment_id) total_tokens = 50 standard_logging_payload["total_tokens"] = total_tokens + standard_logging_payload["hidden_params"]["litellm_model_name"] = "azure/gpt-turbo" kwargs = { "litellm_params": { "metadata": { @@ -237,6 +248,9 @@ def test_router_get_available_deployments(): standard_logging_payload = create_standard_logging_payload() standard_logging_payload["model_group"] = "azure-model" standard_logging_payload["model_id"] = str(deployment_id) + standard_logging_payload["hidden_params"][ + "litellm_model_name" + ] = "azure/gpt-35-turbo" kwargs = { "litellm_params": { "metadata": { @@ -293,10 +307,11 @@ def test_router_skip_rate_limited_deployments(): ## DEPLOYMENT 1 ## deployment_id = 1 total_tokens = 1439 - standard_logging_payload = create_standard_logging_payload() + standard_logging_payload: StandardLoggingPayload = create_standard_logging_payload() standard_logging_payload["model_group"] = "azure-model" standard_logging_payload["model_id"] = str(deployment_id) standard_logging_payload["total_tokens"] = total_tokens + standard_logging_payload["hidden_params"]["litellm_model_name"] = "azure/gpt-turbo" kwargs = { "litellm_params": { "metadata": {