mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Merge 05e8f8a81f into 2fe9feda71
This commit is contained in:
commit
701ffeb501
5 changed files with 801 additions and 7 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,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)
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
243
tests/test_litellm/proxy/auth/test_organization_member_budget.py
Normal file
243
tests/test_litellm/proxy/auth/test_organization_member_budget.py
Normal 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()
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue