fixed _organization_max_budget_check so tests works

This commit is contained in:
shivam 2026-01-10 12:45:41 -08:00
parent 280085edaa
commit d69fb6b726

View file

@ -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(