mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fixed _organization_max_budget_check so tests works
This commit is contained in:
parent
280085edaa
commit
d69fb6b726
1 changed files with 93 additions and 39 deletions
|
|
@ -147,6 +147,7 @@ async def common_checks(
|
|||
# 3.1. If organization is in budget
|
||||
await _organization_max_budget_check(
|
||||
valid_token=valid_token,
|
||||
team_object=team_object,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
|
|
@ -1793,11 +1794,20 @@ async def get_org_object(
|
|||
user_api_key_cache: DualCache,
|
||||
parent_otel_span: Optional[Span] = None,
|
||||
proxy_logging_obj: Optional[ProxyLogging] = None,
|
||||
include_budget_table: bool = False,
|
||||
) -> Optional[LiteLLM_OrganizationTable]:
|
||||
"""
|
||||
- Check if org id in proxy Org Table
|
||||
- if valid, return LiteLLM_OrganizationTable object
|
||||
- if not, then raise an error
|
||||
|
||||
Args:
|
||||
org_id: Organization ID to look up
|
||||
prisma_client: Database client
|
||||
user_api_key_cache: Cache for storing results
|
||||
parent_otel_span: Optional OpenTelemetry span
|
||||
proxy_logging_obj: Optional proxy logging object
|
||||
include_budget_table: If True, includes litellm_budget_table in the query
|
||||
"""
|
||||
if prisma_client is None:
|
||||
raise Exception(
|
||||
|
|
@ -1806,8 +1816,13 @@ async def get_org_object(
|
|||
if not isinstance(org_id, str):
|
||||
return None
|
||||
|
||||
# Use different cache key if budget table is included
|
||||
cache_key = "org_id:{}".format(org_id)
|
||||
if include_budget_table:
|
||||
cache_key = "org_id:{}:with_budget".format(org_id)
|
||||
|
||||
# check if in cache
|
||||
cached_org_obj = user_api_key_cache.async_get_cache(key="org_id:{}".format(org_id))
|
||||
cached_org_obj = user_api_key_cache.async_get_cache(key=cache_key)
|
||||
if cached_org_obj is not None:
|
||||
if isinstance(cached_org_obj, dict):
|
||||
return LiteLLM_OrganizationTable(**cached_org_obj)
|
||||
|
|
@ -1815,13 +1830,24 @@ async def get_org_object(
|
|||
return cached_org_obj
|
||||
# else, check db
|
||||
try:
|
||||
query_kwargs = {"where": {"organization_id": org_id}}
|
||||
if include_budget_table:
|
||||
query_kwargs["include"] = {"litellm_budget_table": True}
|
||||
|
||||
response = await prisma_client.db.litellm_organizationtable.find_unique(
|
||||
where={"organization_id": org_id}
|
||||
**query_kwargs
|
||||
)
|
||||
|
||||
if response is None:
|
||||
raise Exception
|
||||
|
||||
# Cache the result
|
||||
await user_api_key_cache.async_set_cache(
|
||||
key=cache_key,
|
||||
value=response.model_dump() if hasattr(response, "model_dump") else response,
|
||||
ttl=DEFAULT_IN_MEMORY_TTL,
|
||||
)
|
||||
|
||||
return response
|
||||
except Exception:
|
||||
raise Exception(
|
||||
|
|
@ -2310,61 +2336,89 @@ async def _team_max_budget_check(
|
|||
|
||||
async def _organization_max_budget_check(
|
||||
valid_token: Optional[UserAPIKeyAuth],
|
||||
team_object: Optional[LiteLLM_TeamTable],
|
||||
prisma_client: Optional[PrismaClient],
|
||||
user_api_key_cache: DualCache,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
):
|
||||
"""
|
||||
Check if the organization is over its max budget.
|
||||
|
||||
This function checks the organization budget using:
|
||||
1. First, tries to use valid_token.org_id (if key has organization_id set)
|
||||
2. Falls back to team_object.organization_id (if key doesn't have org_id but team does)
|
||||
|
||||
This ensures organization budget checks work even when keys don't have organization_id
|
||||
set directly, as long as their team belongs to an organization.
|
||||
|
||||
Raises:
|
||||
BudgetExceededError if the organization is over its max budget.
|
||||
Triggers a budget alert if the organization is over its max budget.
|
||||
"""
|
||||
# Only check if token has organization info and organization_max_budget is set
|
||||
if (
|
||||
valid_token is None
|
||||
or valid_token.org_id is None
|
||||
or valid_token.organization_max_budget is None
|
||||
or valid_token.organization_max_budget <= 0
|
||||
):
|
||||
if valid_token is None or prisma_client is None:
|
||||
return
|
||||
|
||||
# Get organization object to check current spend
|
||||
if prisma_client is not None:
|
||||
# Determine organization_id: first try from token, then fallback to team
|
||||
org_id: Optional[str] = None
|
||||
if valid_token.org_id is not None:
|
||||
org_id = valid_token.org_id
|
||||
elif team_object is not None and team_object.organization_id is not None:
|
||||
org_id = team_object.organization_id
|
||||
|
||||
# If no organization_id found, skip the check
|
||||
if org_id is None:
|
||||
return
|
||||
|
||||
# Get organization object with budget table - use get_org_object so it can be mocked in tests
|
||||
try:
|
||||
org_table = await get_org_object(
|
||||
org_id=valid_token.org_id,
|
||||
org_id=org_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
include_budget_table=True,
|
||||
)
|
||||
except Exception:
|
||||
# If organization lookup fails, skip the check
|
||||
return
|
||||
|
||||
if org_table is None:
|
||||
return
|
||||
|
||||
# Get max_budget from organization's budget table
|
||||
org_max_budget: Optional[float] = None
|
||||
if org_table.litellm_budget_table is not None:
|
||||
org_max_budget = org_table.litellm_budget_table.max_budget
|
||||
|
||||
# Only check if organization has a valid max_budget set
|
||||
if org_max_budget is None or org_max_budget <= 0:
|
||||
return
|
||||
|
||||
# Check if organization spend exceeds max budget
|
||||
if org_table.spend >= org_max_budget:
|
||||
# Trigger budget alert
|
||||
call_info = CallInfo(
|
||||
token=valid_token.token,
|
||||
spend=org_table.spend,
|
||||
max_budget=org_max_budget,
|
||||
user_id=valid_token.user_id,
|
||||
team_id=valid_token.team_id,
|
||||
team_alias=valid_token.team_alias,
|
||||
organization_id=org_id,
|
||||
event_group=Litellm_EntityType.ORGANIZATION,
|
||||
)
|
||||
asyncio.create_task(
|
||||
proxy_logging_obj.budget_alerts(
|
||||
type="organization_budget",
|
||||
user_info=call_info,
|
||||
)
|
||||
)
|
||||
|
||||
if (
|
||||
org_table is not None
|
||||
and org_table.spend >= valid_token.organization_max_budget
|
||||
):
|
||||
# Trigger budget alert
|
||||
call_info = CallInfo(
|
||||
token=valid_token.token,
|
||||
spend=org_table.spend,
|
||||
max_budget=valid_token.organization_max_budget,
|
||||
user_id=valid_token.user_id,
|
||||
team_id=valid_token.team_id,
|
||||
team_alias=valid_token.team_alias,
|
||||
organization_id=valid_token.org_id,
|
||||
event_group=Litellm_EntityType.ORGANIZATION,
|
||||
)
|
||||
asyncio.create_task(
|
||||
proxy_logging_obj.budget_alerts(
|
||||
type="organization_budget",
|
||||
user_info=call_info,
|
||||
)
|
||||
)
|
||||
|
||||
raise litellm.BudgetExceededError(
|
||||
current_cost=org_table.spend,
|
||||
max_budget=valid_token.organization_max_budget,
|
||||
message=f"Budget has been exceeded! Organization={valid_token.org_id} Current cost: {org_table.spend}, Max budget: {valid_token.organization_max_budget}",
|
||||
)
|
||||
raise litellm.BudgetExceededError(
|
||||
current_cost=org_table.spend,
|
||||
max_budget=org_max_budget,
|
||||
message=f"Budget has been exceeded! Organization={org_id} Current cost: {org_table.spend}, Max budget: {org_max_budget}",
|
||||
)
|
||||
|
||||
|
||||
async def _tag_max_budget_check(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue