feat(jwt): fall back to DB team memberships when JWT has no team claims

This commit is contained in:
mateo-berri 2026-06-25 15:00:37 -07:00
parent 92d0788da2
commit 931fa9af5e
3 changed files with 420 additions and 32 deletions

View file

@ -4312,6 +4312,17 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
"authorization."
),
)
fallback_to_db_teams: bool = Field(
default=False,
description=(
"When True, users whose JWT contains no team claims are authenticated "
"using their database team memberships instead of receiving HTTP 403. "
"Usage is attributed to the user's first resolvable DB team, or to the "
"team specified via the x-litellm-team-id request header (validated "
"against DB membership). Requires user_id_upsert=True so that user "
"records exist before the fallback runs."
),
)
issuers: Optional[List[JWTIssuerConfig]] = Field(
default=None,
description="Optional issuer-bound JWT validation rules. When a token's `iss` matches a configured issuer, validation uses that issuer's JWKS, audience, and claim mappings. Tokens with an unlisted `iss` fall back to the global JWT_AUDIENCE/JWT_ISSUER validation path — this is additive routing, not an allow-list.",

View file

@ -1438,7 +1438,10 @@ class JWTAuthManager:
denied_auth_enforced_pass_through_route = False
if not team_ids:
if jwt_handler.litellm_jwtauth.enforce_team_based_model_access:
if (
jwt_handler.litellm_jwtauth.enforce_team_based_model_access
and not jwt_handler.litellm_jwtauth.fallback_to_db_teams
):
raise HTTPException(
status_code=403,
detail="No teams found in token. `enforce_team_based_model_access` is set to True. Token must belong to a team.",
@ -1700,6 +1703,7 @@ class JWTAuthManager:
def get_team_id_from_header(
request_headers: Optional[dict],
allowed_team_ids: Set[str],
fallback_to_db_teams: bool = False,
) -> Optional[str]:
"""
Extract team_id from x-litellm-team-id header if present.
@ -1708,6 +1712,10 @@ class JWTAuthManager:
Args:
request_headers: Dictionary of request headers
allowed_team_ids: Set of team IDs the user is allowed to access (from JWT)
fallback_to_db_teams: When True and the JWT carries no team claims
(allowed_team_ids is empty), the header value is returned
provisionally and validated against DB memberships later in
auth_builder instead of being rejected here.
Returns:
The team_id from header if valid, None otherwise
@ -1725,8 +1733,8 @@ class JWTAuthManager:
if not header_team_id:
return None
# Validate that the team_id is in the allowed teams
if header_team_id not in allowed_team_ids:
defer_to_db_membership = fallback_to_db_teams and not allowed_team_ids
if not defer_to_db_membership and header_team_id not in allowed_team_ids:
raise HTTPException(
status_code=403,
detail=f"Team '{header_team_id}' from x-litellm-team-id header is not in your JWT's allowed teams. Allowed teams: {list(allowed_team_ids)}",
@ -1953,6 +1961,74 @@ class JWTAuthManager:
)
return None, None, None
@staticmethod
async def _resolve_db_team_fallback(
user_object: Optional[LiteLLM_UserTable],
enforce_team_based_model_access: bool,
team_id_upsert: bool,
prisma_client: Optional[PrismaClient],
user_api_key_cache: UserApiKeyCache,
parent_otel_span: Optional[Span],
proxy_logging_obj: ProxyLogging,
) -> Tuple[Optional[str], Optional[LiteLLM_TeamTable]]:
"""
Resolve a team for a user whose JWT carries no team claims by selecting
the first of their DB team memberships that loads successfully.
Raises HTTP 403 when the user has no usable DB team membership and
`enforce_team_based_model_access` is set; otherwise returns (None, None).
"""
user_team_ids = user_object.teams if user_object else []
for candidate_team_id in user_team_ids:
try:
team_object = await get_team_object(
team_id=candidate_team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
team_id_upsert=team_id_upsert,
)
except Exception:
continue
if team_object:
verbose_proxy_logger.debug(
"JWT DB team fallback: resolved team_id=%s from user DB membership",
candidate_team_id,
)
return candidate_team_id, team_object
if enforce_team_based_model_access:
raise HTTPException(
status_code=403,
detail=(
"User is not a member of any team. Add the user to a team via "
"the LiteLLM UI or API."
),
)
return None, None
@staticmethod
def _validate_header_team_in_db_membership(
team_id: str,
user_object: Optional[LiteLLM_UserTable],
) -> None:
"""
A provisional team_id from the x-litellm-team-id header (accepted without
JWT-team validation when the JWT carries no team claims) must exist in the
user's DB team memberships before it becomes request context.
"""
user_team_ids = user_object.teams if user_object else []
if team_id in user_team_ids:
return
raise HTTPException(
status_code=403,
detail=(
f"Team '{team_id}' (from x-litellm-team-id header) is not in your "
f"team memberships. Your teams: {user_team_ids}"
),
)
@staticmethod
async def auth_builder(
api_key: str,
@ -2065,6 +2141,7 @@ class JWTAuthManager:
header_team_id = JWTAuthManager.get_team_id_from_header(
request_headers=request_headers,
allowed_team_ids=all_team_ids,
fallback_to_db_teams=jwt_handler.litellm_jwtauth.fallback_to_db_teams,
)
if header_team_id:
team_id = header_team_id
@ -2170,8 +2247,21 @@ class JWTAuthManager:
user_api_key_cache=user_api_key_cache,
)
# If JWT did not resolve team_id, attempt single-team DB fallback.
if team_id is None:
# If JWT did not resolve team_id, attempt a team fallback.
db_team_fallback = (
jwt_handler.litellm_jwtauth.fallback_to_db_teams and not all_team_ids
)
if team_id is None and db_team_fallback:
team_id, team_object = await JWTAuthManager._resolve_db_team_fallback(
user_object=user_object,
enforce_team_based_model_access=jwt_handler.litellm_jwtauth.enforce_team_based_model_access,
team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
elif team_id is None:
(
team_id,
team_object,
@ -2185,6 +2275,11 @@ class JWTAuthManager:
proxy_logging_obj=proxy_logging_obj,
team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert,
)
elif db_team_fallback:
JWTAuthManager._validate_header_team_in_db_membership(
team_id=team_id,
user_object=user_object,
)
## MAP USER TO TEAMS
await JWTAuthManager.map_user_to_teams(

View file

@ -1241,24 +1241,24 @@ async def test_auth_builder_returns_team_membership_object():
)
# Verify that team_membership_object is returned
assert (
result["team_membership"] is not None
), "team_membership should be present"
assert (
result["team_membership"] == mock_team_membership
), "team_membership should match the mock object"
assert (
result["team_membership"].user_id == _user_id
), "team_membership user_id should match"
assert (
result["team_membership"].team_id == _team_id
), "team_membership team_id should match"
assert (
result["team_membership"].budget_id == "budget_123"
), "team_membership budget_id should match"
assert (
result["team_membership"].spend == 10.5
), "team_membership spend should match"
assert result["team_membership"] is not None, (
"team_membership should be present"
)
assert result["team_membership"] == mock_team_membership, (
"team_membership should match the mock object"
)
assert result["team_membership"].user_id == _user_id, (
"team_membership user_id should match"
)
assert result["team_membership"].team_id == _team_id, (
"team_membership team_id should match"
)
assert result["team_membership"].budget_id == "budget_123", (
"team_membership budget_id should match"
)
assert result["team_membership"].spend == 10.5, (
"team_membership spend should match"
)
@pytest.mark.asyncio
@ -2717,9 +2717,9 @@ async def test_find_and_validate_specific_team_id_hints_bracket_notation():
error_msg = str(exc_info.value)
# Should mention the bad field name and suggest the fix
assert "roles.0" in error_msg, f"Expected field name in: {error_msg}"
assert (
"roles" in error_msg and "list" in error_msg
), f"Expected hint about using 'roles' instead: {error_msg}"
assert "roles" in error_msg and "list" in error_msg, (
f"Expected hint about using 'roles' instead: {error_msg}"
)
@pytest.mark.asyncio
@ -2747,9 +2747,9 @@ async def test_find_and_validate_specific_team_id_hints_bracket_index_notation()
error_msg = str(exc_info.value)
assert "roles[0]" in error_msg, f"Expected field name in: {error_msg}"
assert (
"roles" in error_msg and "list" in error_msg
), f"Expected hint about using 'roles' instead: {error_msg}"
assert "roles" in error_msg and "list" in error_msg, (
f"Expected hint about using 'roles' instead: {error_msg}"
)
@pytest.mark.asyncio
@ -3164,9 +3164,9 @@ def test_build_decode_kwargs_warns_once_when_unscoped(
if "JWT auth is enabled" in r.getMessage()
and "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage()
]
assert (
len(matching) == 1
), f"Expected exactly one warning across 3 calls, got {len(matching)}"
assert len(matching) == 1, (
f"Expected exactly one warning across 3 calls, got {len(matching)}"
)
def test_build_decode_kwargs_no_warning_when_scoped(
@ -4339,3 +4339,285 @@ def test_build_decode_kwargs_warns_for_unscoped_global_fallback_in_mixed_deploym
if "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage()
]
assert len(matching) == 1
# ---------------------------------------------------------------------------
# fallback_to_db_teams: resolve team from DB memberships when JWT has no team
# claims (config flag on LiteLLM_JWTAuth)
# ---------------------------------------------------------------------------
def test_get_team_id_from_header_defers_to_db_membership_only_without_jwt_claims():
"""With fallback_to_db_teams=True, an x-litellm-team-id header is accepted
provisionally only when the JWT carries no team claims (allowed set empty).
When the JWT does carry team claims, the header must still be validated
against them, and the flag-off behavior must keep rejecting unknown teams."""
deferred = JWTAuthManager.get_team_id_from_header(
request_headers={"x-litellm-team-id": "team-from-db"},
allowed_team_ids=set(),
fallback_to_db_teams=True,
)
assert deferred == "team-from-db"
with pytest.raises(HTTPException) as exc_info:
JWTAuthManager.get_team_id_from_header(
request_headers={"x-litellm-team-id": "team-x"},
allowed_team_ids={"team-1", "team-2"},
fallback_to_db_teams=True,
)
assert exc_info.value.status_code == 403
with pytest.raises(HTTPException):
JWTAuthManager.get_team_id_from_header(
request_headers={"x-litellm-team-id": "team-from-db"},
allowed_team_ids=set(),
fallback_to_db_teams=False,
)
@pytest.mark.asyncio
async def test_find_team_with_model_access_defers_no_team_403_under_db_fallback():
"""find_team_with_model_access raises the early "no teams in token" 403 when
enforcement is on, but defers (returns no team) so auth_builder's DB fallback
can run when fallback_to_db_teams is enabled."""
jwt_handler = JWTHandler()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
enforce_team_based_model_access=True,
fallback_to_db_teams=False,
)
with pytest.raises(HTTPException) as exc_info:
await JWTAuthManager.find_team_with_model_access(
team_ids=set(),
requested_model="gpt-4",
route="/chat/completions",
jwt_handler=jwt_handler,
prisma_client=None,
user_api_key_cache=MagicMock(),
parent_otel_span=None,
proxy_logging_obj=MagicMock(),
)
assert exc_info.value.status_code == 403
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
enforce_team_based_model_access=True,
fallback_to_db_teams=True,
)
team_id, team_object = await JWTAuthManager.find_team_with_model_access(
team_ids=set(),
requested_model="gpt-4",
route="/chat/completions",
jwt_handler=jwt_handler,
prisma_client=None,
user_api_key_cache=MagicMock(),
parent_otel_span=None,
proxy_logging_obj=MagicMock(),
)
assert team_id is None
assert team_object is None
@pytest.mark.asyncio
async def test_resolve_db_team_fallback_skips_unresolvable_membership():
"""An orphaned membership (team row missing/erroring) is skipped and the next
resolvable DB team is selected instead of aborting the fallback."""
user_object = LiteLLM_UserTable(
user_id="u_skip",
user_role=LitellmUserRoles.INTERNAL_USER,
teams=["ghost_team", "real_team"],
)
resolved = LiteLLM_TeamTable(team_id="real_team")
async def fake_get_team(team_id, **kwargs):
if team_id == "ghost_team":
raise HTTPException(status_code=404, detail="missing")
return resolved
with patch(
"litellm.proxy.auth.handle_jwt.get_team_object",
new_callable=AsyncMock,
side_effect=fake_get_team,
):
team_id, team_object = await JWTAuthManager._resolve_db_team_fallback(
user_object=user_object,
enforce_team_based_model_access=True,
team_id_upsert=False,
prisma_client=None,
user_api_key_cache=MagicMock(),
parent_otel_span=None,
proxy_logging_obj=MagicMock(),
)
assert team_id == "real_team"
assert team_object is resolved
@pytest.mark.parametrize(
(
"fallback_to_db_teams",
"user_teams",
"header_team_id",
"expected_team_id",
"expect_403",
),
[
pytest.param(
True, ["team_solo"], None, "team_solo", False, id="flag_on_single_db_team"
),
pytest.param(
True,
["team_a", "team_b"],
None,
"team_a",
False,
id="flag_on_multi_db_team_picks_first",
),
pytest.param(
True,
["team_a", "team_b"],
"team_b",
"team_b",
False,
id="flag_on_header_team_in_membership",
),
pytest.param(
True,
["team_a", "team_b"],
"team_x",
None,
True,
id="flag_on_header_team_not_in_membership_403",
),
pytest.param(True, [], None, None, True, id="flag_on_no_db_team_enforced_403"),
pytest.param(
False,
["team_a", "team_b"],
None,
None,
False,
id="flag_off_multi_db_team_no_fallback",
),
pytest.param(
False,
["team_solo"],
None,
"team_solo",
False,
id="flag_off_single_db_team_upstream_fallback",
),
],
)
@pytest.mark.asyncio
async def test_auth_builder_db_team_fallback_when_jwt_has_no_team(
fallback_to_db_teams: bool,
user_teams: list,
header_team_id: Optional[str],
expected_team_id: Optional[str],
expect_403: bool,
) -> None:
"""End-to-end auth_builder behavior with no JWT team claims.
fallback_to_db_teams=True attributes usage to the user's first resolvable DB
team, honors a valid x-litellm-team-id header, and rejects a header team the
user does not belong to. The default (flag off) preserves the upstream
single-team fallback: a lone DB team is resolved, multiple are ambiguous.
"""
user_id = "u_db_fallback"
user_object = LiteLLM_UserTable(
user_id=user_id,
user_role=LitellmUserRoles.INTERNAL_USER,
teams=user_teams,
)
jwt_handler = JWTHandler()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
enforce_team_based_model_access=True,
fallback_to_db_teams=fallback_to_db_teams,
)
request_headers = {"x-litellm-team-id": header_team_id} if header_team_id else None
async def fake_get_team(team_id, **kwargs):
return LiteLLM_TeamTable(team_id=team_id)
async def call_auth_builder():
with (
patch.object(
jwt_handler, "auth_jwt", new_callable=AsyncMock
) as mock_auth_jwt,
patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock),
patch.object(jwt_handler, "get_rbac_role", return_value=None),
patch.object(jwt_handler, "get_scopes", return_value=[]),
patch.object(jwt_handler, "get_object_id", return_value=None),
patch.object(
JWTAuthManager,
"get_user_info",
new_callable=AsyncMock,
return_value=(user_id, "u@example.com", True),
),
patch.object(jwt_handler, "get_org_id", return_value=None),
patch.object(jwt_handler, "get_end_user_id", return_value=None),
patch.object(
JWTAuthManager,
"check_admin_access",
new_callable=AsyncMock,
return_value=None,
),
patch.object(
JWTAuthManager,
"find_and_validate_specific_team_id",
new_callable=AsyncMock,
return_value=(None, None),
),
patch.object(JWTAuthManager, "get_all_team_ids", return_value=set()),
patch.object(
JWTAuthManager,
"find_team_with_model_access",
new_callable=AsyncMock,
return_value=(None, None),
),
patch.object(
JWTAuthManager,
"get_objects",
new_callable=AsyncMock,
return_value=(user_object, None, None, None, user_id),
),
patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock),
patch.object(JWTAuthManager, "validate_object_id", return_value=True),
patch.object(
JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock
),
patch(
"litellm.proxy.auth.handle_jwt.get_team_object",
new_callable=AsyncMock,
side_effect=fake_get_team,
),
patch(
"litellm.proxy.auth.handle_jwt.get_team_membership",
new_callable=AsyncMock,
return_value=LiteLLM_TeamMembership(
user_id=user_id,
team_id=user_teams[0] if user_teams else "none",
litellm_budget_table=None,
),
),
):
mock_auth_jwt.return_value = {"sub": user_id, "scope": ""}
return await JWTAuthManager.auth_builder(
api_key="test_jwt_token",
jwt_handler=jwt_handler,
request_data={"model": "gpt-4"},
general_settings={"enforce_rbac": False},
route="/chat/completions",
prisma_client=None,
user_api_key_cache=None,
parent_otel_span=None,
proxy_logging_obj=None,
request_headers=request_headers,
)
if expect_403:
with pytest.raises(HTTPException) as exc_info:
await call_auth_builder()
assert exc_info.value.status_code == 403
else:
result = await call_auth_builder()
assert result["team_id"] == expected_team_id