fix(auth): set budget_reset_at when JWT upsert seeds a budget_duration (#34050)

The JWT first-login upsert in get_user_object creates the user row by
merging default_internal_user_params straight into table.create, so a
configured budget_duration landed with budget_reset_at NULL. The reset
sweep now heals such rows (PR #33623), but until the next sweep the row
shows a null reset time and its first window starts at the sweep instead
of one full duration after creation. Compute budget_reset_at at creation
like every other write path (/user/new, UI SSO, /key/generate, /team/new)
already does
This commit is contained in:
ryan-crabbe-berri 2026-07-20 17:23:44 -07:00 • committed by GitHub
parent 2c99b7804e
commit 46a80e1ef5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 54 additions and 1 deletions

View file

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

View file

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