diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index ac300b47635..75db608739c 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -2679,7 +2679,9 @@ def _get_temp_budget_increase(valid_token: UserAPIKeyAuth): valid_token_metadata = valid_token.metadata if "temp_budget_increase" in valid_token_metadata and "temp_budget_expiry" in valid_token_metadata: expiry = datetime.fromisoformat(valid_token_metadata["temp_budget_expiry"]) - if expiry > datetime.now(): + if expiry.tzinfo is None: + expiry = expiry.replace(tzinfo=timezone.utc) + if expiry > datetime.now(timezone.utc): return valid_token_metadata["temp_budget_increase"] return None diff --git a/tests/proxy_unit_tests/test_proxy_utils.py b/tests/proxy_unit_tests/test_proxy_utils.py index d36d73da2c3..d22b343d843 100644 --- a/tests/proxy_unit_tests/test_proxy_utils.py +++ b/tests/proxy_unit_tests/test_proxy_utils.py @@ -1732,6 +1732,35 @@ def test_get_temp_budget_increase(): assert _get_temp_budget_increase(valid_token) == 100 +def test_get_temp_budget_increase_tz_aware_expiry(): + from datetime import datetime, timedelta, timezone + + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import _get_temp_budget_increase + + future_expiry = (datetime.now(timezone.utc) + timedelta(days=1)).isoformat() + valid_token = UserAPIKeyAuth( + max_budget=100, + spend=0, + metadata={ + "temp_budget_increase": 100, + "temp_budget_expiry": future_expiry, + }, + ) + assert _get_temp_budget_increase(valid_token) == 100 + + past_expiry = (datetime.now(timezone.utc) - timedelta(days=1)).isoformat() + expired_token = UserAPIKeyAuth( + max_budget=100, + spend=0, + metadata={ + "temp_budget_increase": 100, + "temp_budget_expiry": past_expiry, + }, + ) + assert _get_temp_budget_increase(expired_token) is None + + def test_update_key_budget_with_temp_budget_increase(): from datetime import datetime, timedelta