e2e budgeting

This commit is contained in:
mubashir1osmani 2026-06-08 19:48:31 -07:00
parent 3448bf79f8
commit 402919102b
No known key found for this signature in database
GPG key ID: AB055FF67D0B4D9A
3 changed files with 391 additions and 5 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,98 @@ 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:
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 +528,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

@ -2409,6 +2409,38 @@ 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")],
ids=["nan", "positive_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.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,
)
@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

@ -17,6 +17,37 @@ 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()
monkeypatch.setattr(ps, "prisma_client", ps.prisma_client)
@pytest.fixture
def client_and_mocks(monkeypatch):
# Setup MagicMock Prisma
@ -265,3 +296,139 @@ 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_row = types.SimpleNamespace(
budget_id="budget-info-1",
max_budget=10.0,
budget_duration="30d",
dict=lambda: {
"budget_id": "budget-info-1",
"max_budget": 10.0,
"budget_duration": "30d",
},
)
mock_table.find_many = AsyncMock(return_value=[mock_row])
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_awaited()
finally:
app.dependency_overrides.clear()