feat(proxy): bind JWT claims to registered agents via agent_id_jwt_field
Some checks are pending
ai-gateway image / ai-gateway release image (push) Waiting to run
LiteLLM Rust / rust-lint (push) Waiting to run
LiteLLM Rust / rust-test (push) Waiting to run

JWT auth validated Entra app tokens but never carried an agent identity
into the authenticated principal, so agent policies (trace id
requirement, per-agent MCP restrictions, agent spend attribution) only
applied to virtual keys bound to an agent. A new litellm_jwtauth field,
agent_id_jwt_field, names the claim (dot notation supported) that is
matched against a registered agent's id, then name; the canonical
agent_id flows through the standard and proxy-admin JWT paths, and a
configured claim naming no registered agent fails closed with 403

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yassin 2026-09-12 21:26:30 +00:00
parent 30f33a949b
commit eb48850a1c
5 changed files with 305 additions and 2 deletions

View file

@ -4694,6 +4694,7 @@ class JWTAuthBuilderResult(TypedDict):
org_id: str | None
team_membership: LiteLLM_TeamMembership | None
jwt_claims: dict # Decoded JWT token claims (avoids re-decoding)
agent_id: ReadOnly[str | None]
class ClientSideFallbackModel(TypedDict, total=False):
@ -4924,6 +4925,14 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
user_allowed_roles: list[str] | None = None
user_id_upsert: bool = Field(default=False, description="If user doesn't exist, upsert them into the db.")
end_user_id_jwt_field: str | None = None
agent_id_jwt_field: str | None = Field(
default=None,
description=(
"The field in the JWT token that identifies the calling agent (e.g. 'azp' for a Microsoft Entra ID "
"app token). Supports dot notation. The value is matched against a registered agent's agent_id, "
"then agent_name, and the request is rejected when it matches neither."
),
)
public_key_ttl: float = 600
public_key_stale_ttl: float = Field(
default=DEFAULT_JWKS_STALE_TTL,

View file

@ -14,7 +14,7 @@ import hashlib
import os
import re
import time
from collections.abc import Awaitable, Callable, Sequence
from collections.abc import Awaitable, Callable, Mapping, Sequence
from typing import Any, Final, Literal, NoReturn, Protocol, TypeVar, cast
import httpx
@ -51,6 +51,7 @@ from litellm.proxy._types import (
TeamMemberAddRequest,
UserAPIKeyAuth,
)
from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry, global_agent_registry
from litellm.proxy.auth.auth_checks import can_team_access_model
from litellm.proxy.auth.resolvers.grants import GrantResolver, UserLookup, canonical_user_id
from litellm.proxy.auth.route_checks import RouteChecks
@ -623,6 +624,12 @@ class JWTHandler:
object_id = default_value
return object_id
def get_agent_claim(self, token: Mapping[str, object]) -> str | None:
if self.litellm_jwtauth.agent_id_jwt_field is None:
return None
claim: Final[object] = get_nested_value(data=token, key_path=self.litellm_jwtauth.agent_id_jwt_field)
return claim if isinstance(claim, str) and claim else None
def get_org_id(self, token: dict, default_value: str | None) -> str | None:
if self._has_trusted_issuer_normalized_claim(token=token, claim=self.LITELLM_ORG_ID_CLAIM):
return token.get(self.LITELLM_ORG_ID_CLAIM)
@ -1380,6 +1387,7 @@ class JWTAuthManager:
api_key: str,
jwt_valid_token: dict | None = None,
user_email: str | None = None,
agent_id: str | None = None,
) -> JWTAuthBuilderResult | None:
"""Check admin status and route access permissions"""
if not jwt_handler.is_admin(scopes=scopes):
@ -1409,8 +1417,28 @@ class JWTAuthManager:
org_id=org_id,
team_membership=None,
jwt_claims=jwt_valid_token or {},
agent_id=agent_id,
)
@staticmethod
def resolve_agent_id(
jwt_handler: JWTHandler,
jwt_valid_token: Mapping[str, object],
agent_registry: AgentRegistry,
) -> str | None:
agent_claim: Final = jwt_handler.get_agent_claim(token=jwt_valid_token)
if agent_claim is None:
return None
agent: Final = agent_registry.get_agent_by_id(agent_id=agent_claim) or agent_registry.get_agent_by_name(
agent_name=agent_claim
)
if agent is None:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=f"No registered agent matches JWT claim {jwt_handler.litellm_jwtauth.agent_id_jwt_field}={agent_claim}",
)
return agent.agent_id
@staticmethod
async def find_and_validate_specific_team_id(
jwt_handler: JWTHandler,
@ -2209,6 +2237,7 @@ class JWTAuthManager:
proxy_logging_obj: ProxyLogging,
request_headers: dict | None = None,
request_method: str | None = None,
agent_registry: AgentRegistry = global_agent_registry,
) -> JWTAuthBuilderResult:
"""Main authentication and authorization builder"""
# Check if OIDC UserInfo endpoint is enabled, but fall back to standard
@ -2268,9 +2297,21 @@ class JWTAuthManager:
elif rbac_role == LitellmUserRoles.INTERNAL_USER:
user_id = object_id
agent_id: Final = JWTAuthManager.resolve_agent_id(
jwt_handler=jwt_handler, jwt_valid_token=jwt_valid_token, agent_registry=agent_registry
)
# Check admin access
admin_result: Final = await JWTAuthManager.check_admin_access(
jwt_handler, scopes, route, user_id, org_id, api_key, jwt_valid_token, user_email=user_email
jwt_handler,
scopes,
route,
user_id,
org_id,
api_key,
jwt_valid_token,
user_email=user_email,
agent_id=agent_id,
)
if admin_result:
await JWTAuthManager._attach_team_from_header_for_admin(
@ -2514,4 +2555,5 @@ class JWTAuthManager:
token=api_key,
team_membership=team_membership_object,
jwt_claims=jwt_valid_token,
agent_id=agent_id,
)

View file

@ -1559,6 +1559,7 @@ async def _user_api_key_auth_builder(
org_id: Final = result["org_id"]
team_membership: Final[LiteLLM_TeamMembership | None] = result.get("team_membership", None)
jwt_claims = result.get("jwt_claims", None)
agent_id: Final[str | None] = result.get("agent_id")
if is_proxy_admin:
# Proxy admins authenticate via auth_builder (full
@ -1584,6 +1585,7 @@ async def _user_api_key_auth_builder(
end_user_id=end_user_id,
parent_otel_span=parent_otel_span,
jwt_claims=jwt_claims,
agent_id=agent_id,
**team_grants(team_object=team_object, team_membership=team_membership, user_id=user_id),
)
@ -1604,6 +1606,7 @@ async def _user_api_key_auth_builder(
user_rpm_limit=(user_object.rpm_limit if user_object is not None else None),
user_model_max_budget=(user_object.model_max_budget if user_object is not None else None),
jwt_claims=jwt_claims,
agent_id=agent_id,
**team_grants(team_object=team_object, team_membership=team_membership, user_id=user_id),
)

View file

@ -23,6 +23,7 @@ from litellm.proxy._types import (
ProxyException,
)
from litellm.caching.dual_cache import DualCache
from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry
from litellm.proxy.auth.handle_jwt import (
JWKS_FETCH_ATTEMPTS,
STALE_CACHE_KEY_PREFIX,
@ -32,6 +33,7 @@ from litellm.proxy.auth.handle_jwt import (
JWTHandler,
NoMatchingJWTPublicKeyError,
)
from litellm.types.agents import AgentResponse
@pytest.mark.asyncio
@ -6786,3 +6788,180 @@ async def test_sync_user_role_and_teams_singular_claim_only_recognized_under_fla
}
assert mock_patch.call_args.kwargs["teams_ids_to_add_user_to"] == []
assert user.teams == []
def _entra_agent_registry() -> AgentRegistry:
registry = AgentRegistry()
registry.register_agent(
AgentResponse(
agent_id="canonical-agent-id",
agent_name="2f5c9b1e-6a4d-4c8e-9f0b-7d1a3e5c9b21",
agent_card_params={"name": "research-agent", "url": "http://localhost:9999/a2a", "version": "1.0.0"},
litellm_params={"require_trace_id_on_calls_by_agent": True},
)
)
return registry
def _entra_agent_jwt_handler(agent_id_jwt_field: str | None) -> JWTHandler:
jwt_handler = JWTHandler()
jwt_handler.update_environment(
prisma_client=None,
user_api_key_cache=DualCache(),
litellm_jwtauth=LiteLLM_JWTAuth(user_id_jwt_field="sub", agent_id_jwt_field=agent_id_jwt_field),
)
return jwt_handler
@pytest.mark.parametrize(
"claim_value",
["canonical-agent-id", "2f5c9b1e-6a4d-4c8e-9f0b-7d1a3e5c9b21"],
ids=["matches_agent_id", "matches_agent_name"],
)
def test_resolve_agent_id_returns_canonical_agent_id(claim_value: str):
"""An Entra app token's azp claim binds to the registered agent by id or by name and yields its canonical id."""
jwt_handler = _entra_agent_jwt_handler(agent_id_jwt_field="azp")
resolved = JWTAuthManager.resolve_agent_id(
jwt_handler=jwt_handler,
jwt_valid_token={"sub": "sp-object-id-1234", "azp": claim_value},
agent_registry=_entra_agent_registry(),
)
assert resolved == "canonical-agent-id"
def test_resolve_agent_id_reads_nested_claim():
jwt_handler = _entra_agent_jwt_handler(agent_id_jwt_field="entra.client_id")
resolved = JWTAuthManager.resolve_agent_id(
jwt_handler=jwt_handler,
jwt_valid_token={"sub": "sp-object-id-1234", "entra": {"client_id": "canonical-agent-id"}},
agent_registry=_entra_agent_registry(),
)
assert resolved == "canonical-agent-id"
def test_resolve_agent_id_rejects_claim_for_unregistered_agent():
"""A configured agent claim naming no registered agent fails closed with 403 instead of falling back to an unbound identity."""
jwt_handler = _entra_agent_jwt_handler(agent_id_jwt_field="azp")
with pytest.raises(HTTPException) as exc_info:
JWTAuthManager.resolve_agent_id(
jwt_handler=jwt_handler,
jwt_valid_token={"sub": "sp-object-id-1234", "azp": "00000000-0000-0000-0000-000000000000"},
agent_registry=_entra_agent_registry(),
)
assert exc_info.value.status_code == 403
@pytest.mark.parametrize(
"token",
[
{"sub": "sp-object-id-1234"},
{"sub": "sp-object-id-1234", "azp": ""},
{"sub": "sp-object-id-1234", "azp": ["canonical-agent-id"]},
],
ids=["claim_absent", "claim_empty", "claim_not_a_string"],
)
def test_resolve_agent_id_returns_none_when_claim_unusable(token: dict):
jwt_handler = _entra_agent_jwt_handler(agent_id_jwt_field="azp")
assert (
JWTAuthManager.resolve_agent_id(
jwt_handler=jwt_handler, jwt_valid_token=token, agent_registry=_entra_agent_registry()
)
is None
)
def test_resolve_agent_id_ignores_claim_when_field_not_configured():
"""Without agent_id_jwt_field an azp claim (even an unknown one) leaves JWT auth behaviour unchanged."""
jwt_handler = _entra_agent_jwt_handler(agent_id_jwt_field=None)
resolved = JWTAuthManager.resolve_agent_id(
jwt_handler=jwt_handler,
jwt_valid_token={"sub": "sp-object-id-1234", "azp": "00000000-0000-0000-0000-000000000000"},
agent_registry=_entra_agent_registry(),
)
assert resolved is None
def _entra_signed_app_token(monkeypatch, azp: str, scope: str) -> tuple[JWTHandler, str]:
"""A JWTHandler that verifies RS256 tokens against a pre-cached JWKS, plus a signed Entra-style app token."""
jwks_url = "https://login.microsoftonline.test/discovery/v2.0/keys"
monkeypatch.setenv("JWT_PUBLIC_KEY_URL", jwks_url)
monkeypatch.delenv("JWT_AUDIENCE", raising=False)
private_key, jwk = _get_rsa_key_and_jwk(kid="entra-kid")
cache = DualCache()
cache.set_cache(key=f"litellm_jwt_auth_keys_{jwks_url}", value=[jwk])
jwt_handler = JWTHandler()
jwt_handler.update_environment(
prisma_client=None,
user_api_key_cache=cache,
litellm_jwtauth=LiteLLM_JWTAuth(agent_id_jwt_field="azp"),
)
token = _encode_rsa_jwt(
private_key,
issuer="https://login.microsoftonline.test/lit7664-tenant/v2.0",
audience="api://litellm",
kid="entra-kid",
extra_claims={"sub": "sp-object-id-1234", "azp": azp, "scope": scope},
)
return jwt_handler, token
@pytest.mark.asyncio
@pytest.mark.parametrize("is_admin_token", [False, True], ids=["standard_jwt", "proxy_admin_jwt"])
async def test_auth_builder_propagates_agent_id_from_jwt_claim(monkeypatch, is_admin_token: bool):
"""auth_builder carries the resolved agent id into JWTAuthBuilderResult on both the admin and standard paths."""
jwt_handler, token = _entra_signed_app_token(
monkeypatch,
azp="2f5c9b1e-6a4d-4c8e-9f0b-7d1a3e5c9b21",
scope=LiteLLM_JWTAuth().admin_jwt_scope if is_admin_token else "",
)
result = await JWTAuthManager.auth_builder(
api_key=token,
jwt_handler=jwt_handler,
request_data={"model": "gpt-5.6"},
general_settings={"enforce_rbac": False},
route="/key/info" if is_admin_token else "/chat/completions",
prisma_client=None,
user_api_key_cache=None,
parent_otel_span=None,
proxy_logging_obj=None,
agent_registry=_entra_agent_registry(),
)
assert result["is_proxy_admin"] is is_admin_token
assert result["agent_id"] == "canonical-agent-id"
@pytest.mark.asyncio
async def test_auth_builder_denies_jwt_naming_unregistered_agent_before_admin_check(monkeypatch):
"""An unknown agent claim is rejected even when the token would otherwise be a proxy admin."""
jwt_handler, token = _entra_signed_app_token(
monkeypatch,
azp="00000000-0000-0000-0000-000000000000",
scope=LiteLLM_JWTAuth().admin_jwt_scope,
)
with pytest.raises(HTTPException) as exc_info:
await JWTAuthManager.auth_builder(
api_key=token,
jwt_handler=jwt_handler,
request_data={"model": "gpt-5.6"},
general_settings={"enforce_rbac": False},
route="/key/info",
prisma_client=None,
user_api_key_cache=None,
parent_otel_span=None,
proxy_logging_obj=None,
agent_registry=_entra_agent_registry(),
)
assert exc_info.value.status_code == 403

View file

@ -1937,6 +1937,76 @@ async def test_standard_jwt_auth_propagates_user_email():
assert result.api_key is None
@pytest.mark.asyncio
@pytest.mark.parametrize("is_proxy_admin", [False, True], ids=["standard_jwt", "proxy_admin_jwt"])
async def test_jwt_auth_propagates_agent_id_to_user_api_key_auth(is_proxy_admin: bool):
"""The agent id resolved by auth_builder must land on UserAPIKeyAuth.agent_id so
agent-scoped checks (trace id requirement, MCP server/tool restrictions, spend
attribution) apply to JWT callers the same way they apply to agent-bound keys."""
jwt_token = "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.signature"
general_settings = {"enable_jwt_auth": True}
user_api_key_cache = DualCache()
jwt_handler = MagicMock()
jwt_handler.is_jwt.return_value = True
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(agent_id_jwt_field="azp")
user_object = LiteLLM_UserTable(user_id="sp-object-id-1234", user_role="internal_user")
mock_jwt_result = {
"is_proxy_admin": is_proxy_admin,
"team_object": None,
"user_object": user_object,
"end_user_object": None,
"org_object": None,
"token": jwt_token,
"team_id": None,
"user_id": "sp-object-id-1234",
"user_email": None,
"end_user_id": None,
"org_id": None,
"team_membership": None,
"jwt_claims": {"sub": "sp-object-id-1234", "azp": "2f5c9b1e-6a4d-4c8e-9f0b-7d1a3e5c9b21"},
"agent_id": "canonical-agent-id",
}
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.multiple( # test-quality-ok: production auth reads these module globals; no dependency injection seam exists
"litellm.proxy.proxy_server",
general_settings=general_settings,
premium_user=True,
master_key="sk-master",
prisma_client=None,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=MagicMock(),
jwt_handler=jwt_handler,
),
patch( # test-quality-ok: the builder calls this static method directly; no dependency injection seam exists
"litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder",
new_callable=AsyncMock,
return_value=mock_jwt_result,
),
):
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-5.6"},
)
assert result.agent_id == "canonical-agent-id"
assert result.user_id == "sp-object-id-1234"
assert result.api_key is None
@pytest.mark.asyncio
async def test_auto_register_binds_api_key_to_token_hash():
"""