test(budgets): avoid mutable fixture state

This commit is contained in:
Yuneng Jiang 2026-09-16 22:06:31 -07:00
parent 5c41e0b8dc
commit a0869fe835
No known key found for this signature in database
3 changed files with 60 additions and 32 deletions

View file

@ -408,14 +408,21 @@ def test_update_customer_budget_omission_and_null_preserve_existing_budget(
):
from litellm.proxy._types import LiteLLM_BudgetTable
budget_state = {"budget_id": "budget-1", "max_budget": 100.0}
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_state),
litellm_budget_table=LiteLLM_BudgetTable(budget_id="budget-1", max_budget=budget_state.max_budget),
)
def response_row():
@ -426,15 +433,15 @@ def test_update_customer_budget_omission_and_null_preserve_existing_budget(
"budget_id": "budget-1",
"litellm_budget_table": {
"budget_id": "budget-1",
"max_budget": budget_state["max_budget"],
"max_budget": budget_state.max_budget,
"created_at": "2024-01-01T00:00:00",
},
}
return row
async def update_budget(*, where, data):
budget_state.update(data)
return LiteLLM_BudgetTable(**budget_state)
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)

View file

@ -691,7 +691,14 @@ async def test_organization_member_update_budget_omission_and_null_preserve_exis
from litellm.proxy._types import LitellmUserRoles, OrganizationMemberUpdateRequest, UserAPIKeyAuth
from litellm.proxy.management_endpoints import organization_endpoints
budget_state = {"max_budget": 100.0}
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()
@ -705,14 +712,14 @@ async def test_organization_member_update_budget_omission_and_null_preserve_exis
"budget_id": "budget-1",
"created_at": datetime(2024, 1, 1),
"updated_at": datetime(2024, 1, 1),
"litellm_budget_table": {"budget_id": "budget-1", **budget_state},
"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["max_budget"] = budget_obj.max_budget
budget_state.store(budget_obj.max_budget)
mock_db = SimpleNamespace(
litellm_organizationtable=SimpleNamespace(find_unique=AsyncMock(return_value=SimpleNamespace())),

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``.
@ -219,11 +234,10 @@ async def test_update_tag():
@pytest.mark.asyncio
async def test_new_tag_persists_a_budget():
from datetime import datetime
from types import SimpleNamespace
from litellm.proxy.management_endpoints.tag_management_endpoints import new_tag
budget_state = {"budget_id": "budget-1", "max_budget": None}
budget_state = _BudgetState({"budget_id": "budget-1", "max_budget": None})
created_tag = SimpleNamespace(
tag_name="budget-tag",
description=None,
@ -238,8 +252,8 @@ async def test_new_tag_persists_a_budget():
mock_db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
async def create_budget(data, **_):
budget_state.update(data)
return SimpleNamespace(**budget_state)
budget_state.store(data)
return budget_state.row()
async def create_tag(data, **_):
created_tag.budget_id = data["budget_id"]
@ -266,7 +280,7 @@ async def test_new_tag_persists_a_budget():
user_api_key_dict=UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN),
)
assert budget_state["max_budget"] == 25.0
assert budget_state.get("max_budget") == 25.0
assert created_tag.budget_id == "budget-1"
@ -277,20 +291,21 @@ async def test_new_tag_persists_a_budget():
)
async def test_update_tag_explicit_null_preserves_general_budget_fields(field):
from datetime import datetime
from types import SimpleNamespace
from litellm.proxy.management_endpoints.tag_management_endpoints import update_tag
from litellm.types.tag_management import TagUpdateRequest
budget_state = {
"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",
}
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",
@ -307,8 +322,8 @@ async def test_update_tag_explicit_null_preserves_general_budget_fields(field):
mock_db.litellm_tagtable.update = AsyncMock(return_value=updated_tag)
async def update_budget(where, data, **_):
budget_state.update(data)
return SimpleNamespace(**budget_state)
budget_state.store(data)
return budget_state.row()
mock_db.litellm_budgettable.update = update_budget
with (
@ -334,18 +349,17 @@ async def test_update_tag_explicit_null_preserves_general_budget_fields(field):
"tpm_limit": 1000,
"rpm_limit": 100,
}
assert budget_state[field] == expected_values[field]
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 types import SimpleNamespace
from litellm.proxy.management_endpoints.tag_management_endpoints import update_tag
from litellm.types.tag_management import TagUpdateRequest
budget_state = {"budget_id": "budget-1", "budget_duration": "30d"}
budget_state = _BudgetState({"budget_id": "budget-1", "budget_duration": "30d"})
existing_tag = SimpleNamespace(budget_id="budget-1")
updated_tag = SimpleNamespace(
tag_name="budget-tag",
@ -362,8 +376,8 @@ async def test_update_tag_explicit_null_clears_budget_duration():
mock_db.litellm_tagtable.update = AsyncMock(return_value=updated_tag)
async def update_budget(where, data, **_):
budget_state.update(data)
return SimpleNamespace(**budget_state)
budget_state.store(data)
return budget_state.row()
mock_db.litellm_budgettable.update = update_budget
with (
@ -382,7 +396,7 @@ async def test_update_tag_explicit_null_clears_budget_duration():
user_api_key_dict=UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN),
)
assert budget_state["budget_duration"] is None
assert budget_state.get("budget_duration") is None
@pytest.mark.asyncio