From 402919102bb6606936b3824b7c07a7e631871c15 Mon Sep 17 00:00:00 2001 From: mubashir1osmani Date: Mon, 8 Jun 2026 19:48:31 -0700 Subject: [PATCH 1/5] e2e budgeting --- tests/otel_tests/test_e2e_budgeting.py | 197 +++++++++++++++++- .../proxy/auth/test_auth_checks.py | 32 +++ .../test_budget_endpoints.py | 167 +++++++++++++++ 3 files changed, 391 insertions(+), 5 deletions(-) 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() From 5f954f1f4a62c849c9a1e5b87f082f9b89386918 Mon Sep 17 00:00:00 2001 From: mubashir1osmani Date: Tue, 9 Jun 2026 13:22:00 -0700 Subject: [PATCH 2/5] org member budget enforcement --- .../auth/test_organization_member_budget.py | 243 ++++++++++++++++++ 1 file changed, 243 insertions(+) create mode 100644 tests/test_litellm/proxy/auth/test_organization_member_budget.py diff --git a/tests/test_litellm/proxy/auth/test_organization_member_budget.py b/tests/test_litellm/proxy/auth/test_organization_member_budget.py new file mode 100644 index 00000000000..345a06d4f09 --- /dev/null +++ b/tests/test_litellm/proxy/auth/test_organization_member_budget.py @@ -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() From 5f3541c68bde09b76359e66aa9e03e94939a3560 Mon Sep 17 00:00:00 2001 From: mubashir1osmani Date: Tue, 9 Jun 2026 13:55:46 -0700 Subject: [PATCH 3/5] fix greptile review --- .../proxy/auth/test_auth_checks.py | 51 ++++++++++++++++--- 1 file changed, 44 insertions(+), 7 deletions(-) diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 6360d4a4e1b..cc05a1073ac 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -1,5 +1,6 @@ import asyncio import json +import math import os import sys from unittest.mock import AsyncMock, MagicMock, patch @@ -2411,8 +2412,8 @@ async def test_virtual_key_budget_check_fallback_no_counter(): @pytest.mark.parametrize( "non_finite_max_budget", - [float("nan"), float("inf")], - ids=["nan", "positive_infinity"], + [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( @@ -2434,11 +2435,47 @@ async def test_virtual_key_max_budget_check_non_finite_max_budget_does_not_bypas 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, - ) + 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 From 52a3f64846a781e33ba5893f1e5b65614adefe00 Mon Sep 17 00:00:00 2001 From: mubashir1osmani Date: Tue, 9 Jun 2026 14:00:02 -0700 Subject: [PATCH 4/5] fix' : --- tests/otel_tests/test_e2e_budgeting.py | 1 + .../test_budget_endpoints.py | 23 ++++++++----------- 2 files changed, 10 insertions(+), 14 deletions(-) diff --git a/tests/otel_tests/test_e2e_budgeting.py b/tests/otel_tests/test_e2e_budgeting.py index 0b40a081cb1..ab2874190c8 100644 --- a/tests/otel_tests/test_e2e_budgeting.py +++ b/tests/otel_tests/test_e2e_budgeting.py @@ -122,6 +122,7 @@ async def generate_key_with_budget_id(session, budget_id: str): 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() 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 d5edb08877c..cfce2c14e53 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py @@ -45,7 +45,6 @@ def admin_client_and_mocks(monkeypatch): yield client, mock_prisma, mock_table app.dependency_overrides.clear() - monkeypatch.setattr(ps, "prisma_client", ps.prisma_client) @pytest.fixture @@ -75,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 @@ -302,17 +299,15 @@ async def test_new_budget_invalid_model_max_budget(client_and_mocks, monkeypatch 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=[ + { + "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 @@ -429,6 +424,6 @@ async def test_delete_budget_rejects_non_admin(client_and_mocks): 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() + mock_table.delete.assert_not_called() finally: app.dependency_overrides.clear() From 05e8f8a81fa599d3abae2122f6581e568f9d04cd Mon Sep 17 00:00:00 2001 From: mubashir1osmani Date: Tue, 9 Jun 2026 14:42:07 -0700 Subject: [PATCH 5/5] add org budgets --- .../test_organization_budget_enforcement.py | 132 ++++++++++++++++++ 1 file changed, 132 insertions(+) diff --git a/tests/test_litellm/proxy/auth/test_organization_budget_enforcement.py b/tests/test_litellm/proxy/auth/test_organization_budget_enforcement.py index 45e24832274..f4095cee312 100644 --- a/tests/test_litellm/proxy/auth/test_organization_budget_enforcement.py +++ b/tests/test_litellm/proxy/auth/test_organization_budget_enforcement.py @@ -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(): """