mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(lowest_tpm_rpm_v2.py): fix updating limits
This commit is contained in:
parent
cfe94c86cc
commit
39ac9e3eca
3 changed files with 33 additions and 12 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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": {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue