mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
e2e budgeting
This commit is contained in:
parent
3448bf79f8
commit
402919102b
3 changed files with 391 additions and 5 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue