From d32c508ed60285c3a101011338941cf918ea4656 Mon Sep 17 00:00:00 2001 From: Ryan Crabbe Date: Fri, 8 May 2026 12:25:55 -0700 Subject: [PATCH] fix(sso): honor singular team_id_jwt_field in SSO callback MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The SSO callback flow (`generic_response_convertor` → `add_missing_team_member`) was reading only the plural `team_ids_jwt_field` claim. IdPs that populate the singular `team_id_jwt_field` instead (Okta/Auth0 group → primary team) had their users created via SSO with `teams: []` — never assigned to the team named in the JWT, regardless of how the user_metadata changed across logins. Adds `JWTHandler.get_all_jwt_team_ids()` returning the deduplicated union of both claim sources, and threads it through both branches that build `all_teams` in `generic_response_convertor`. JWT bearer auth path is intentionally unchanged — pending customer feedback on whether membership reconciliation there should be moved or stay additive. Reads the singular team_id_jwt_field claim directly via get_nested_value so the helper returns only what's actually in the token, avoiding silent fallback to litellm_jwtauth.team_id_default (which would have been an undocumented backwards-incompatible change for tenants with team_id_default configured). Tests pin behavior across both configurations of the default (team_id_jwt_field set + missing claim, and team_id_jwt_field unset). --- litellm/proxy/auth/handle_jwt.py | 30 ++++++++ litellm/proxy/management_endpoints/ui_sso.py | 4 +- .../proxy/auth/test_handle_jwt.py | 71 ++++++++++++++++++- 3 files changed, 102 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index d1fd5818f35..a270031c33c 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -224,6 +224,36 @@ class JWTHandler: return [] + def get_all_jwt_team_ids(self, token: dict) -> List[str]: + """ + Return team IDs from both the plural ``team_ids_jwt_field`` and the + singular ``team_id_jwt_field`` claim, as a deduplicated list preserving + plural-first order. + + Membership-reconciliation paths (SSO callback, JWT-bearer sync) need + to consider both claim shapes. Reading only the plural field — as + callers historically did — silently dropped users whose IdP populates + the singular field, which is what Okta and Auth0 default to when a + user has a single primary team. + + This intentionally does NOT consult ``team_id_default``: that fallback + is a property of how the JWT-bearer auth flow resolves a single + request-bound team, not of the token's claims. Callers that want the + default-team behavior should still go through ``get_team_id``. + """ + team_ids: List[str] = list(self.get_team_ids_from_jwt(token)) + if self.litellm_jwtauth.team_id_jwt_field is not None: + singular = get_nested_value( + data=token, + key_path=self.litellm_jwtauth.team_id_jwt_field, + default=None, + ) + if isinstance(singular, list): + singular = singular[0] if singular else None + if singular and singular not in team_ids: + team_ids.append(singular) + return team_ids + def get_end_user_id( self, token: dict, default_value: Optional[str] ) -> Optional[str]: diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 74ee7c7220d..ea629f24485 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -740,7 +740,7 @@ def generic_response_convertor( all_teams = [] if sso_jwt_handler is not None: - team_ids = sso_jwt_handler.get_team_ids_from_jwt(cast(dict, response)) + team_ids = sso_jwt_handler.get_all_jwt_team_ids(cast(dict, response)) all_teams.extend(team_ids) if team_mappings is not None and team_mappings.team_ids_jwt_field is not None: @@ -755,7 +755,7 @@ def generic_response_convertor( f"Loaded team_ids from DB team_mappings.team_ids_jwt_field='{team_mappings.team_ids_jwt_field}': {team_ids_from_db_mapping}" ) else: - team_ids = jwt_handler.get_team_ids_from_jwt(cast(dict, response)) + team_ids = jwt_handler.get_all_jwt_team_ids(cast(dict, response)) all_teams.extend(team_ids) # Determine user role based on role_mappings if available diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index b7dba9c1d16..c09c303deee 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -1,5 +1,5 @@ from typing import Optional -from unittest.mock import AsyncMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -494,6 +494,75 @@ async def test_sync_user_role_and_teams_no_cache_write_when_nothing_changes(): mock_cache.async_set_cache.assert_not_called() +def test_get_all_jwt_team_ids_unions_singular_and_plural(): + """get_all_jwt_team_ids must include the singular team_id_jwt_field claim + in addition to the plural team_ids_jwt_field, deduplicated.""" + jwt_handler = JWTHandler() + jwt_handler.update_environment( + prisma_client=None, + user_api_key_cache=MagicMock(), + litellm_jwtauth=LiteLLM_JWTAuth( + team_id_jwt_field="team_id", + team_ids_jwt_field="teams", + ), + ) + + # singular only — Okta/Auth0 default shape + assert jwt_handler.get_all_jwt_team_ids({"team_id": "team-low"}) == ["team-low"] + + # plural only — pre-fix shape + assert jwt_handler.get_all_jwt_team_ids({"teams": ["a", "b"]}) == ["a", "b"] + + # both populated, no overlap + assert jwt_handler.get_all_jwt_team_ids( + {"team_id": "primary", "teams": ["a", "b"]} + ) == ["a", "b", "primary"] + + # both populated with overlap — singular dedup'd + assert jwt_handler.get_all_jwt_team_ids({"team_id": "a", "teams": ["a", "b"]}) == [ + "a", + "b", + ] + + # neither populated + assert jwt_handler.get_all_jwt_team_ids({}) == [] + + +def test_get_all_jwt_team_ids_does_not_use_team_id_default(): + """team_id_default is a JWT-bearer-flow auth-builder fallback, not a token + claim. It must NOT leak into get_all_jwt_team_ids — otherwise SSO logins + would silently start adding users to the default team for any tenant that + has team_id_default configured.""" + jwt_handler = JWTHandler() + jwt_handler.update_environment( + prisma_client=None, + user_api_key_cache=MagicMock(), + litellm_jwtauth=LiteLLM_JWTAuth( + team_id_jwt_field="team_id", + team_ids_jwt_field="teams", + team_id_default="default-team", + ), + ) + + # team_id claim missing — must not fall back to default-team + assert jwt_handler.get_all_jwt_team_ids({"teams": []}) == [] + assert jwt_handler.get_all_jwt_team_ids({}) == [] + + # only the plural is populated — default still must not be added + assert jwt_handler.get_all_jwt_team_ids({"teams": ["a"]}) == ["a"] + + # team_id_jwt_field unset entirely + only default configured: still no default + jwt_handler.update_environment( + prisma_client=None, + user_api_key_cache=MagicMock(), + litellm_jwtauth=LiteLLM_JWTAuth( + team_ids_jwt_field="teams", + team_id_default="default-team", + ), + ) + assert jwt_handler.get_all_jwt_team_ids({"teams": []}) == [] + + @pytest.mark.asyncio async def test_map_jwt_role_to_litellm_role(): """Test JWT role mapping to LiteLLM roles with various patterns"""