mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
Merge pull request #41563 from BerriAI/litellm_budget-null-clear-tests
test(budgets): cover management null handling
This commit is contained in:
commit
f58389c0e8
4 changed files with 402 additions and 1 deletions
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue