mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-15 23:31:29 +00:00
Merge pull request #40927 from BerriAI/litellm_per_issuer_virtual_key_claim_field
feat(jwt): allow virtual_key_claim_field per issuer
This commit is contained in:
commit
9e1ed40db3
4 changed files with 324 additions and 8 deletions
|
|
@ -4837,6 +4837,14 @@ class JWTIssuerConfig(BaseModel):
|
|||
default=None,
|
||||
description="Issuer-specific claim path to normalize into LiteLLM's end-user id.",
|
||||
)
|
||||
virtual_key_claim_field: str | None = Field(
|
||||
default=None,
|
||||
description="Issuer-specific claim path used for the virtual key mapping lookup. Falls back to the global field.",
|
||||
)
|
||||
unregistered_jwt_client_behavior: UnregisteredJWTClientBehavior | None = Field(
|
||||
default=None,
|
||||
description="Issuer-specific policy when the virtual key claim has no mapping. Falls back to the global policy.",
|
||||
)
|
||||
|
||||
model_config = {
|
||||
"extra": "forbid",
|
||||
|
|
@ -5063,6 +5071,28 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
|
|||
|
||||
super().__init__(**kwargs)
|
||||
|
||||
def get_issuer_config(self, issuer: str | None) -> JWTIssuerConfig | None:
|
||||
if issuer is None or self.issuers is None:
|
||||
return None
|
||||
return next((config for config in self.issuers if config.issuer == issuer), None)
|
||||
|
||||
def is_virtual_key_mapping_configured(self) -> bool:
|
||||
if self.virtual_key_claim_field is not None:
|
||||
return True
|
||||
return any(config.virtual_key_claim_field is not None for config in self.issuers or ())
|
||||
|
||||
def get_virtual_key_claim_field(self, issuer: str | None) -> str | None:
|
||||
issuer_config: Final = self.get_issuer_config(issuer)
|
||||
if issuer_config is not None and issuer_config.virtual_key_claim_field is not None:
|
||||
return issuer_config.virtual_key_claim_field
|
||||
return self.virtual_key_claim_field
|
||||
|
||||
def get_unregistered_jwt_client_behavior(self, issuer: str | None) -> UnregisteredJWTClientBehavior:
|
||||
issuer_config: Final = self.get_issuer_config(issuer)
|
||||
if issuer_config is not None and issuer_config.unregistered_jwt_client_behavior is not None:
|
||||
return issuer_config.unregistered_jwt_client_behavior
|
||||
return self.unregistered_jwt_client_behavior
|
||||
|
||||
|
||||
class PrismaCompatibleUpdateDBModel(TypedDict, total=False):
|
||||
model_name: str
|
||||
|
|
|
|||
|
|
@ -996,9 +996,12 @@ async def _resolve_jwt_to_virtual_key(
|
|||
- Raises HTTPException: REJECT policy hit, missing claim under
|
||||
REJECT/AUTO_REGISTER, or other policy violations.
|
||||
"""
|
||||
virtual_key_claim_field: Final = jwt_handler.litellm_jwtauth.virtual_key_claim_field
|
||||
raw_issuer: Final = jwt_claims.get(JWTHandler.LITELLM_JWT_ISSUER_CLAIM)
|
||||
normalized_issuer: Final = raw_issuer if isinstance(raw_issuer, str) else None
|
||||
virtual_key_claim_field: Final = jwt_handler.litellm_jwtauth.get_virtual_key_claim_field(normalized_issuer)
|
||||
if virtual_key_claim_field is None:
|
||||
return None
|
||||
behavior: Final = jwt_handler.litellm_jwtauth.get_unregistered_jwt_client_behavior(normalized_issuer)
|
||||
|
||||
claim_value: Final = get_nested_value(
|
||||
data=jwt_claims,
|
||||
|
|
@ -1015,7 +1018,6 @@ async def _resolve_jwt_to_virtual_key(
|
|||
# simply by presenting a JWT that omits the configured field. For
|
||||
# AUTO_REGISTER there is no stable identity to map without a claim
|
||||
# value, so we deny rather than create a sentinel-keyed record.
|
||||
behavior = jwt_handler.litellm_jwtauth.unregistered_jwt_client_behavior
|
||||
if behavior in (
|
||||
UnregisteredJWTClientBehavior.REJECT,
|
||||
UnregisteredJWTClientBehavior.AUTO_REGISTER,
|
||||
|
|
@ -1030,7 +1032,13 @@ async def _resolve_jwt_to_virtual_key(
|
|||
return None
|
||||
|
||||
cache_key: Final = jwt_key_mapping_cache_key(virtual_key_claim_field, str(claim_value))
|
||||
cached_mapping: Final = await user_api_key_cache.async_get_cache(cache_key)
|
||||
raw_cached_mapping: Final = await user_api_key_cache.async_get_cache(cache_key)
|
||||
sentinel_written_by_this_policy: Final = behavior == UnregisteredJWTClientBehavior.AUTO_REGISTER
|
||||
cached_mapping: Final = (
|
||||
None
|
||||
if raw_cached_mapping == _JWT_PROXY_ADMIN_SENTINEL and not sentinel_written_by_this_policy
|
||||
else raw_cached_mapping
|
||||
)
|
||||
|
||||
if cached_mapping == _JWT_PROXY_ADMIN_SENTINEL:
|
||||
# Previously resolved to a proxy admin via auth_builder; skip the
|
||||
|
|
@ -1039,7 +1047,6 @@ async def _resolve_jwt_to_virtual_key(
|
|||
return None
|
||||
|
||||
if cached_mapping == "__NO_MAPPING__":
|
||||
behavior = jwt_handler.litellm_jwtauth.unregistered_jwt_client_behavior
|
||||
if behavior == UnregisteredJWTClientBehavior.REJECT:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
|
|
@ -1102,8 +1109,6 @@ async def _resolve_jwt_to_virtual_key(
|
|||
)
|
||||
|
||||
# No mapping found (DB miss or no DB) — apply no-match policy.
|
||||
behavior = jwt_handler.litellm_jwtauth.unregistered_jwt_client_behavior
|
||||
|
||||
if behavior == UnregisteredJWTClientBehavior.REJECT:
|
||||
# Cache the miss before raising so repeated rejections are served from
|
||||
# cache and don't re-query the DB on every request.
|
||||
|
|
@ -1483,7 +1488,7 @@ async def _user_api_key_auth_builder(
|
|||
# unnecessary DB queries in auth_builder
|
||||
do_standard_jwt_auth = True
|
||||
pending_auto_register: _PendingAutoRegister | None = None
|
||||
if jwt_handler.litellm_jwtauth.virtual_key_claim_field is not None:
|
||||
if jwt_handler.litellm_jwtauth.is_virtual_key_mapping_configured():
|
||||
# Decode JWT to get claims without running full auth_builder
|
||||
jwt_claims: dict | None
|
||||
if jwt_handler.litellm_jwtauth.oidc_userinfo_enabled and not is_jwt:
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ from unittest.mock import ANY, AsyncMock, MagicMock, patch
|
|||
|
||||
|
||||
import pytest
|
||||
from fastapi import status
|
||||
from fastapi import HTTPException, status
|
||||
|
||||
import litellm
|
||||
import litellm.proxy.proxy_server
|
||||
|
|
@ -7628,3 +7628,226 @@ async def test_claude_view_never_reinterprets_explicit_names(monkeypatch, layer)
|
|||
assert data["model"] == ("foo" if layer == "unclaimed" else encoded)
|
||||
await _normalize_claude_model(data, token, request, "/v1/messages")
|
||||
assert data["model"] == ("foo" if layer == "unclaimed" else encoded)
|
||||
|
||||
|
||||
ISSUER_ONE = "https://issuer-one.example.com"
|
||||
ISSUER_TWO = "https://issuer-two.example.com"
|
||||
|
||||
|
||||
def _per_issuer_virtual_key_jwt_handler(
|
||||
global_claim_field: str | None, global_behavior: str = "fallback_team_mapping"
|
||||
) -> MagicMock:
|
||||
jwt_handler = MagicMock()
|
||||
jwt_handler.is_jwt.return_value = True
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
|
||||
virtual_key_claim_field=global_claim_field,
|
||||
unregistered_jwt_client_behavior=global_behavior,
|
||||
issuers=[
|
||||
{
|
||||
"issuer": ISSUER_ONE,
|
||||
"jwks_url": f"{ISSUER_ONE}/keys",
|
||||
"audience": "audience-one",
|
||||
"team_id_jwt_field": "sub",
|
||||
},
|
||||
{
|
||||
"issuer": ISSUER_TWO,
|
||||
"jwks_url": f"{ISSUER_TWO}/keys",
|
||||
"audience": "audience-two",
|
||||
"virtual_key_claim_field": "sub",
|
||||
"unregistered_jwt_client_behavior": "reject",
|
||||
},
|
||||
],
|
||||
)
|
||||
return jwt_handler
|
||||
|
||||
|
||||
def _fake_prisma_with_jwt_key_mapping(hashed_token: str | None) -> tuple[SimpleNamespace, AsyncMock]:
|
||||
find_first = AsyncMock(return_value=None if hashed_token is None else SimpleNamespace(token=hashed_token))
|
||||
prisma_client = SimpleNamespace(db=SimpleNamespace(litellm_jwtkeymapping=SimpleNamespace(find_first=find_first)))
|
||||
return prisma_client, find_first
|
||||
|
||||
|
||||
def _mapping_where(claim_name: str, claim_value: str) -> dict[str, str | bool]:
|
||||
return {"jwt_claim_name": claim_name, "jwt_claim_value": claim_value, "is_active": True}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_per_issuer_virtual_key_claim_field_selects_the_issuer_mapping_for_the_db_lookup():
|
||||
from litellm.proxy.auth.user_api_key_auth import _resolve_jwt_to_virtual_key
|
||||
|
||||
jwt_handler = _per_issuer_virtual_key_jwt_handler(global_claim_field=None)
|
||||
prisma_client, find_first = _fake_prisma_with_jwt_key_mapping("hashed-mapped-key")
|
||||
user_api_key_cache = DualCache()
|
||||
await user_api_key_cache.async_set_cache(
|
||||
key="hashed-mapped-key",
|
||||
value=UserAPIKeyAuth(token="hashed-mapped-key", api_key="hashed-mapped-key", team_id="svc-team"),
|
||||
)
|
||||
|
||||
resolved = await _resolve_jwt_to_virtual_key(
|
||||
jwt_claims={JWTHandler.LITELLM_JWT_ISSUER_CLAIM: ISSUER_TWO, "sub": "svc-account-7"},
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
find_first.assert_awaited_once_with(where=_mapping_where("sub", "svc-account-7"))
|
||||
assert isinstance(resolved, UserAPIKeyAuth)
|
||||
assert resolved.token == "hashed-mapped-key"
|
||||
assert resolved.team_id == "svc-team"
|
||||
assert await user_api_key_cache.async_get_cache("jwt_key_mapping:sub:svc-account-7") == "hashed-mapped-key"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_per_issuer_reject_behavior_does_not_leak_into_the_team_issuer():
|
||||
from litellm.proxy.auth.user_api_key_auth import _resolve_jwt_to_virtual_key
|
||||
|
||||
jwt_handler = _per_issuer_virtual_key_jwt_handler(global_claim_field=None)
|
||||
prisma_client, find_first = _fake_prisma_with_jwt_key_mapping(None)
|
||||
|
||||
team_issuer_result = await _resolve_jwt_to_virtual_key(
|
||||
jwt_claims={JWTHandler.LITELLM_JWT_ISSUER_CLAIM: ISSUER_ONE, "sub": "team-alpha"},
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=DualCache(),
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
assert team_issuer_result is None
|
||||
find_first.assert_not_awaited()
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await _resolve_jwt_to_virtual_key(
|
||||
jwt_claims={JWTHandler.LITELLM_JWT_ISSUER_CLAIM: ISSUER_TWO, "sub": "unknown-svc"},
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=DualCache(),
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
assert exc.value.status_code == 403
|
||||
assert "No registered mapping for sub='unknown-svc'" in str(exc.value.detail)
|
||||
find_first.assert_awaited_once_with(where=_mapping_where("sub", "unknown-svc"))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_proxy_admin_sentinel_cached_by_another_issuer_does_not_bypass_reject():
|
||||
from litellm.proxy.auth.user_api_key_auth import _JWT_PROXY_ADMIN_SENTINEL, _resolve_jwt_to_virtual_key
|
||||
|
||||
jwt_handler = _per_issuer_virtual_key_jwt_handler(global_claim_field="sub", global_behavior="auto_register")
|
||||
prisma_client, find_first = _fake_prisma_with_jwt_key_mapping(None)
|
||||
user_api_key_cache = DualCache()
|
||||
await user_api_key_cache.async_set_cache(key="jwt_key_mapping:sub:admin-7", value=_JWT_PROXY_ADMIN_SENTINEL)
|
||||
|
||||
auto_register_issuer_result = await _resolve_jwt_to_virtual_key(
|
||||
jwt_claims={JWTHandler.LITELLM_JWT_ISSUER_CLAIM: ISSUER_ONE, "sub": "admin-7"},
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
assert auto_register_issuer_result is None
|
||||
find_first.assert_not_awaited()
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await _resolve_jwt_to_virtual_key(
|
||||
jwt_claims={JWTHandler.LITELLM_JWT_ISSUER_CLAIM: ISSUER_TWO, "sub": "admin-7"},
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
assert exc.value.status_code == 403
|
||||
assert "No registered mapping for sub='admin-7'" in str(exc.value.detail)
|
||||
find_first.assert_awaited_once_with(where=_mapping_where("sub", "admin-7"))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_issuer_without_virtual_key_claim_field_falls_back_to_the_global_field():
|
||||
from litellm.proxy.auth.user_api_key_auth import _resolve_jwt_to_virtual_key
|
||||
|
||||
jwt_handler = _per_issuer_virtual_key_jwt_handler(global_claim_field="client_id")
|
||||
prisma_client, find_first = _fake_prisma_with_jwt_key_mapping(None)
|
||||
|
||||
with_claim = await _resolve_jwt_to_virtual_key(
|
||||
jwt_claims={JWTHandler.LITELLM_JWT_ISSUER_CLAIM: ISSUER_ONE, "sub": "team-alpha", "client_id": "app-9"},
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=DualCache(),
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
without_claim = await _resolve_jwt_to_virtual_key(
|
||||
jwt_claims={JWTHandler.LITELLM_JWT_ISSUER_CLAIM: ISSUER_ONE, "sub": "team-alpha"},
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=DualCache(),
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
assert with_claim is None
|
||||
assert without_claim is None
|
||||
find_first.assert_awaited_once_with(where=_mapping_where("client_id", "app-9"))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auth_flow_enters_virtual_key_mapping_when_only_an_issuer_configures_the_claim_field():
|
||||
jwt_token = "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJzdmMtYWNjb3VudC03In0.signature"
|
||||
jwt_handler = _per_issuer_virtual_key_jwt_handler(global_claim_field=None)
|
||||
jwt_handler.auth_jwt = AsyncMock(
|
||||
return_value={JWTHandler.LITELLM_JWT_ISSUER_CLAIM: ISSUER_TWO, "sub": "svc-account-7"}
|
||||
)
|
||||
mapped_key = UserAPIKeyAuth(token="hashed-mapped-key", api_key="hashed-mapped-key", team_id="svc-team")
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request.url.path = "/v1/chat/completions"
|
||||
mock_request.method = "POST"
|
||||
mock_request.headers = {"authorization": f"Bearer {jwt_token}"}
|
||||
mock_request.query_params = {}
|
||||
mock_request.state = SimpleNamespace()
|
||||
|
||||
with (
|
||||
patch( # test-quality-ok: the builder reads proxy settings from module globals, no injection seam
|
||||
"litellm.proxy.proxy_server.general_settings", {"enable_jwt_auth": True}
|
||||
),
|
||||
patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: module-global proxy state
|
||||
patch("litellm.proxy.proxy_server.master_key", "sk-master"), # test-quality-ok: module-global proxy state
|
||||
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), # test-quality-ok: module-global proxy state
|
||||
patch( # test-quality-ok: module-global proxy state
|
||||
"litellm.proxy.proxy_server.user_api_key_cache", DualCache()
|
||||
),
|
||||
patch( # test-quality-ok: module-global proxy state
|
||||
"litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()
|
||||
),
|
||||
patch("litellm.proxy.proxy_server.jwt_handler", jwt_handler), # test-quality-ok: module-global proxy state
|
||||
patch( # test-quality-ok: the regression is whether the builder reaches this seam at all
|
||||
"litellm.proxy.auth.user_api_key_auth._resolve_jwt_to_virtual_key",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mapped_key,
|
||||
) as resolve_mock,
|
||||
patch( # test-quality-ok: a mapped key must short-circuit standard JWT auth; reaching it is the failure
|
||||
"litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder",
|
||||
new_callable=AsyncMock,
|
||||
side_effect=AssertionError("standard JWT auth must not run for a mapped virtual key"),
|
||||
),
|
||||
):
|
||||
result = await _user_api_key_auth_builder(
|
||||
request=mock_request,
|
||||
api_key=jwt_token,
|
||||
azure_api_key_header="",
|
||||
anthropic_api_key_header=None,
|
||||
google_ai_studio_api_key_header=None,
|
||||
azure_apim_header=None,
|
||||
request_data={"model": "gpt-4o-mini"},
|
||||
)
|
||||
|
||||
resolve_mock.assert_awaited_once()
|
||||
assert resolve_mock.await_args.kwargs["jwt_claims"][JWTHandler.LITELLM_JWT_ISSUER_CLAIM] == ISSUER_TWO
|
||||
assert result.api_key == "hashed-mapped-key"
|
||||
assert result.team_id == "svc-team"
|
||||
|
|
|
|||
|
|
@ -277,3 +277,61 @@ def test_team_membership_budget_table_present_still_works():
|
|||
}
|
||||
result = LiteLLM_TeamMembership.model_validate(data)
|
||||
assert result.litellm_budget_table is None
|
||||
|
||||
|
||||
def test_a_jwt_issuer_can_override_the_virtual_key_claim_field_while_other_issuers_keep_the_global_one():
|
||||
from litellm.proxy._types import LiteLLM_JWTAuth, UnregisteredJWTClientBehavior
|
||||
|
||||
jwt_auth = LiteLLM_JWTAuth(
|
||||
virtual_key_claim_field="client_id",
|
||||
issuers=[
|
||||
{
|
||||
"issuer": "https://team-idp.example.com",
|
||||
"jwks_url": "https://team-idp.example.com/keys",
|
||||
"audience": "litellm",
|
||||
"team_id_jwt_field": "sub",
|
||||
},
|
||||
{
|
||||
"issuer": "https://service-idp.example.com",
|
||||
"jwks_url": "https://service-idp.example.com/keys",
|
||||
"audience": "litellm",
|
||||
"virtual_key_claim_field": "sub",
|
||||
"unregistered_jwt_client_behavior": "reject",
|
||||
},
|
||||
],
|
||||
)
|
||||
|
||||
assert jwt_auth.get_virtual_key_claim_field("https://service-idp.example.com") == "sub"
|
||||
assert jwt_auth.get_unregistered_jwt_client_behavior("https://service-idp.example.com") is (
|
||||
UnregisteredJWTClientBehavior.REJECT
|
||||
)
|
||||
assert jwt_auth.get_virtual_key_claim_field("https://team-idp.example.com") == "client_id"
|
||||
assert jwt_auth.get_unregistered_jwt_client_behavior("https://team-idp.example.com") is (
|
||||
UnregisteredJWTClientBehavior.FALLBACK_TEAM_MAPPING
|
||||
)
|
||||
assert jwt_auth.get_virtual_key_claim_field(None) == "client_id"
|
||||
assert jwt_auth.get_virtual_key_claim_field("https://unknown-idp.example.com") == "client_id"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("global_field", "issuer_field", "is_configured"),
|
||||
((None, None, False), ("sub", None, True), (None, "sub", True)),
|
||||
)
|
||||
def test_virtual_key_mapping_counts_as_configured_when_any_issuer_sets_the_claim_field(
|
||||
global_field, issuer_field, is_configured
|
||||
):
|
||||
from litellm.proxy._types import LiteLLM_JWTAuth
|
||||
|
||||
jwt_auth = LiteLLM_JWTAuth(
|
||||
virtual_key_claim_field=global_field,
|
||||
issuers=[
|
||||
{
|
||||
"issuer": "https://idp.example.com",
|
||||
"jwks_url": "https://idp.example.com/keys",
|
||||
"audience": "litellm",
|
||||
"virtual_key_claim_field": issuer_field,
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
assert jwt_auth.is_virtual_key_mapping_configured() is is_configured
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue