Merge pull request #41563 from BerriAI/litellm_budget-null-clear-tests

test(budgets): cover management null handling
This commit is contained in:
yuneng-jiang 2026-09-16 23:02:04 -07:00 • committed by GitHub
commit f58389c0e8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 402 additions and 1 deletions

View file

@ -929,6 +929,34 @@ async def test_put_access_group_budget_rejects_an_empty_body():
assert cache.deleted_keys == []
@pytest.mark.asyncio
async def test_put_access_group_budget_rejects_explicit_null_max_budget():
from fastapi import HTTPException
from litellm.proxy.management_endpoints.model_access_group_management_endpoints import (
set_access_group_budget,
)
from litellm.types.proxy.management_endpoints.model_management_endpoints import (
AccessGroupBudgetRequest,
)
prisma = _FakePrismaClient([], deployments=[_deployment()])
cache = _FakeAuthCache()
with _proxy(prisma), pytest.raises(HTTPException) as exc_info:
await set_access_group_budget(
access_group="prod-models",
data=AccessGroupBudgetRequest(max_budget=None),
user_api_key_dict=_admin(),
auth_cache=cache,
)
assert exc_info.value.status_code == 400
assert prisma.access_group_budget_table.rows == {}
assert prisma.budget_table.create_calls == []
assert cache.deleted_keys == []
@pytest.mark.asyncio
async def test_put_access_group_budget_rejects_an_unparseable_duration():
"""An unparseable duration can only be discovered by the reset job, long after the write."""

View file

@ -398,6 +398,65 @@ def test_update_customer_response_preserves_budget_id(mock_prisma_client, mock_u
assert response.json()["budget_id"] == "budget-123"
@pytest.mark.parametrize(
"budget_payload",
[{"max_budget": None}, {}],
ids=["explicit-null", "omitted"],
)
def test_update_customer_budget_omission_and_null_preserve_existing_budget(
mock_prisma_client, mock_user_api_key_auth, budget_payload
):
from litellm.proxy._types import LiteLLM_BudgetTable
class BudgetState:
def __init__(self) -> None:
self.max_budget: float | None = 100.0
def store(self, data) -> None:
self.max_budget = data.get("max_budget", self.max_budget)
budget_state = BudgetState()
def end_user_row():
return LiteLLM_EndUserTable(
user_id="cust-1",
blocked=False,
budget_id="budget-1",
litellm_budget_table=LiteLLM_BudgetTable(budget_id="budget-1", max_budget=budget_state.max_budget),
)
def response_row():
row = MagicMock()
row.model_dump.return_value = {
"user_id": "cust-1",
"blocked": False,
"budget_id": "budget-1",
"litellm_budget_table": {
"budget_id": "budget-1",
"max_budget": budget_state.max_budget,
"created_at": "2024-01-01T00:00:00",
},
}
return row
async def update_budget(*, where, data):
budget_state.store(data)
return LiteLLM_BudgetTable(budget_id="budget-1", max_budget=budget_state.max_budget)
mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(return_value=end_user_row())
mock_prisma_client.db.litellm_budgettable.update = AsyncMock(side_effect=update_budget)
mock_prisma_client.db.litellm_endusertable.update = AsyncMock(side_effect=lambda **_: response_row())
response = client.post(
"/customer/update",
json={"user_id": "cust-1", **budget_payload},
headers={"Authorization": "Bearer test-key"},
)
assert response.status_code == 200, response.text
assert response.json()["litellm_budget_table"]["max_budget"] == 100.0
def test_update_customer_response_keeps_nested_budget_server_fields(mock_prisma_client, mock_user_api_key_auth):
"""
Faithfulness regression: /customer/update embeds the full budget row. The

View file

@ -621,6 +621,137 @@ async def test_organization_member_update_rejects_unauthorized_caller(patched_or
assert exc.value.status_code == 403
@pytest.mark.asyncio
@pytest.mark.parametrize(
"budget_payload",
[{"max_budget_in_organization": None}, {}],
ids=["explicit-null", "omitted"],
)
async def test_organization_member_add_budget_omission_and_null_leave_budget_unset(budget_payload, monkeypatch):
from datetime import datetime
from types import SimpleNamespace
from litellm.proxy._types import (
LiteLLM_OrganizationMembershipTable,
LiteLLM_UserTable,
LitellmUserRoles,
OrganizationMemberAddRequest,
UserAPIKeyAuth,
)
from litellm.proxy.management_endpoints.organization_endpoints import organization_member_add
user = LiteLLM_UserTable(user_id="user-1", user_role="internal_user")
async def create_membership(data):
return LiteLLM_OrganizationMembershipTable(
user_id="user-1",
organization_id="org-1",
user_role="internal_user",
budget_id=data.get("budget_id"),
created_at=datetime(2024, 1, 1),
updated_at=datetime(2024, 1, 1),
)
mock_db = SimpleNamespace(
litellm_organizationtable=SimpleNamespace(find_unique=AsyncMock(return_value=SimpleNamespace())),
litellm_usertable=SimpleNamespace(find_unique=AsyncMock(return_value=user)),
litellm_organizationmembership=SimpleNamespace(create=create_membership),
)
mock_prisma = SimpleNamespace(db=mock_db)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.organization_endpoints._verify_org_access",
AsyncMock(),
)
response = await organization_member_add(
data=OrganizationMemberAddRequest(
organization_id="org-1",
member={"role": "internal_user", "user_id": "user-1"},
**budget_payload,
),
http_request=MagicMock(),
user_api_key_dict=UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN),
)
assert response.updated_organization_memberships[0].budget_id is None
@pytest.mark.asyncio
@pytest.mark.parametrize(
"budget_payload",
[{"max_budget_in_organization": None}, {}],
ids=["explicit-null", "omitted"],
)
async def test_organization_member_update_budget_omission_and_null_preserve_existing_budget(
budget_payload, monkeypatch
):
from datetime import datetime
from types import SimpleNamespace
from litellm.proxy._types import LitellmUserRoles, OrganizationMemberUpdateRequest, UserAPIKeyAuth
from litellm.proxy.management_endpoints import organization_endpoints
class BudgetState:
def __init__(self) -> None:
self.max_budget: float | None = 100.0
def store(self, max_budget: float | None) -> None:
self.max_budget = max_budget
budget_state = BudgetState()
def membership_row():
row = MagicMock()
row.budget_id = "budget-1"
def dump(**_):
return {
"user_id": "user-1",
"organization_id": "org-1",
"user_role": "internal_user",
"budget_id": "budget-1",
"created_at": datetime(2024, 1, 1),
"updated_at": datetime(2024, 1, 1),
"litellm_budget_table": {"budget_id": "budget-1", "max_budget": budget_state.max_budget},
}
row.model_dump.side_effect = dump
return row
async def update_budget(*, budget_obj, user_api_key_dict):
budget_state.store(budget_obj.max_budget)
mock_db = SimpleNamespace(
litellm_organizationtable=SimpleNamespace(find_unique=AsyncMock(return_value=SimpleNamespace())),
litellm_organizationmembership=SimpleNamespace(
find_unique=AsyncMock(side_effect=[membership_row(), membership_row()]),
update=AsyncMock(),
),
litellm_usertable=SimpleNamespace(
find_unique=AsyncMock(return_value=SimpleNamespace(user_role="internal_user"))
),
)
mock_prisma = SimpleNamespace(db=mock_db)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
monkeypatch.setattr(organization_endpoints, "update_budget", update_budget)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.organization_endpoints._verify_org_access",
AsyncMock(),
)
response = await organization_endpoints.organization_member_update(
data=OrganizationMemberUpdateRequest(
organization_id="org-1",
user_id="user-1",
**budget_payload,
),
user_api_key_dict=UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN),
)
assert response.litellm_budget_table is not None
assert response.litellm_budget_table.max_budget == 100.0
@pytest.mark.asyncio
async def test_organization_member_delete_rejects_unauthorized_caller(patched_org_prisma, unauthorized_caller):
from litellm.proxy._types import OrganizationMemberDeleteRequest

View file

@ -1,7 +1,8 @@
import inspect
import json
from collections.abc import Sequence
from typing import Optional
from types import MappingProxyType, SimpleNamespace
from typing import Mapping, Optional
import pytest
from fastapi import HTTPException
@ -20,6 +21,20 @@ from litellm.types.tag_management import TagDeleteRequest, TagInfoRequest, TagNe
client = TestClient(app)
class _BudgetState:
def __init__(self, values: Mapping[str, object]) -> None:
self._values: Mapping[str, object] = MappingProxyType(dict(values))
def store(self, values: Mapping[str, object]) -> None:
self._values = MappingProxyType({**self._values, **values})
def get(self, field: str) -> object:
return self._values[field]
def row(self) -> SimpleNamespace:
return SimpleNamespace(**self._values)
class FakeVerificationTokenTable:
"""Stand-in for ``prisma_client.db.litellm_verificationtoken``.
@ -216,6 +231,174 @@ async def test_update_tag():
app.dependency_overrides.clear()
@pytest.mark.asyncio
async def test_new_tag_persists_a_budget():
from datetime import datetime
from litellm.proxy.management_endpoints.tag_management_endpoints import new_tag
budget_state = _BudgetState({"budget_id": "budget-1", "max_budget": None})
created_tag = SimpleNamespace(
tag_name="budget-tag",
description=None,
models=[],
created_at=datetime(2024, 1, 1),
updated_at=datetime(2024, 1, 1),
created_by="admin",
)
mock_db = Mock()
mock_prisma = SimpleNamespace(db=mock_db, jsonify_object=lambda data: dict(data))
mock_db.litellm_tagtable.find_unique = AsyncMock(return_value=None)
mock_db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
async def create_budget(data, **_):
budget_state.store(data)
return budget_state.row()
async def create_tag(data, **_):
created_tag.budget_id = data["budget_id"]
return created_tag
mock_db.litellm_budgettable.create = create_budget
mock_db.litellm_tagtable.create = create_tag
with (
patch( # test-quality-ok: endpoint resolves the fake database through proxy_server
"litellm.proxy.proxy_server.prisma_client", mock_prisma
),
patch( # test-quality-ok: endpoint reads the audit actor from proxy_server
"litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"
),
patch( # test-quality-ok: endpoint requires a router before the budget write
"litellm.proxy.proxy_server.llm_router", object()
),
patch( # test-quality-ok: cache invalidation is outside this budget contract
"litellm.proxy.management_endpoints.tag_management_endpoints._evict_tag_cache_keys", new=AsyncMock()
),
):
await new_tag(
tag=TagNewRequest(name="budget-tag", max_budget=25.0),
user_api_key_dict=UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN),
)
assert budget_state.get("max_budget") == 25.0
assert created_tag.budget_id == "budget-1"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"field",
["max_budget", "soft_budget", "model_max_budget", "tpm_limit", "rpm_limit"],
)
async def test_update_tag_explicit_null_preserves_general_budget_fields(field):
from datetime import datetime
from litellm.proxy.management_endpoints.tag_management_endpoints import update_tag
from litellm.types.tag_management import TagUpdateRequest
budget_state = _BudgetState(
{
"budget_id": "budget-1",
"max_budget": 100.0,
"soft_budget": 80.0,
"model_max_budget": {"model-a": {"max_budget": 50.0}},
"tpm_limit": 1000,
"rpm_limit": 100,
"budget_duration": "30d",
}
)
existing_tag = SimpleNamespace(budget_id="budget-1")
updated_tag = SimpleNamespace(
tag_name="budget-tag",
description=None,
models=[],
created_at=datetime(2024, 1, 1),
updated_at=datetime(2024, 1, 1),
created_by="admin",
)
mock_db = Mock()
mock_prisma = SimpleNamespace(db=mock_db)
mock_db.litellm_tagtable.find_unique = AsyncMock(return_value=existing_tag)
mock_db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
mock_db.litellm_tagtable.update = AsyncMock(return_value=updated_tag)
async def update_budget(where, data, **_):
budget_state.store(data)
return budget_state.row()
mock_db.litellm_budgettable.update = update_budget
with (
patch( # test-quality-ok: endpoint resolves the fake database through proxy_server
"litellm.proxy.proxy_server.prisma_client", mock_prisma
),
patch( # test-quality-ok: endpoint reads the audit actor from proxy_server
"litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"
),
patch( # test-quality-ok: cache invalidation is outside this budget contract
"litellm.proxy.management_endpoints.tag_management_endpoints._evict_tag_cache_keys", new=AsyncMock()
),
):
await update_tag(
tag=TagUpdateRequest(name="budget-tag", **{field: None}),
user_api_key_dict=UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN),
)
expected_values = {
"max_budget": 100.0,
"soft_budget": 80.0,
"model_max_budget": {"model-a": {"max_budget": 50.0}},
"tpm_limit": 1000,
"rpm_limit": 100,
}
assert budget_state.get(field) == expected_values[field]
@pytest.mark.asyncio
async def test_update_tag_explicit_null_clears_budget_duration():
from datetime import datetime
from litellm.proxy.management_endpoints.tag_management_endpoints import update_tag
from litellm.types.tag_management import TagUpdateRequest
budget_state = _BudgetState({"budget_id": "budget-1", "budget_duration": "30d"})
existing_tag = SimpleNamespace(budget_id="budget-1")
updated_tag = SimpleNamespace(
tag_name="budget-tag",
description=None,
models=[],
created_at=datetime(2024, 1, 1),
updated_at=datetime(2024, 1, 1),
created_by="admin",
)
mock_db = Mock()
mock_prisma = SimpleNamespace(db=mock_db)
mock_db.litellm_tagtable.find_unique = AsyncMock(return_value=existing_tag)
mock_db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
mock_db.litellm_tagtable.update = AsyncMock(return_value=updated_tag)
async def update_budget(where, data, **_):
budget_state.store(data)
return budget_state.row()
mock_db.litellm_budgettable.update = update_budget
with (
patch( # test-quality-ok: endpoint resolves the fake database through proxy_server
"litellm.proxy.proxy_server.prisma_client", mock_prisma
),
patch( # test-quality-ok: endpoint reads the audit actor from proxy_server
"litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"
),
patch( # test-quality-ok: cache invalidation is outside this budget contract
"litellm.proxy.management_endpoints.tag_management_endpoints._evict_tag_cache_keys", new=AsyncMock()
),
):
await update_tag(
tag=TagUpdateRequest(name="budget-tag", budget_duration=None),
user_api_key_dict=UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN),
)
assert budget_state.get("budget_duration") is None
@pytest.mark.asyncio
async def test_delete_tag():
"""