From cfd8c186161068bef2ed5faae002dbfb3b2ab63e Mon Sep 17 00:00:00 2001 From: jesus Date: Thu, 17 Sep 2026 21:24:36 +0000 Subject: [PATCH] fix(auth): inherit org budget, tpm and rpm limits for JWT and team-linked keys Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/user_api_key_auth.py | 31 ++++++++++++----- .../proxy/auth/test_user_api_key_auth.py | 34 +++++++++++++++---- 2 files changed, 50 insertions(+), 15 deletions(-) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 110c524ecdf..df5257908ce 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -2409,11 +2409,17 @@ async def _inherit_org_identity( ) -> None: if user_api_key_auth_obj.org_id is None and team_object is not None and team_object.organization_id is not None: user_api_key_auth_obj.org_id = team_object.organization_id - if ( - user_api_key_auth_obj.org_id is None - or user_api_key_auth_obj.organization_alias is not None - or prisma_client is None - ): + already_populated: Final = any( + value is not None + for value in ( + user_api_key_auth_obj.organization_alias, + user_api_key_auth_obj.organization_max_budget, + user_api_key_auth_obj.organization_tpm_limit, + user_api_key_auth_obj.organization_rpm_limit, + user_api_key_auth_obj.organization_metadata, + ) + ) + if user_api_key_auth_obj.org_id is None or already_populated or prisma_client is None: return try: org_object: Final = await get_org_object( @@ -2422,12 +2428,21 @@ async def _inherit_org_identity( user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, + include_budget_table=True, ) except Exception: - verbose_proxy_logger.debug("org alias lookup failed for org_id=%s", user_api_key_auth_obj.org_id, exc_info=True) + verbose_proxy_logger.debug("org lookup failed for org_id=%s", user_api_key_auth_obj.org_id, exc_info=True) return - if org_object is not None: - user_api_key_auth_obj.organization_alias = org_object.organization_alias + if org_object is None: + return + user_api_key_auth_obj.organization_alias = org_object.organization_alias + user_api_key_auth_obj.organization_metadata = org_object.metadata + budget: Final = org_object.litellm_budget_table + if budget is None: + return + user_api_key_auth_obj.organization_max_budget = budget.max_budget + user_api_key_auth_obj.organization_tpm_limit = budget.tpm_limit + user_api_key_auth_obj.organization_rpm_limit = budget.rpm_limit @tracer.wrap() diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 03efbfa7185..fd9b6f09678 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -5296,22 +5296,26 @@ async def test_centralized_common_checks_backfills_org_id_from_team(key_org_id, @pytest.mark.asyncio @pytest.mark.parametrize( - "key_org_id,team_id,team_org_id,existing_alias,lookup_mode,expected_org_id,expected_alias", + "key_org_id,team_id,team_org_id,existing_alias,existing_rpm,lookup_mode,expected_org_id,expected_alias,expected_limits", [ - (None, "t1", "org-from-team", None, "success", "org-from-team", "acme-org"), - ("org-jwt", None, None, None, "success", "org-jwt", "acme-org"), - ("org-pinned", None, None, "preset", "success", "org-pinned", "preset"), - ("org-missing", None, None, None, "missing", "org-missing", None), + (None, "t1", "org-from-team", None, None, "success", "org-from-team", "acme-org", (12.5, 700, 7)), + ("org-jwt", None, None, None, None, "success", "org-jwt", "acme-org", (12.5, 700, 7)), + ("org-pinned", None, None, "preset", None, "success", "org-pinned", "preset", (None, None, None)), + ("org-view", None, None, None, 3, "success", "org-view", None, (None, None, 3)), + ("org-missing", None, None, None, None, "missing", "org-missing", None, (None, None, None)), + ("org-nobudget", None, None, None, None, "no_budget", "org-nobudget", "acme-org", (None, None, None)), ], ) -async def test_centralized_common_checks_inherits_org_alias( +async def test_centralized_common_checks_inherits_org_identity( key_org_id, team_id, team_org_id, existing_alias, + existing_rpm, lookup_mode, expected_org_id, expected_alias, + expected_limits, ): import litellm.proxy.proxy_server as _proxy_server_mod from fastapi import Request @@ -5325,6 +5329,7 @@ async def test_centralized_common_checks_inherits_org_alias( team_id=team_id, org_id=key_org_id, organization_alias=existing_alias, + organization_rpm_limit=existing_rpm, ) request = Request(scope={"type": "http"}) request._url = URL(url="/chat/completions") @@ -5336,9 +5341,15 @@ async def test_centralized_common_checks_inherits_org_alias( organization_id=expected_org_id, organization_alias="acme-org", budget_id="budget-id", + metadata={"model_rpm_limit": {"gpt-4o": 2}}, models=[], created_by="test", updated_by="test", + litellm_budget_table=( + None + if lookup_mode == "no_budget" + else LiteLLM_BudgetTable(budget_id="budget-id", max_budget=12.5, tpm_limit=700, rpm_limit=7) + ), ) attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=None) @@ -5380,16 +5391,25 @@ async def test_centralized_common_checks_inherits_org_alias( mock_checks.assert_awaited_once() assert token.org_id == expected_org_id assert token.organization_alias == expected_alias + assert ( + token.organization_max_budget, + token.organization_tpm_limit, + token.organization_rpm_limit, + ) == expected_limits assert identity_seen_by_common_checks == [(expected_org_id, expected_alias)] if team_id is None: mock_get_team_object.assert_not_awaited() else: mock_get_team_object.assert_awaited_once() - if existing_alias is not None: + if existing_alias is not None or existing_rpm is not None: mock_get_org_object.assert_not_awaited() + assert token.organization_metadata is None else: mock_get_org_object.assert_awaited_once() assert mock_get_org_object.await_args.kwargs["org_id"] == expected_org_id + assert mock_get_org_object.await_args.kwargs["include_budget_table"] is True + if lookup_mode != "missing": + assert token.organization_metadata == {"model_rpm_limit": {"gpt-4o": 2}} finally: for k, v in originals.items(): setattr(_proxy_server_mod, k, v)