diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 9046d522280..4b3c3df81f4 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -4524,6 +4524,16 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase): default=None, description="Optional claim-based routing overrides for JWT-shaped tokens. Matching rules route requests to oauth2 before default JWT flow.", ) + team_claim_fallback: bool = Field( + default=False, + description=( + "If True, when a configured team_id_jwt_field / team_ids_jwt_field " + "claim is present but does not resolve to any known team, defer to " + "the single-team DB fallback (caller's only team membership) " + "instead of raising. Default False preserves strict claim-based " + "authorization." + ), + ) ######################################################### def __init__(self, **kwargs: Any) -> None: diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 2deb7d40004..06ad36031e8 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -1048,7 +1048,10 @@ class JWTAuthManager: ) return individual_team_id, team_object except HTTPException as e: - if e.status_code != 404: + if ( + e.status_code != 404 + or not jwt_handler.litellm_jwtauth.team_claim_fallback + ): raise # Claim doesn't map to a known team — defer to fallback. verbose_proxy_logger.debug( @@ -1181,14 +1184,17 @@ class JWTAuthManager: except Exception: continue - if requested_model and any_claim_team_resolved: - # Claim teams resolved but none grant the model — deny. + if requested_model and ( + any_claim_team_resolved + or not jwt_handler.litellm_jwtauth.team_claim_fallback + ): + # Claim resolved but no model access, or fallback disabled — deny. raise HTTPException( status_code=403, detail=f"No team has access to the requested model: {requested_model}. Checked teams={team_ids}. Check `/models` to see all available models.", ) - # No claim team resolved — defer to fallback. + # No claim team resolved and fallback enabled — defer to fallback. return None, None @staticmethod diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index 8cbe48cf9cb..e106b171539 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -2764,13 +2764,17 @@ def test_build_decode_kwargs_no_warning_when_scoped( @pytest.mark.asyncio async def test_find_and_validate_specific_team_id_unresolved_claim_returns_none(): - """team_id claim is present in the JWT but the team is missing in the - DB — return (None, None) so the auth_builder single-team fallback can run, - instead of raising and failing auth.""" + """With `team_claim_fallback=True`: team_id claim is present in the JWT + but the team is missing in the DB — return (None, None) so the + auth_builder single-team fallback can run, instead of raising and + failing auth.""" from fastapi import HTTPException jwt_handler = JWTHandler() - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_id_jwt_field="team_id") + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + team_id_jwt_field="team_id", + team_claim_fallback=True, + ) token = {"sub": "user-1", "team_id": "claim-team-not-in-db"} with patch( @@ -2796,8 +2800,9 @@ async def test_find_and_validate_specific_team_id_unresolved_claim_returns_none( async def test_find_team_with_model_access_unresolved_group_claim_returns_none( monkeypatch, ): - """Group claim resolves to team_ids that don't exist in the DB — return - (None, None) instead of raising 403, so the single-team fallback can run.""" + """With `team_claim_fallback=True`: group claim resolves to team_ids that + don't exist in the DB — return (None, None) instead of raising 403, so + the single-team fallback can run.""" import sys import types @@ -2820,7 +2825,7 @@ async def test_find_team_with_model_access_unresolved_group_claim_returns_none( monkeypatch.setattr("litellm.proxy.auth.handle_jwt.get_team_object", raise_404) jwt_handler = JWTHandler() - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_claim_fallback=True) team_id, team_object = await JWTAuthManager.find_team_with_model_access( team_ids={"idp-group-a", "idp-group-b"}, @@ -2977,3 +2982,79 @@ async def test_find_team_with_model_access_resolved_team_without_model_still_rai assert exc_info.value.status_code == 403 assert "No team has access to the requested model" in str(exc_info.value.detail) + + +@pytest.mark.asyncio +async def test_find_and_validate_specific_team_id_unresolved_claim_default_raises(): + """Default `team_claim_fallback=False`: unresolved team_id claim must + still raise — preserves the strict claim-based authorization boundary + when the operator has not opted in to the fallback.""" + from fastapi import HTTPException + + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_id_jwt_field="team_id") + token = {"sub": "user-1", "team_id": "claim-team-not-in-db"} + + with patch( + "litellm.proxy.auth.handle_jwt.get_team_object", + new_callable=AsyncMock, + ) as mock_get_team: + mock_get_team.side_effect = HTTPException(status_code=404, detail="missing") + + with pytest.raises(HTTPException) as exc_info: + await JWTAuthManager.find_and_validate_specific_team_id( + jwt_handler=jwt_handler, + jwt_valid_token=token, + prisma_client=None, + user_api_key_cache=None, + parent_otel_span=None, + proxy_logging_obj=None, + ) + + assert exc_info.value.status_code == 404 + + +@pytest.mark.asyncio +async def test_find_team_with_model_access_unresolved_group_claim_default_raises( + monkeypatch, +): + """Default `team_claim_fallback=False`: group claims that don't resolve + to any LiteLLM team must still raise 403 — preserves the strict + claim-based authorization boundary.""" + import sys + import types + + from fastapi import HTTPException + + from litellm.router import Router + + router = Router( + model_list=[ + {"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}} + ] + ) + proxy_server_module = types.ModuleType("proxy_server") + proxy_server_module.llm_router = router + monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_module) + + async def raise_404(*_args, **_kwargs): + raise HTTPException(status_code=404, detail="missing") + + monkeypatch.setattr("litellm.proxy.auth.handle_jwt.get_team_object", raise_404) + + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() + + with pytest.raises(HTTPException) as exc_info: + await JWTAuthManager.find_team_with_model_access( + team_ids={"idp-group-a", "idp-group-b"}, + requested_model="gpt-4o-mini", + route="/chat/completions", + jwt_handler=jwt_handler, + prisma_client=None, + user_api_key_cache=None, + parent_otel_span=None, + proxy_logging_obj=None, + ) + + assert exc_info.value.status_code == 403