mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
feat(proxy): bind JWT claims to registered agents via agent_id_jwt_field
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:
parent
30f33a949b
commit
eb48850a1c
5 changed files with 305 additions and 2 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue