fix(lowest_tpm_rpm_v2.py): fix updating limits

This commit is contained in:
Krrish Dholakia 2025-03-18 17:10:17 -07:00
parent cfe94c86cc
commit 39ac9e3eca
3 changed files with 33 additions and 12 deletions

View file

@ -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)

View file

@ -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):

View file

@ -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": {