mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(jwt): AUTO_REGISTER inherits team_id so keys are bounded by team limits
Auto-registered virtual keys were created with no team, model, route, rate, or budget constraints — broader access than the standard team-based JWT auth path the same client would have taken. Under AUTO_REGISTER, resolve the team_id from the JWT (via the operator-configured team_id_jwt_field / team_id_default) and stamp it on the new key. Downstream auth then applies the team's budget/models/tpm/rpm/allowed_routes via the existing virtual-key flow. Policy when team_id_jwt_field is configured: - JWT carries team claim → stamp resolved team_id - JWT lacks claim + team_id_default set → stamp default - JWT lacks claim + no default → 403 (refuse to create an unbounded key) When neither team_id_jwt_field nor team_id_default is configured, the operator has explicitly opted out of team-based limits — the auto-created key has no team_id (matches what team-auth would do in the same config). Adds 4 tests covering each branch. Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
parent
2c23811880
commit
2c42c5bd4b
2 changed files with 250 additions and 6 deletions
|
|
@ -587,6 +587,40 @@ async def check_api_key_for_custom_headers_or_pass_through_endpoints(
|
|||
return api_key
|
||||
|
||||
|
||||
def _resolve_inherited_team_id(
|
||||
jwt_handler: JWTHandler,
|
||||
jwt_claims: dict,
|
||||
) -> Optional[str]:
|
||||
"""
|
||||
Resolve the team_id the JWT would have been bound to under standard team-
|
||||
based auth. AUTO_REGISTER stamps this on the new key so it inherits the
|
||||
team's budget/model/tpm/rpm/route limits — auto-registered clients then
|
||||
have the same constraints they'd have via the JWT team-auth path.
|
||||
|
||||
If ``team_id_jwt_field`` is configured but the JWT lacks the claim and no
|
||||
``team_id_default`` is set, raise 403: the operator has declared every JWT
|
||||
must carry a team, and registering a teamless (unbounded) key here would
|
||||
silently grant broader access than the team-auth path allows.
|
||||
"""
|
||||
if (
|
||||
jwt_handler.litellm_jwtauth.team_id_jwt_field is None
|
||||
and jwt_handler.litellm_jwtauth.team_id_default is None
|
||||
):
|
||||
return None
|
||||
|
||||
team_id = jwt_handler.get_team_id(token=jwt_claims, default_value=None)
|
||||
if team_id is None:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=(
|
||||
"JWT Key Mapping: AUTO_REGISTER requires a team_id on the JWT "
|
||||
"(via team_id_jwt_field) or a configured team_id_default — "
|
||||
"refusing to create an unbounded key. Access denied."
|
||||
),
|
||||
)
|
||||
return team_id
|
||||
|
||||
|
||||
async def _auto_register_jwt_mapping(
|
||||
virtual_key_claim_field: str,
|
||||
claim_value: str,
|
||||
|
|
@ -596,15 +630,19 @@ async def _auto_register_jwt_mapping(
|
|||
parent_otel_span: Optional[Span],
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
cache_key: str,
|
||||
team_id: Optional[str] = None,
|
||||
) -> Optional[UserAPIKeyAuth]:
|
||||
"""
|
||||
Auto-register: create a new virtual key + mapping for an unrecognised JWT claim value.
|
||||
The new key carries no model/budget restrictions; admins can tighten it later.
|
||||
Auto-register: create a new virtual key + mapping for an unrecognised JWT
|
||||
claim value. The key is stamped with the JWT's resolved ``team_id`` (when
|
||||
available) so it inherits the team's budget/model/tpm/rpm/route limits via
|
||||
the standard auth path — auto-registered clients are then no more permissive
|
||||
than the team-based JWT auth path they would otherwise have taken.
|
||||
|
||||
Race safety: if two concurrent requests both reach here simultaneously (both saw
|
||||
no mapping in the DB), one will win the unique-constraint race on
|
||||
litellm_jwtkeymapping. The loser catches the conflict, fetches the winner's
|
||||
mapping, and proceeds — no orphaned keys and no error surfaced to the caller.
|
||||
Race safety: if two concurrent requests both reach here simultaneously (both
|
||||
saw no mapping in the DB), one will win the unique-constraint race on
|
||||
litellm_jwtkeymapping. The loser catches the conflict, deletes its orphaned
|
||||
key, fetches the winner's mapping, and proceeds — no error surfaced.
|
||||
"""
|
||||
# Inline import required: key_management_endpoints imports user_api_key_auth
|
||||
# (line 51) so a module-level import here would create a circular dependency.
|
||||
|
|
@ -614,6 +652,7 @@ async def _auto_register_jwt_mapping(
|
|||
|
||||
key_data = await generate_key_helper_fn(
|
||||
request_type="key",
|
||||
team_id=team_id,
|
||||
metadata={
|
||||
"auto_registered": True,
|
||||
"jwt_claim_field": virtual_key_claim_field,
|
||||
|
|
@ -757,6 +796,7 @@ async def _resolve_jwt_to_virtual_key(
|
|||
"Configure a database or change unregistered_jwt_client_behavior."
|
||||
),
|
||||
)
|
||||
inherited_team_id = _resolve_inherited_team_id(jwt_handler, jwt_claims)
|
||||
await user_api_key_cache.async_delete_cache(cache_key)
|
||||
return await _auto_register_jwt_mapping(
|
||||
virtual_key_claim_field=virtual_key_claim_field,
|
||||
|
|
@ -767,6 +807,7 @@ async def _resolve_jwt_to_virtual_key(
|
|||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
cache_key=cache_key,
|
||||
team_id=inherited_team_id,
|
||||
)
|
||||
return None
|
||||
elif cached_mapping is not None:
|
||||
|
|
@ -827,6 +868,7 @@ async def _resolve_jwt_to_virtual_key(
|
|||
"Configure a database or change unregistered_jwt_client_behavior."
|
||||
),
|
||||
)
|
||||
inherited_team_id = _resolve_inherited_team_id(jwt_handler, jwt_claims)
|
||||
return await _auto_register_jwt_mapping(
|
||||
virtual_key_claim_field=virtual_key_claim_field,
|
||||
claim_value=str(claim_value),
|
||||
|
|
@ -836,6 +878,7 @@ async def _resolve_jwt_to_virtual_key(
|
|||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
cache_key=cache_key,
|
||||
team_id=inherited_team_id,
|
||||
)
|
||||
|
||||
# FALLBACK_TEAM_MAPPING (default): cache the miss and return None so the
|
||||
|
|
|
|||
|
|
@ -1069,6 +1069,207 @@ async def test_auto_register_race_conflict_tolerates_delete_failure():
|
|||
prisma_client.db.litellm_verificationtoken.delete.assert_called_once()
|
||||
|
||||
|
||||
# ──────────────────────────────────────────────
|
||||
# Tests: AUTO_REGISTER inherits team_id from JWT
|
||||
# ──────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auto_register_stamps_team_id_from_jwt_claim():
|
||||
"""
|
||||
When team_id_jwt_field is configured and the JWT carries the claim,
|
||||
AUTO_REGISTER must stamp the resolved team_id on the new key so it inherits
|
||||
the team's limits via standard auth (budget, models, tpm/rpm, allowed_routes).
|
||||
"""
|
||||
from litellm.proxy._types import UnregisteredJWTClientBehavior
|
||||
|
||||
jwt_handler = JWTHandler()
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
|
||||
virtual_key_claim_field="sub",
|
||||
team_id_jwt_field="team_id",
|
||||
unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.AUTO_REGISTER,
|
||||
)
|
||||
jwt_claims = {"sub": "new-user", "team_id": "team-engineering"}
|
||||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None)
|
||||
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
|
||||
|
||||
user_api_key_cache = DualCache()
|
||||
mock_key_obj = UserAPIKeyAuth(token="hashed", team_id="team-engineering")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.get_key_object",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get_key,
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_gen_key,
|
||||
):
|
||||
mock_gen_key.return_value = {"token": "sk-newkey", "key": "sk-newkey"}
|
||||
mock_get_key.return_value = mock_key_obj
|
||||
|
||||
result = await _resolve_jwt_to_virtual_key(
|
||||
jwt_claims=jwt_claims,
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
|
||||
assert result == mock_key_obj
|
||||
# generate_key_helper_fn must have been called with team_id from the JWT,
|
||||
# so the key inherits team-level limits via the standard auth path.
|
||||
mock_gen_key.assert_called_once()
|
||||
assert mock_gen_key.call_args.kwargs["team_id"] == "team-engineering"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auto_register_rejects_when_team_id_required_but_missing_from_jwt():
|
||||
"""
|
||||
Security: when team_id_jwt_field is configured (operator wants every JWT
|
||||
bound to a team), AUTO_REGISTER must refuse a JWT that lacks the team
|
||||
claim. Otherwise the auto-created key would be teamless and inherit no
|
||||
budget/model/rate limits — broader access than the team-auth path allows.
|
||||
"""
|
||||
from litellm.proxy._types import UnregisteredJWTClientBehavior
|
||||
|
||||
jwt_handler = JWTHandler()
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
|
||||
virtual_key_claim_field="sub",
|
||||
team_id_jwt_field="team_id",
|
||||
unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.AUTO_REGISTER,
|
||||
)
|
||||
jwt_claims = {"sub": "new-user"} # no "team_id"
|
||||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_gen_key,
|
||||
pytest.raises(HTTPException) as exc_info,
|
||||
):
|
||||
await _resolve_jwt_to_virtual_key(
|
||||
jwt_claims=jwt_claims,
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=DualCache(),
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "team_id" in exc_info.value.detail
|
||||
# No key was created — we refused before generate_key_helper_fn ran
|
||||
mock_gen_key.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auto_register_uses_team_id_default_when_jwt_lacks_team_claim():
|
||||
"""
|
||||
When team_id_jwt_field is configured but the JWT lacks the claim and
|
||||
team_id_default is set, the default team_id is stamped on the new key —
|
||||
the explicit operator-chosen fallback team bounds the auto-registered key.
|
||||
"""
|
||||
from litellm.proxy._types import UnregisteredJWTClientBehavior
|
||||
|
||||
jwt_handler = JWTHandler()
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
|
||||
virtual_key_claim_field="sub",
|
||||
team_id_jwt_field="team_id",
|
||||
team_id_default="default-team",
|
||||
unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.AUTO_REGISTER,
|
||||
)
|
||||
jwt_claims = {"sub": "new-user"} # no "team_id"
|
||||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None)
|
||||
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
|
||||
|
||||
mock_key_obj = UserAPIKeyAuth(token="hashed", team_id="default-team")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.get_key_object",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get_key,
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_gen_key,
|
||||
):
|
||||
mock_gen_key.return_value = {"token": "sk-newkey", "key": "sk-newkey"}
|
||||
mock_get_key.return_value = mock_key_obj
|
||||
|
||||
result = await _resolve_jwt_to_virtual_key(
|
||||
jwt_claims=jwt_claims,
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=DualCache(),
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
|
||||
assert result == mock_key_obj
|
||||
assert mock_gen_key.call_args.kwargs["team_id"] == "default-team"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auto_register_no_team_id_when_team_field_not_configured():
|
||||
"""
|
||||
When the operator has NOT configured team_id_jwt_field or team_id_default,
|
||||
they have explicitly opted out of team-based bounding. AUTO_REGISTER then
|
||||
creates a teamless key (preserving prior behavior) — there is no team to
|
||||
inherit from. The key being unbounded matches what the team-auth path
|
||||
would have produced in the same config.
|
||||
"""
|
||||
from litellm.proxy._types import UnregisteredJWTClientBehavior
|
||||
|
||||
jwt_handler = JWTHandler()
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
|
||||
virtual_key_claim_field="sub",
|
||||
unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.AUTO_REGISTER,
|
||||
)
|
||||
jwt_claims = {"sub": "new-user", "team_id": "ignored-because-not-configured"}
|
||||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None)
|
||||
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
|
||||
|
||||
mock_key_obj = UserAPIKeyAuth(token="hashed", team_id=None)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.get_key_object",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get_key,
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_gen_key,
|
||||
):
|
||||
mock_gen_key.return_value = {"token": "sk-newkey", "key": "sk-newkey"}
|
||||
mock_get_key.return_value = mock_key_obj
|
||||
|
||||
await _resolve_jwt_to_virtual_key(
|
||||
jwt_claims=jwt_claims,
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=DualCache(),
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
|
||||
# team_id is None because the operator opted out of team-based config
|
||||
assert mock_gen_key.call_args.kwargs["team_id"] is None
|
||||
|
||||
|
||||
# ──────────────────────────────────────────────
|
||||
# Tests: backward-compat alias jwt_client_id_field
|
||||
# ──────────────────────────────────────────────
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue