diff --git a/tests/otel_tests/test_e2e_budgeting.py b/tests/otel_tests/test_e2e_budgeting.py index f61befac4fb..0b40a081cb1 100644 --- a/tests/otel_tests/test_e2e_budgeting.py +++ b/tests/otel_tests/test_e2e_budgeting.py @@ -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) diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 42c76c4671d..6360d4a4e1b 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -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.""" diff --git a/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py index d924d5ecdfe..d5edb08877c 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py @@ -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()