fix(jwt-auth): gate claim-mismatch fallback behind opt-in flag
Some checks failed
Unit Tests: Proxy DB Operations / assert-shard-coverage (push) Has been cancelled
Unit Tests: Security / security (push) Has been cancelled
Unit Tests: Proxy DB Operations / auth-checks (push) Has been cancelled
Unit Tests: Proxy DB Operations / budgets (push) Has been cancelled
Unit Tests: Proxy DB Operations / custom-logging (push) Has been cancelled
Unit Tests: Proxy DB Operations / db-and-spend (push) Has been cancelled
Unit Tests: Proxy DB Operations / endpoints-and-responses (push) Has been cancelled
Unit Tests: Proxy DB Operations / guardrails-hooks (push) Has been cancelled
Unit Tests: Proxy DB Operations / jwt-and-keys (push) Has been cancelled
Unit Tests: Proxy DB Operations / key-generation (push) Has been cancelled
Unit Tests: Proxy DB Operations / logging-misc (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-runtime (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-server-core (push) Has been cancelled
Unit Tests: Proxy DB Operations / schema-migration (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-utils (push) Has been cancelled

The unresolved-team-claim fallback added in the previous commit
weakened the strict claim-based authorization contract by default —
an authenticated user whose JWT carries a stale or invalid team
claim could still consume their single DB team's models/quota via
the fallback.

Gate both soft-fail paths in `find_and_validate_specific_team_id`
and `find_team_with_model_access` behind a new opt-in flag
`team_claim_fallback` on `LiteLLM_JWTAuth` (default False).

Default-off preserves the pre-existing strict behavior. Operators
who intentionally treat IdP team claims as advisory (e.g. machine
tokens whose group claims live in a separate namespace from
LiteLLM team_ids) opt in via config.

Adds two regression tests guarding the default-off behavior.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Milan 2026-05-27 12:35:59 +03:00
parent 639a18af1d
commit 491d79850b
No known key found for this signature in database
3 changed files with 108 additions and 11 deletions

View file

@ -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:

View file

@ -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

View file

@ -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