This commit is contained in:
mubashir1osmani 2026-06-10 23:07:11 +08:00 • committed by GitHub
commit 701ffeb501
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 801 additions and 7 deletions

View file

@ -3,7 +3,9 @@ import asyncio
import aiohttp
import json
from httpx import AsyncClient
from typing import Any, Optional
from typing import Any, List, Optional
from litellm._uuid import uuid
async def make_calls_until_budget_exceeded(session, key: str, call_function, **kwargs):
@ -59,22 +61,99 @@ async def generate_key(
return await response.json()
async def chat_completion(session, key: str, model: str):
async def chat_completion(
session,
key: str,
model: str,
tags: Optional[List[str]] = None,
):
"""Make a chat completion request using OpenAI SDK"""
from openai import AsyncOpenAI
from litellm._uuid import uuid
client = AsyncOpenAI(
api_key=key, base_url="http://0.0.0.0:4000/v1" # Point to our local proxy
api_key=key, base_url="http://0.0.0.0:4000/v1"
)
extra_headers = None
if tags:
extra_headers = {"x-litellm-tags": ",".join(tags)}
response = await client.chat.completions.create(
model=model,
messages=[{"role": "user", "content": f"Say hello! {uuid.uuid4()}" * 100}],
extra_headers=extra_headers,
)
return response
async def create_budget(session, budget_id: str, max_budget: float, budget_duration: str):
url = "http://0.0.0.0:4000/budget/new"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
data = {
"budget_id": budget_id,
"max_budget": max_budget,
"budget_duration": budget_duration,
}
async with session.post(url, headers=headers, json=data) as response:
assert response.status == 200, await response.text()
return await response.json()
async def update_budget_max(session, budget_id: str, max_budget: float):
url = "http://0.0.0.0:4000/budget/update"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
data = {"budget_id": budget_id, "max_budget": max_budget}
async with session.post(url, headers=headers, json=data) as response:
assert response.status == 200, await response.text()
return await response.json()
async def delete_budget_by_id(session, budget_id: str):
url = "http://0.0.0.0:4000/budget/delete"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
data = {"id": budget_id}
async with session.post(url, headers=headers, json=data) as response:
assert response.status == 200, await response.text()
return await response.json()
async def generate_key_with_budget_id(session, budget_id: str):
url = "http://0.0.0.0:4000/key/generate"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
data = {"budget_id": budget_id}
async with session.post(url, headers=headers, json=data) as response:
assert response.status == 200, await response.text()
return await response.json()
async def create_tag_with_budget(
session,
tag_name: str,
max_budget: float,
models: Optional[List[str]] = None,
):
url = "http://0.0.0.0:4000/tag/new"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
data = {
"name": tag_name,
"description": "budget e2e tag",
"models": models or ["fake-openai-endpoint"],
"max_budget": max_budget,
}
async with session.post(url, headers=headers, json=data) as response:
assert response.status == 200, await response.text()
return await response.json()
async def reset_key_spend(session, key: str, reset_to: float = 0.0):
url = f"http://0.0.0.0:4000/key/{key}/reset_spend"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
data = {"reset_to": reset_to}
async with session.post(url, headers=headers, json=data) as response:
assert response.status == 200, await response.text()
return await response.json()
async def update_key_budget(session, key: str, max_budget: float):
"""Helper function to update a key's max budget"""
url = "http://0.0.0.0:4000/key/update"
@ -450,4 +529,113 @@ async def test_team_budget_update():
f"Request should succeed after team budget update but got error: {e}"
)
# Verify it was the team budget that was exceeded
@pytest.mark.asyncio
async def test_key_budget_duration_reset_unblocks_requests():
"""
Create a key with a 1d budget window, exhaust it, reset spend, verify requests work again.
"""
async with aiohttp.ClientSession() as session:
key_gen = await generate_key(
session=session,
max_budget=0.0000000005,
)
key = key_gen["key"]
url = "http://0.0.0.0:4000/key/update"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
async with session.post(
url,
headers=headers,
json={"key": key, "budget_duration": "1d"},
) as response:
assert response.status == 200, await response.text()
calls_made = await make_calls_until_budget_exceeded(
session=session,
key=key,
call_function=chat_completion,
model="fake-openai-endpoint",
)
assert calls_made > 0
await reset_key_spend(session, key, reset_to=0.0)
response = await chat_completion(
session=session,
key=key,
model="fake-openai-endpoint",
)
assert response is not None
@pytest.mark.asyncio
async def test_tag_budget_enforcement_with_header():
"""
Tag budget enforcement via x-litellm-tags on live proxy requests.
"""
async with aiohttp.ClientSession() as session:
tag_name = f"budget-tag-{uuid.uuid4()}"
await create_tag_with_budget(
session=session,
tag_name=tag_name,
max_budget=0.0000000005,
)
key_gen = await generate_key(session=session, max_budget=None)
key = key_gen["key"]
calls_made = await make_calls_until_budget_exceeded(
session=session,
key=key,
call_function=chat_completion,
model="fake-openai-endpoint",
tags=[tag_name],
)
assert calls_made > 0
@pytest.mark.asyncio
async def test_shared_budget_linking_unblocks_both_entities():
"""
Two keys linked to the same budget tier inherit the shared max_budget.
Updating the shared budget row unblocks both keys.
"""
async with aiohttp.ClientSession() as session:
budget_id = f"shared-budget-{uuid.uuid4()}"
await create_budget(
session=session,
budget_id=budget_id,
max_budget=0.0000000005,
budget_duration="1d",
)
key_gen_1 = await generate_key_with_budget_id(session, budget_id)
key_gen_2 = await generate_key_with_budget_id(session, budget_id)
key_1 = key_gen_1["key"]
key_2 = key_gen_2["key"]
await make_calls_until_budget_exceeded(
session=session,
key=key_1,
call_function=chat_completion,
model="fake-openai-endpoint",
)
await make_calls_until_budget_exceeded(
session=session,
key=key_2,
call_function=chat_completion,
model="fake-openai-endpoint",
)
await update_budget_max(session, budget_id, max_budget=0.001)
for key in (key_1, key_2):
response = await chat_completion(
session=session,
key=key,
model="fake-openai-endpoint",
)
assert response is not None
await delete_budget_by_id(session, budget_id)

View file

@ -1,5 +1,6 @@
import asyncio
import json
import math
import os
import sys
from unittest.mock import AsyncMock, MagicMock, patch
@ -2409,6 +2410,74 @@ async def test_virtual_key_budget_check_fallback_no_counter():
assert exc_info.value.current_cost == 15.0
@pytest.mark.parametrize(
"non_finite_max_budget",
[float("nan"), float("inf"), float("-inf")],
ids=["nan", "positive_infinity", "negative_infinity"],
)
@pytest.mark.asyncio
async def test_virtual_key_max_budget_check_non_finite_max_budget_does_not_bypass(
non_finite_max_budget,
):
"""Non-finite max_budget must not disable enforcement (GHSA-2rv4-xv66-fpjg)."""
from litellm.proxy.utils import ProxyLogging
valid_token = UserAPIKeyAuth(
token="test-hashed-token",
spend=0.0,
max_budget=non_finite_max_budget,
user_id="test-user",
)
proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
proxy_logging_obj.budget_alerts = AsyncMock()
async def mock_get_current_spend(counter_key, fallback_spend):
return 999.0
with patch(
"litellm.proxy.auth.auth_checks.math.isfinite", wraps=math.isfinite
) as mock_isfinite:
with patch(
"litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend
):
await _virtual_key_max_budget_check(
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
)
mock_isfinite.assert_called_once_with(non_finite_max_budget)
@pytest.mark.asyncio
async def test_virtual_key_max_budget_check_negative_infinity_blocks_without_isfinite_guard():
"""Removing math.isfinite would wrongly block keys with max_budget=-inf."""
from litellm.proxy.utils import ProxyLogging
valid_token = UserAPIKeyAuth(
token="test-hashed-token",
spend=0.0,
max_budget=float("-inf"),
user_id="test-user",
)
proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
proxy_logging_obj.budget_alerts = AsyncMock()
async def mock_get_current_spend(counter_key, fallback_spend):
return 999.0
with patch("litellm.proxy.auth.auth_checks.math.isfinite", return_value=True):
with patch(
"litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend
):
with pytest.raises(litellm.BudgetExceededError):
await _virtual_key_max_budget_check(
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
)
@pytest.mark.asyncio
async def test_team_budget_check_reads_from_spend_counter():
"""Team budget check should use get_current_spend when counter exists."""

View file

@ -30,6 +30,138 @@ from litellm.proxy.auth.auth_checks import common_checks
from litellm.proxy.utils import ProxyLogging
def _org_over_budget_object(
org_id: str,
*,
spend: float,
max_budget: float,
) -> LiteLLM_OrganizationTable:
return LiteLLM_OrganizationTable(
organization_id=org_id,
budget_id=f"budget-{org_id}",
spend=spend,
models=["gpt-4"],
created_by="test",
updated_by="test",
litellm_budget_table=LiteLLM_BudgetTable(max_budget=max_budget),
)
async def _assert_org_budget_blocks_request(
*,
valid_token: UserAPIKeyAuth,
org_object: LiteLLM_OrganizationTable,
team_object: Optional[LiteLLM_TeamTable] = None,
org_counter_spend: Optional[float] = None,
):
mock_request = MagicMock()
mock_request.url.path = "/v1/chat/completions"
mock_proxy_logging = MagicMock(spec=ProxyLogging)
mock_proxy_logging.budget_alerts = AsyncMock()
org_id = org_object.organization_id
counter_spend = (
org_counter_spend
if org_counter_spend is not None
else org_object.spend or 0.0
)
async def mock_get_current_spend(counter_key, fallback_spend):
if counter_key == f"spend:org:{org_id}":
return counter_spend
return fallback_spend
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
patch(
"litellm.proxy.auth.auth_checks.get_org_object",
new_callable=AsyncMock,
return_value=org_object,
),
patch(
"litellm.proxy.proxy_server.get_current_spend",
new_callable=AsyncMock,
side_effect=mock_get_current_spend,
),
):
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await common_checks(
request_body={"model": "gpt-4"},
team_object=team_object,
user_object=None,
end_user_object=None,
global_proxy_spend=None,
general_settings={},
route="/v1/chat/completions",
llm_router=None,
proxy_logging_obj=mock_proxy_logging,
valid_token=valid_token,
request=mock_request,
)
assert "Organization" in str(exc_info.value.message)
assert exc_info.value.current_cost == counter_spend
assert (
exc_info.value.max_budget
== org_object.litellm_budget_table.max_budget
)
@pytest.mark.asyncio
async def test_org_key_blocks_when_org_spend_exceeds_max_budget():
"""Org-scoped key with no team is blocked when org spend >= org max_budget."""
org_id = "org-key-only"
await _assert_org_budget_blocks_request(
valid_token=UserAPIKeyAuth(
token="sk-org-key-only",
org_id=org_id,
team_id=None,
spend=50.0,
max_budget=200.0,
),
org_object=_org_over_budget_object(
org_id, spend=120.0, max_budget=100.0
),
team_object=None,
org_counter_spend=120.0,
)
@pytest.mark.asyncio
async def test_all_keys_with_same_org_id_share_org_budget_enforcement():
"""Every key with the same org_id is blocked once org spend hits max_budget."""
org_id = "shared-org-budget"
org_object = _org_over_budget_object(org_id, spend=150.0, max_budget=100.0)
org_keys = [
UserAPIKeyAuth(
token="sk-org-key-alpha",
org_id=org_id,
team_id=None,
spend=10.0,
max_budget=500.0,
),
UserAPIKeyAuth(
token="sk-org-key-beta",
org_id=org_id,
team_id=None,
spend=80.0,
max_budget=500.0,
),
]
for valid_token in org_keys:
await _assert_org_budget_blocks_request(
valid_token=valid_token,
org_object=org_object,
team_object=None,
org_counter_spend=150.0,
)
@pytest.mark.asyncio
async def test_organization_budget_exceeded_blocks_request():
"""

View file

@ -0,0 +1,243 @@
"""
Regression coverage for per-organization-member budget enforcement.
Mirrors test_team_member_budget.py. These tests call _check_organization_member_budget
directly; they fail until that helper exists and is wired into common_checks.
"""
import pytest
from datetime import datetime, timezone
from unittest.mock import AsyncMock, MagicMock, patch
import litellm
from litellm.proxy._types import (
LiteLLM_BudgetTable,
LiteLLM_OrganizationMembershipTable,
LiteLLM_TeamTable,
UserAPIKeyAuth,
)
def _org_member_budget_check():
import litellm.proxy.auth.auth_checks as auth_checks
fn = getattr(auth_checks, "_check_organization_member_budget", None)
if fn is None:
pytest.fail("_check_organization_member_budget is not implemented")
return fn
@pytest.mark.asyncio
async def test_organization_member_budget_check_exceeds_budget():
team_object = LiteLLM_TeamTable(
team_id="test-team-1",
organization_id="test-org-1",
spend=0.0,
max_budget=None,
)
valid_token = UserAPIKeyAuth(
token="test-token",
user_id="test-user-1",
org_id="test-org-1",
team_id="test-team-1",
models=["gpt-3.5-turbo"],
)
now = datetime.now(timezone.utc)
org_membership = LiteLLM_OrganizationMembershipTable(
user_id="test-user-1",
organization_id="test-org-1",
spend=0.0000002,
created_at=now,
updated_at=now,
litellm_budget_table=LiteLLM_BudgetTable(max_budget=0.0000001),
)
mock_prisma_client = MagicMock()
mock_user_api_key_cache = MagicMock()
mock_proxy_logging_obj = MagicMock()
with (
patch(
"litellm.proxy.auth.auth_checks.get_organization_membership",
new_callable=AsyncMock,
return_value=org_membership,
),
patch(
"litellm.proxy.proxy_server.get_current_spend",
new_callable=AsyncMock,
return_value=0.0000002,
),
):
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await _org_member_budget_check()(
team_object=team_object,
valid_token=valid_token,
prisma_client=mock_prisma_client,
user_api_key_cache=mock_user_api_key_cache,
proxy_logging_obj=mock_proxy_logging_obj,
)
assert "Budget has been exceeded" in str(exc_info.value)
assert "test-user-1" in str(exc_info.value)
assert "test-org-1" in str(exc_info.value)
@pytest.mark.asyncio
async def test_organization_member_budget_check_within_budget():
team_object = LiteLLM_TeamTable(
team_id="test-team-1",
organization_id="test-org-1",
spend=0.0,
max_budget=None,
)
valid_token = UserAPIKeyAuth(
token="test-token",
user_id="test-user-1",
org_id="test-org-1",
team_id="test-team-1",
models=["gpt-3.5-turbo"],
)
now = datetime.now(timezone.utc)
org_membership = LiteLLM_OrganizationMembershipTable(
user_id="test-user-1",
organization_id="test-org-1",
spend=0.00000005,
created_at=now,
updated_at=now,
litellm_budget_table=LiteLLM_BudgetTable(max_budget=0.0000001),
)
mock_prisma_client = MagicMock()
mock_user_api_key_cache = MagicMock()
mock_proxy_logging_obj = MagicMock()
with (
patch(
"litellm.proxy.auth.auth_checks.get_organization_membership",
new_callable=AsyncMock,
return_value=org_membership,
),
patch(
"litellm.proxy.proxy_server.get_current_spend",
new_callable=AsyncMock,
return_value=0.00000005,
),
):
await _org_member_budget_check()(
team_object=team_object,
valid_token=valid_token,
prisma_client=mock_prisma_client,
user_api_key_cache=mock_user_api_key_cache,
proxy_logging_obj=mock_proxy_logging_obj,
)
@pytest.mark.asyncio
async def test_organization_member_budget_check_no_budget_set():
team_object = LiteLLM_TeamTable(
team_id="test-team-1",
organization_id="test-org-1",
spend=0.0,
max_budget=None,
)
valid_token = UserAPIKeyAuth(
token="test-token",
user_id="test-user-1",
org_id="test-org-1",
team_id="test-team-1",
models=["gpt-3.5-turbo"],
)
now = datetime.now(timezone.utc)
org_membership = LiteLLM_OrganizationMembershipTable(
user_id="test-user-1",
organization_id="test-org-1",
spend=0.0,
created_at=now,
updated_at=now,
litellm_budget_table=None,
)
mock_prisma_client = MagicMock()
mock_user_api_key_cache = MagicMock()
mock_proxy_logging_obj = MagicMock()
with patch(
"litellm.proxy.auth.auth_checks.get_organization_membership",
new_callable=AsyncMock,
return_value=org_membership,
):
await _org_member_budget_check()(
team_object=team_object,
valid_token=valid_token,
prisma_client=mock_prisma_client,
user_api_key_cache=mock_user_api_key_cache,
proxy_logging_obj=mock_proxy_logging_obj,
)
@pytest.mark.asyncio
async def test_organization_member_budget_check_no_membership():
team_object = LiteLLM_TeamTable(
team_id="test-team-1",
organization_id="test-org-1",
spend=0.0,
max_budget=None,
)
valid_token = UserAPIKeyAuth(
token="test-token",
user_id="test-user-1",
org_id="test-org-1",
team_id="test-team-1",
models=["gpt-3.5-turbo"],
)
mock_prisma_client = MagicMock()
mock_user_api_key_cache = MagicMock()
mock_proxy_logging_obj = MagicMock()
with patch(
"litellm.proxy.auth.auth_checks.get_organization_membership",
new_callable=AsyncMock,
return_value=None,
):
await _org_member_budget_check()(
team_object=team_object,
valid_token=valid_token,
prisma_client=mock_prisma_client,
user_api_key_cache=mock_user_api_key_cache,
proxy_logging_obj=mock_proxy_logging_obj,
)
@pytest.mark.asyncio
async def test_organization_member_budget_check_skipped_without_org_context():
valid_token = UserAPIKeyAuth(
token="test-token",
user_id="test-user-1",
org_id=None,
team_id=None,
models=["gpt-3.5-turbo"],
)
mock_prisma_client = MagicMock()
mock_user_api_key_cache = MagicMock()
mock_proxy_logging_obj = MagicMock()
with patch(
"litellm.proxy.auth.auth_checks.get_organization_membership",
new_callable=AsyncMock,
) as mock_get_org_membership:
await _org_member_budget_check()(
team_object=None,
valid_token=valid_token,
prisma_client=mock_prisma_client,
user_api_key_cache=mock_user_api_key_cache,
proxy_logging_obj=mock_proxy_logging_obj,
)
mock_get_org_membership.assert_not_called()

View file

@ -17,6 +17,36 @@ sys.path.insert(
) # Adds the parent directory to the system path
@pytest.fixture
def admin_client_and_mocks(monkeypatch):
mock_prisma = MagicMock()
mock_table = MagicMock()
mock_table.create = AsyncMock(side_effect=lambda *, data: data)
mock_table.update = AsyncMock(side_effect=lambda *, where, data: {**where, **data})
mock_table.delete = AsyncMock(side_effect=lambda *, where: where)
mock_table.find_many = AsyncMock(return_value=[])
mock_table.find_first = AsyncMock(return_value=None)
mock_prisma.db = types.SimpleNamespace(
litellm_budgettable=mock_table,
litellm_dailyspend=mock_table,
)
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
fake_user = UserAPIKeyAuth(
user_id="admin_user",
user_role=LitellmUserRoles.PROXY_ADMIN,
)
app.dependency_overrides[ps.user_api_key_auth] = lambda: fake_user
client = TestClient(app)
yield client, mock_prisma, mock_table
app.dependency_overrides.clear()
@pytest.fixture
def client_and_mocks(monkeypatch):
# Setup MagicMock Prisma
@ -44,9 +74,7 @@ def client_and_mocks(monkeypatch):
yield client, mock_prisma, mock_table
# teardown
app.dependency_overrides.clear()
monkeypatch.setattr(ps, "prisma_client", ps.prisma_client)
@pytest.mark.asyncio
@ -265,3 +293,137 @@ async def test_new_budget_invalid_model_max_budget(client_and_mocks, monkeypatch
assert resp.status_code in (400, 422), resp.text
detail = resp.json()["detail"]
assert "model_max_budget" in str(detail) or "dictionary" in str(detail).lower()
@pytest.mark.asyncio
async def test_info_budget_success(admin_client_and_mocks):
client, _, mock_table = admin_client_and_mocks
mock_table.find_many = AsyncMock(
return_value=[
{
"budget_id": "budget-info-1",
"max_budget": 10.0,
"budget_duration": "30d",
}
]
)
resp = client.post("/budget/info", json={"budgets": ["budget-info-1"]})
assert resp.status_code == 200, resp.text
body = resp.json()
assert len(body) == 1
assert body[0]["budget_id"] == "budget-info-1"
mock_table.find_many.assert_awaited_once()
@pytest.mark.asyncio
async def test_info_budget_empty_list_rejected(admin_client_and_mocks):
client, _, _ = admin_client_and_mocks
resp = client.post("/budget/info", json={"budgets": []})
assert resp.status_code == 400, resp.text
assert "Specify list of budget id" in str(resp.json()["detail"])
@pytest.mark.asyncio
async def test_info_budget_db_not_connected(admin_client_and_mocks, monkeypatch):
client, _, _ = admin_client_and_mocks
monkeypatch.setattr(ps, "prisma_client", None)
resp = client.post("/budget/info", json={"budgets": ["budget-info-1"]})
assert resp.status_code == 500
assert resp.json()["detail"]["error"] == "No db connected"
@pytest.mark.asyncio
async def test_list_budget_success(admin_client_and_mocks):
client, _, mock_table = admin_client_and_mocks
mock_table.find_many = AsyncMock(
return_value=[
{"budget_id": "budget-a", "max_budget": 1.0},
{"budget_id": "budget-b", "max_budget": 2.0},
]
)
resp = client.get("/budget/list")
assert resp.status_code == 200, resp.text
body = resp.json()
assert len(body) == 2
mock_table.find_many.assert_awaited_once()
@pytest.mark.asyncio
async def test_list_budget_rejects_non_admin(client_and_mocks):
client, _, _ = client_and_mocks
resp = client.get("/budget/list")
assert resp.status_code == 400, resp.text
assert CommonProxyErrors.not_allowed_access.value in str(resp.json()["detail"])
@pytest.mark.asyncio
async def test_budget_settings_success(admin_client_and_mocks):
client, _, mock_table = admin_client_and_mocks
mock_row = types.SimpleNamespace(
model_dump=lambda exclude_none=True: {
"budget_id": "budget-settings-1",
"max_budget": 25.0,
"soft_budget": 20.0,
"budget_duration": "7d",
}
)
mock_table.find_first = AsyncMock(return_value=mock_row)
resp = client.get("/budget/settings", params={"budget_id": "budget-settings-1"})
assert resp.status_code == 200, resp.text
body = resp.json()
field_names = {item["field_name"] for item in body}
assert "max_budget" in field_names
assert "soft_budget" in field_names
max_budget_field = next(item for item in body if item["field_name"] == "max_budget")
assert max_budget_field["field_value"] == 25.0
@pytest.mark.asyncio
async def test_budget_settings_rejects_non_admin(client_and_mocks):
client, _, _ = client_and_mocks
resp = client.get("/budget/settings", params={"budget_id": "budget-settings-1"})
assert resp.status_code == 400, resp.text
assert CommonProxyErrors.not_allowed_access.value in str(resp.json()["detail"])
@pytest.mark.asyncio
async def test_delete_budget_success(admin_client_and_mocks):
client, _, mock_table = admin_client_and_mocks
mock_table.delete = AsyncMock(return_value={"budget_id": "budget-delete-1"})
resp = client.post("/budget/delete", json={"id": "budget-delete-1"})
assert resp.status_code == 200, resp.text
assert resp.json()["budget_id"] == "budget-delete-1"
mock_table.delete.assert_awaited_once_with(
where={"budget_id": "budget-delete-1"}
)
@pytest.mark.asyncio
async def test_delete_budget_rejects_non_admin(client_and_mocks):
client, _, mock_table = client_and_mocks
fake_viewer = UserAPIKeyAuth(
user_id="viewer_user",
user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
)
app.dependency_overrides[ps.user_api_key_auth] = lambda: fake_viewer
try:
resp = client.post("/budget/delete", json={"id": "budget-delete-1"})
assert resp.status_code == 400, resp.text
assert CommonProxyErrors.not_allowed_access.value in str(resp.json()["detail"])
mock_table.delete.assert_not_called()
finally:
app.dependency_overrides.clear()