mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
2c99b7804e
commit
46a80e1ef5
2 changed files with 54 additions and 1 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue