From d69fb6b72694bc252ffb1d3a179f9a2990251d72 Mon Sep 17 00:00:00 2001 From: shivam Date: Sat, 10 Jan 2026 12:45:41 -0800 Subject: [PATCH] fixed _organization_max_budget_check so tests works --- litellm/proxy/auth/auth_checks.py | 132 +++++++++++++++++++++--------- 1 file changed, 93 insertions(+), 39 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index de4973ecc69..db5a21b0a85 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -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(