mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(budgets): preserve explicit reset interval clears
This commit is contained in:
parent
559247fa84
commit
f5f81e973a
15 changed files with 63 additions and 24 deletions
|
|
@ -14,6 +14,7 @@ All /budget management endpoints
|
|||
#### BUDGET TABLE MANAGEMENT ####
|
||||
import math
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
|
@ -176,6 +177,10 @@ async def update_budget(
|
|||
recomputed_reset_at: Final = (
|
||||
{"budget_reset_at": get_budget_reset_time(budget_duration=budget_obj.budget_duration)}
|
||||
if budget_obj.budget_duration is not None and "budget_reset_at" not in budget_obj.model_fields_set
|
||||
else MappingProxyType({"budget_reset_at": None})
|
||||
if "budget_duration" in budget_obj.model_fields_set
|
||||
and budget_obj.budget_duration is None
|
||||
and "budget_reset_at" not in budget_obj.model_fields_set
|
||||
else {}
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -1254,8 +1254,8 @@ def _update_internal_user_params(data_json: dict, data: UpdateUserRequest | Upda
|
|||
fields_set: Final = data.fields_set() if hasattr(data, "fields_set") else set()
|
||||
|
||||
for k, v in data_json.items():
|
||||
if k == "max_budget":
|
||||
if "max_budget" in fields_set:
|
||||
if k in ("max_budget", "budget_duration"):
|
||||
if k in fields_set:
|
||||
non_default_values[k] = v
|
||||
elif k == "model_max_budget":
|
||||
if k in fields_set:
|
||||
|
|
@ -1283,8 +1283,10 @@ def _update_internal_user_params(data_json: dict, data: UpdateUserRequest | Upda
|
|||
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
|
||||
|
||||
validate_budget_duration(non_default_values["budget_duration"])
|
||||
non_default_values["budget_reset_at"] = get_budget_reset_time(
|
||||
budget_duration=non_default_values["budget_duration"]
|
||||
non_default_values["budget_reset_at"] = (
|
||||
get_budget_reset_time(budget_duration=non_default_values["budget_duration"])
|
||||
if non_default_values["budget_duration"] is not None
|
||||
else None
|
||||
)
|
||||
|
||||
if "max_budget" not in non_default_values:
|
||||
|
|
|
|||
|
|
@ -438,6 +438,7 @@ async def update_tag(
|
|||
user_api_key_dict=user_api_key_dict,
|
||||
prisma_client=prisma_client,
|
||||
litellm_proxy_admin_name=litellm_proxy_admin_name,
|
||||
budget_duration_cleared="budget_duration" in tag.model_fields_set and tag.budget_duration is None,
|
||||
)
|
||||
|
||||
# Get model names for model_info
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@
|
|||
from collections.abc import Callable, Mapping, MutableMapping, Sequence
|
||||
from datetime import datetime
|
||||
from functools import wraps
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Final, Protocol
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
|
|
@ -180,6 +181,7 @@ async def handle_budget_for_entity(
|
|||
user_api_key_dict: UserAPIKeyAuth,
|
||||
prisma_client: PrismaClient,
|
||||
litellm_proxy_admin_name: str,
|
||||
budget_duration_cleared: bool = False,
|
||||
) -> str | None:
|
||||
"""
|
||||
Common helper to handle budget creation/updates for entities (organizations, tags, etc).
|
||||
|
|
@ -208,7 +210,14 @@ async def handle_budget_for_entity(
|
|||
|
||||
# Extract budget fields from data
|
||||
_json_data: Final = data.model_dump(exclude_none=True) if hasattr(data, "model_dump") else data
|
||||
_budget_data: Final = {k: v for k, v in _json_data.items() if k in budget_params}
|
||||
_budget_data: Final = MappingProxyType(
|
||||
{
|
||||
k: _json_data.get(k)
|
||||
for k in budget_params
|
||||
if k in _json_data
|
||||
or (k == "budget_duration" and existing_budget_id is not None and budget_duration_cleared)
|
||||
}
|
||||
)
|
||||
|
||||
# Check if budget_id is explicitly provided in the data
|
||||
data_budget_id: Final[str | None] = getattr(data, "budget_id", None)
|
||||
|
|
|
|||
|
|
@ -340,7 +340,8 @@ async def test_update_budget_recomputes_reset_at_when_duration_changes(
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_budget_preserves_explicit_reset_at(client_and_mocks):
|
||||
@pytest.mark.parametrize("budget_duration", ["1d", None])
|
||||
async def test_update_budget_preserves_explicit_reset_at(client_and_mocks, budget_duration):
|
||||
"""An explicit budget_reset_at from the caller always wins over recompute."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
captured = _capture_update_data(mock_table)
|
||||
|
|
@ -350,7 +351,7 @@ async def test_update_budget_preserves_explicit_reset_at(client_and_mocks):
|
|||
"/budget/update",
|
||||
json={
|
||||
"budget_id": "budget_explicit_reset",
|
||||
"budget_duration": "1d",
|
||||
"budget_duration": budget_duration,
|
||||
"budget_reset_at": explicit.isoformat(),
|
||||
},
|
||||
)
|
||||
|
|
@ -377,8 +378,7 @@ async def test_update_budget_without_duration_leaves_reset_at_untouched(
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_budget_duration_none_does_not_recompute(client_and_mocks):
|
||||
"""Clearing budget_duration (explicit null) must not recompute against a None duration."""
|
||||
async def test_update_budget_duration_none_clears_obsolete_reset(client_and_mocks):
|
||||
client, _, mock_table = client_and_mocks
|
||||
captured = _capture_update_data(mock_table)
|
||||
|
||||
|
|
@ -389,7 +389,7 @@ async def test_update_budget_duration_none_does_not_recompute(client_and_mocks):
|
|||
assert resp.status_code == 200, resp.text
|
||||
|
||||
assert "budget_duration" in captured and captured["budget_duration"] is None
|
||||
assert "budget_reset_at" not in captured
|
||||
assert captured["budget_reset_at"] is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -2097,6 +2097,22 @@ def test_update_internal_user_params_reset_max_budget_with_none():
|
|||
assert non_default_values["user_id"] == "test_user"
|
||||
|
||||
|
||||
def test_update_internal_user_params_explicit_duration_clear_overrides_role_default(monkeypatch):
|
||||
import litellm
|
||||
|
||||
monkeypatch.setattr(litellm, "internal_user_budget_duration", "30d")
|
||||
data = UpdateUserRequest(
|
||||
user_id="duration-clear-test",
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
budget_duration=None,
|
||||
)
|
||||
|
||||
updated = _update_internal_user_params(data_json=data.model_dump(exclude_unset=True), data=data)
|
||||
|
||||
assert updated["budget_duration"] is None
|
||||
assert updated["budget_reset_at"] is None
|
||||
|
||||
|
||||
def test_update_internal_user_params_ignores_other_nones():
|
||||
"""
|
||||
Test that other fields are still filtered out if None
|
||||
|
|
|
|||
|
|
@ -98,7 +98,11 @@ const AccessGroupBudgetModal: React.FC<AccessGroupBudgetModalProps> = ({
|
|||
)}
|
||||
>
|
||||
{({ id, value, onChange }) => (
|
||||
<BudgetDurationDropdown id={id} value={value || null} onChange={onChange} />
|
||||
<BudgetDurationDropdown
|
||||
id={id}
|
||||
value={value || null}
|
||||
onChange={(next) => onChange(next ?? undefined)}
|
||||
/>
|
||||
)}
|
||||
</FormField>
|
||||
</FieldGroup>
|
||||
|
|
|
|||
|
|
@ -143,7 +143,11 @@ const CreateTagModal: React.FC<CreateTagModalProps> = ({ visible, onCancel, onSu
|
|||
)}
|
||||
>
|
||||
{({ id, value, onChange }) => (
|
||||
<BudgetDurationDropdown id={id} value={value ?? null} onChange={onChange} />
|
||||
<BudgetDurationDropdown
|
||||
id={id}
|
||||
value={value ?? null}
|
||||
onChange={(next) => onChange(next ?? undefined)}
|
||||
/>
|
||||
)}
|
||||
</FormField>
|
||||
</FieldGroup>
|
||||
|
|
|
|||
|
|
@ -28,7 +28,7 @@ const tagEditShape = {
|
|||
description: z.string().optional(),
|
||||
models: z.array(z.string()).optional(),
|
||||
max_budget: z.union([z.string(), z.number()]).optional(),
|
||||
budget_duration: z.string().optional(),
|
||||
budget_duration: z.string().nullish(),
|
||||
};
|
||||
|
||||
const tagEditSchema = z.object(tagEditShape);
|
||||
|
|
|
|||
|
|
@ -807,7 +807,7 @@ const Teams: React.FC<TeamProps> = ({ accessToken, userID, userRole, premiumUser
|
|||
showNeverResets
|
||||
placeholder={budgetDurationPlaceholder}
|
||||
value={value}
|
||||
onChange={onChange}
|
||||
onChange={(next) => onChange(next ?? undefined)}
|
||||
/>
|
||||
)}
|
||||
</FormField>
|
||||
|
|
|
|||
|
|
@ -14,7 +14,7 @@ const DURATION_LABELS: Record<string, string> = {
|
|||
interface BudgetDurationDropdownProps {
|
||||
id?: string;
|
||||
value?: string | null;
|
||||
onChange?: (value: string | undefined) => void;
|
||||
onChange?: (value: string | null) => void;
|
||||
className?: string;
|
||||
style?: React.CSSProperties;
|
||||
placeholder?: string;
|
||||
|
|
@ -31,11 +31,7 @@ const BudgetDurationDropdown: React.FC<BudgetDurationDropdownProps> = ({
|
|||
showNeverResets = false,
|
||||
}) => {
|
||||
return (
|
||||
<Select
|
||||
items={DURATION_LABELS}
|
||||
value={value || null}
|
||||
onValueChange={(next: string | null) => onChange?.(next ?? undefined)}
|
||||
>
|
||||
<Select items={DURATION_LABELS} value={value || null} onValueChange={onChange}>
|
||||
<SelectTrigger id={id} className={`w-full ${className}`} style={style}>
|
||||
<SelectValue placeholder={placeholder} />
|
||||
</SelectTrigger>
|
||||
|
|
|
|||
|
|
@ -1021,7 +1021,7 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey, autoOp
|
|||
value={control.value as string | null | undefined}
|
||||
showNeverResets
|
||||
placeholder="Not set"
|
||||
onChange={control.onChange}
|
||||
onChange={(next) => control.onChange(next ?? undefined)}
|
||||
/>
|
||||
)}
|
||||
</MountedFormField>
|
||||
|
|
|
|||
|
|
@ -41,7 +41,7 @@ export interface TagUpdateRequest {
|
|||
soft_budget?: number;
|
||||
tpm_limit?: number;
|
||||
rpm_limit?: number;
|
||||
budget_duration?: string;
|
||||
budget_duration?: string | null;
|
||||
}
|
||||
|
||||
export interface TagDeleteRequest {
|
||||
|
|
|
|||
|
|
@ -157,7 +157,7 @@ const MemberModal = <T extends BaseMember>({
|
|||
<BudgetDurationDropdown
|
||||
id={id}
|
||||
value={typeof value === "string" ? value : null}
|
||||
onChange={(next) => onChange(next)}
|
||||
onChange={(next) => onChange(mode === "add" ? next ?? undefined : next)}
|
||||
/>
|
||||
);
|
||||
default:
|
||||
|
|
|
|||
|
|
@ -1467,7 +1467,9 @@ const TeamInfoView: React.FC<TeamInfoProps> = ({
|
|||
showNeverResets
|
||||
placeholder="Inherit team reset period"
|
||||
value={value === null ? NEVER_RESETS_BUDGET_DURATION : value}
|
||||
onChange={(next) => onChange(next === NEVER_RESETS_BUDGET_DURATION ? null : next)}
|
||||
onChange={(next) =>
|
||||
onChange(next === NEVER_RESETS_BUDGET_DURATION ? null : next ?? undefined)
|
||||
}
|
||||
/>
|
||||
)}
|
||||
</FormField>
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue