diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 6ed283d898b..99a867a5d07 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -66,6 +66,7 @@ from litellm.proxy.auth.budget_throttle import ( ) from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec +from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time from litellm.proxy.common_utils.http_parsing_utils import ( _safe_get_request_headers, _safe_get_request_query_params, @@ -1677,6 +1678,13 @@ async def get_user_object( new_user_params["user_email"] = user_email if litellm.default_internal_user_params is not None: new_user_params.update(litellm.default_internal_user_params) + if ( + new_user_params.get("budget_duration") is not None + and new_user_params.get("budget_reset_at") is None + ): + new_user_params["budget_reset_at"] = get_budget_reset_time( + budget_duration=new_user_params["budget_duration"] + ) response = await UserRepository(prisma_client).table.create( data=new_user_params, diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 2da645bf4e1..5e07d1bcbc5 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -8,7 +8,7 @@ sys.path.insert( 0, os.path.abspath("../../..") ) # Adds the parent directory to the system path -from datetime import datetime, timedelta +from datetime import datetime, timedelta, timezone import httpx import pytest @@ -744,6 +744,51 @@ async def test_default_internal_user_params_with_get_user_object(monkeypatch): assert creation_args["user_role"] == "internal_user" +@pytest.mark.asyncio +@pytest.mark.parametrize("has_budget_duration", [True, False]) +async def test_get_user_object_upsert_sets_budget_reset_at(monkeypatch, has_budget_duration): + """The JWT first-login upsert must compute budget_reset_at when + default_internal_user_params carries a budget_duration; otherwise the row + lands with budget_reset_at=NULL and shows a null reset time until the next + reset sweep heals it. Without a budget_duration, no reset time is written.""" + default_params = {"max_budget": 300.0} + if has_budget_duration: + default_params["budget_duration"] = "24h" + monkeypatch.setattr(litellm, "default_internal_user_params", default_params) + + mock_prisma_client = MagicMock() + mock_prisma_client.db = AsyncMock() + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None) + mock_prisma_client.db.litellm_usertable.create = AsyncMock(return_value=MagicMock(organization_memberships=[])) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + user_id = f"jwt_upsert_reset_at_{has_budget_duration}" + try: + await get_user_object( + user_id=user_id, + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + user_id_upsert=True, + proxy_logging_obj=None, + ) + except Exception as e: + print(e) + + mock_prisma_client.db.litellm_usertable.create.assert_called_once() + creation_args = mock_prisma_client.db.litellm_usertable.create.call_args[1]["data"] + + if has_budget_duration: + reset_at = creation_args.get("budget_reset_at") + assert isinstance(reset_at, datetime), f"expected a computed budget_reset_at, got {creation_args!r}" + assert reset_at > datetime.now(timezone.utc) + else: + assert "budget_reset_at" not in creation_args + + @pytest.mark.asyncio async def test_get_user_object_wraps_db_outage_as_valueerror_preserving_context(): """Pin get_user_object's exception contract: it catches every DB failure in a broad except and