feat(agents): validate identity lifecycle and enroll trusted SSO subjects

This commit is contained in:
Joshua Valluru 2026-09-26 12:02:55 -07:00
parent 397567f462
commit 724accfd0c
6 changed files with 477 additions and 6 deletions

View file

@ -1,15 +1,189 @@
from collections.abc import Mapping
from typing import Final, NoReturn
from datetime import datetime
from typing import Final, NoReturn, TypedDict
from uuid import uuid4
from fastapi import HTTPException
from pydantic import TypeAdapter, ValidationError
from typing_extensions import ReadOnly
from litellm.types.agents import AgentResponse
from litellm.types.proxy.agent_identity import (
AgentExecutionMode,
AgentIdentityBinding,
AgentIdentityFailure,
AgentSubject,
EntraIdentityConfig,
)
_MODE: Final = TypeAdapter(AgentExecutionMode)
class IdentityFields(TypedDict, total=False):
provider: ReadOnly[str]
tenant_id: ReadOnly[str]
client_id: ReadOnly[str]
issuer: ReadOnly[str]
service_principal_id: ReadOnly[str | None]
required_roles: ReadOnly[tuple[str, ...]]
required_scopes: ReadOnly[tuple[str, ...]]
active: ReadOnly[bool]
revision: ReadOnly[str]
last_authenticated_at: ReadOnly[datetime | None]
class IdentityUpsert(TypedDict):
create: ReadOnly[IdentityFields]
update: ReadOnly[IdentityFields]
class IdentityRelationWrite(TypedDict, total=False):
create: ReadOnly[IdentityFields]
update: ReadOnly[IdentityFields]
upsert: ReadOnly[IdentityUpsert]
class IdentityHistoryKey(TypedDict):
provider: ReadOnly[str]
tenant_id: ReadOnly[str]
client_id: ReadOnly[str]
class IdentityHistoryWhere(TypedDict):
provider_tenant_id_client_id: ReadOnly[IdentityHistoryKey]
class IdentityHistoryEntry(IdentityHistoryKey):
issuer: ReadOnly[str]
class IdentityHistoryConnect(TypedDict):
where: ReadOnly[IdentityHistoryWhere]
create: ReadOnly[IdentityHistoryEntry]
class IdentityHistoryWrite(TypedDict):
connectOrCreate: ReadOnly[IdentityHistoryConnect]
class ManagedWriteFields(TypedDict, total=False):
enabled: ReadOnly[bool]
execution_mode: ReadOnly[AgentExecutionMode]
identity_managed: ReadOnly[bool]
identity: ReadOnly[IdentityRelationWrite]
retired_identities: ReadOnly[IdentityHistoryWrite]
def raise_identity_failure(failure: AgentIdentityFailure, status_code: int = 403) -> NoReturn:
raise HTTPException(503 if failure.code == "policy_unavailable" else status_code, failure.message)
def _configuration_failure(
identity: EntraIdentityConfig | AgentIdentityBinding | None,
mode: AgentExecutionMode,
enabling_without_binding: bool,
) -> AgentIdentityFailure | None:
if identity is not None and mode != "delegated" and not identity.service_principal_id:
return AgentIdentityFailure(
message="Autonomous mode requires the Enterprise application service-principal object ID"
)
if identity is not None and mode != "autonomous" and not identity.required_scopes:
return AgentIdentityFailure(message="Delegated mode requires at least one delegated scope")
if enabling_without_binding and (
identity is None or isinstance(identity, AgentIdentityBinding) and not identity.active
):
return AgentIdentityFailure(message="Bind an identity before enabling this managed agent")
return None
def managed_write_fields(
incoming: Mapping[str, object],
existing: AgentResponse | None,
updated_by: str,
) -> ManagedWriteFields | AgentIdentityFailure:
try:
identity: Final = (
EntraIdentityConfig.model_validate(incoming["identity"]) if incoming.get("identity") is not None else None
)
mode: Final = _MODE.validate_python(
incoming.get("execution_mode", existing.execution_mode if existing else "autonomous")
)
current_identity: Final = identity if "identity" in incoming else existing.identity if existing else None
failure: Final = _configuration_failure(
current_identity,
mode,
incoming.get("enabled") is True
and "identity" not in incoming
and bool(existing and existing.identity_managed),
)
if failure is not None:
return failure
empty: Final[ManagedWriteFields] = {}
identity_fields: Final = _identity_write(identity, existing) if "identity" in incoming else empty
result: Final[ManagedWriteFields] = {
**({"enabled": incoming["enabled"] is True} if "enabled" in incoming else {}),
**({"execution_mode": mode} if "execution_mode" in incoming else {}),
**identity_fields,
}
return result
except (ValidationError, ValueError) as exc:
return AgentIdentityFailure(message=f"Invalid agent identity configuration: {exc}")
def _identity_write(identity: EntraIdentityConfig | None, existing: AgentResponse | None) -> ManagedWriteFields:
if identity is None:
unbind: Final[ManagedWriteFields] = {
**(
{"identity": {"update": {"active": False, "revision": str(uuid4()), "last_authenticated_at": None}}}
if existing and existing.identity
else {}
),
**({"identity_managed": True, "enabled": False} if existing and existing.identity_managed else {}),
}
return unbind
if (
existing
and existing.identity
and existing.identity.active
and all(getattr(existing.identity, name) == value for name, value in identity.model_dump().items())
):
unchanged: Final[ManagedWriteFields] = {}
return unchanged
binding: Final[IdentityFields] = {
"provider": identity.provider,
"tenant_id": identity.tenant_id,
"client_id": identity.client_id,
"service_principal_id": identity.service_principal_id,
"required_roles": identity.required_roles,
"required_scopes": identity.required_scopes,
"issuer": identity.issuer,
"active": True,
"revision": str(uuid4()),
"last_authenticated_at": None,
}
result: Final[ManagedWriteFields] = {
"retired_identities": {
"connectOrCreate": {
"where": {
"provider_tenant_id_client_id": {
"provider": identity.provider,
"tenant_id": identity.tenant_id,
"client_id": identity.client_id,
}
},
"create": {
"provider": identity.provider,
"issuer": identity.issuer,
"tenant_id": identity.tenant_id,
"client_id": identity.client_id,
},
}
},
"identity_managed": True,
"identity": {"upsert": {"create": binding, "update": binding}} if existing else {"create": binding},
}
return result
def classify_agent_subject(
binding: AgentIdentityBinding,
@ -46,7 +220,3 @@ def classify_agent_subject(
if not frozenset(binding.required_roles).issubset(roles):
return AgentIdentityFailure(message="Token lacks the required application roles")
return AgentSubject(kind="application", oid=oid, mode="autonomous")
def raise_identity_failure(failure: AgentIdentityFailure, status_code: int = 403) -> NoReturn:
raise HTTPException(503 if failure.code == "policy_unavailable" else status_code, failure.message)

View file

@ -0,0 +1,49 @@
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final
from uuid import UUID
from litellm.types.proxy.agent_identity import MicrosoftInteractiveSubject
def microsoft_interactive_subject(
tenant: str | None,
response: Mapping[str, object],
endpoints: Mapping[str, str | None],
) -> MicrosoftInteractiveSubject | None:
if tenant is None:
return None
try:
tenant_id: Final = str(UUID(tenant))
object_id: Final = response.get("id")
if not isinstance(object_id, str):
return None
oid: Final = str(UUID(object_id))
except ValueError:
return None
expected: Final = MappingProxyType(
{
"MICROSOFT_AUTHORIZATION_ENDPOINT": f"https://login.microsoftonline.com/{tenant_id}/oauth2/v2.0/authorize",
"MICROSOFT_TOKEN_ENDPOINT": f"https://login.microsoftonline.com/{tenant_id}/oauth2/v2.0/token",
"MICROSOFT_USERINFO_ENDPOINT": "https://graph.microsoft.com/v1.0/me",
}
)
if any(value and value != expected.get(name) for name, value in endpoints.items()):
return None
return MicrosoftInteractiveSubject(
issuer=f"https://login.microsoftonline.com/{tenant_id}/v2.0",
tenant_id=tenant_id,
oid=oid,
)
async def enroll_microsoft_subject(subject: object, user_id: object, client: object) -> None:
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
from litellm.types.proxy.agent_identity import AgentIdentityFailure
if not isinstance(subject, MicrosoftInteractiveSubject) or not isinstance(user_id, str) or not user_id:
return
result: Final = await AgentIdentityStore.from_client(client).enroll_interactive_human(subject, user_id)
if isinstance(result, AgentIdentityFailure):
raise_identity_failure(result)

View file

@ -3631,6 +3631,12 @@ class SSOAuthenticationHandler:
},
)
from litellm.proxy.management_endpoints.sso.agent_subject_enrollment import enroll_microsoft_subject
await enroll_microsoft_subject(
request.scope.get("litellm_microsoft_interactive_subject"), user_id, prisma_client
)
if isinstance(user_id, str) and user_id:
await retain_sso_identity_assertion_for_ema(user_id=user_id, assertion=sso_assertion)
await warn_if_id_jag_assertion_uncaptured(sso_assertion)
@ -4300,6 +4306,22 @@ class MicrosoftSSOHandler:
original_msft_result["app_roles"] = app_roles
return original_msft_result or {}
from litellm.proxy.management_endpoints.sso.agent_subject_enrollment import microsoft_interactive_subject
request.scope["litellm_microsoft_interactive_subject"] = microsoft_interactive_subject(
microsoft_tenant,
original_msft_result,
MappingProxyType(
{
name: os.getenv(name)
for name in (
"MICROSOFT_AUTHORIZATION_ENDPOINT",
"MICROSOFT_TOKEN_ENDPOINT",
"MICROSOFT_USERINFO_ENDPOINT",
)
}
),
)
result: Final = MicrosoftSSOHandler.openid_from_response(
response=original_msft_result,
team_ids=user_team_ids,

View file

@ -2,7 +2,8 @@ from typing import Final
import pytest
from litellm.proxy.agent_endpoints.managed_identity import classify_agent_subject
from litellm.proxy.agent_endpoints.managed_identity import classify_agent_subject, managed_write_fields
from litellm.types.agents import AgentResponse
from litellm.types.proxy.agent_identity import (
AgentExecutionMode,
AgentIdentityBinding,
@ -93,6 +94,118 @@ def test_native_facet_absence_does_not_establish_human_identity() -> None:
assert result.kind == "delegated_subject"
def managed_agent() -> AgentResponse:
return AgentResponse(
agent_id="agent-one", agent_name="Research", agent_card_params={}, identity=BINDING, identity_managed=True
)
def test_unbinding_keeps_managed_state_and_disables_agent() -> None:
result: Final = managed_write_fields({"identity": None, "enabled": True}, managed_agent(), "admin")
assert not isinstance(result, AgentIdentityFailure)
assert result["identity_managed"] is True
assert result["enabled"] is False
assert result["identity"]["update"]["active"] is False
assert result["identity"]["update"]["last_authenticated_at"] is None
assert result["identity"]["update"]["revision"] != BINDING.revision
def test_rename_does_not_rewrite_binding_or_evidence() -> None:
assert managed_write_fields({"agent_name": "Renamed"}, managed_agent(), "admin") == {}
def test_autonomous_binding_requires_enterprise_application_object_id() -> None:
result: Final = managed_write_fields(
{"identity": {"provider": "microsoft_entra", "tenant_id": TENANT, "client_id": CLIENT}}, None, "admin"
)
assert isinstance(result, AgentIdentityFailure)
assert "service-principal" in result.message
def test_rebinding_clears_evidence_and_uses_atomic_nested_write() -> None:
result: Final = managed_write_fields(
{
"identity": {
"provider": "microsoft_entra",
"tenant_id": TENANT,
"client_id": CLIENT,
"service_principal_id": PRINCIPAL,
}
},
managed_agent(),
"admin",
)
assert not isinstance(result, AgentIdentityFailure)
assert result["identity_managed"] is True
assert "upsert" in result["identity"]
assert result["identity"]["upsert"]["update"]["revision"] != BINDING.revision
assert result["identity"]["upsert"]["update"]["last_authenticated_at"] is None
def test_unbound_identity_can_be_reactivated_with_the_same_application() -> None:
disabled: Final = managed_agent().model_copy(
update={"identity": BINDING.model_copy(update={"active": False}), "enabled": False}
)
configuration: Final = BINDING.model_dump(
exclude={"agent_id", "issuer", "revision", "last_authenticated_at", "active"}
)
result: Final = managed_write_fields({"identity": configuration, "enabled": True}, disabled, "admin")
assert not isinstance(result, AgentIdentityFailure)
assert result["enabled"] is True
assert result["identity"]["upsert"]["update"]["active"] is True
assert result["identity"]["upsert"]["update"]["revision"] != BINDING.revision
def test_each_application_binding_records_its_history_atomically() -> None:
configuration: Final = BINDING.model_dump(
exclude={"agent_id", "issuer", "revision", "last_authenticated_at", "active"}
)
created: Final = managed_write_fields({"identity": configuration}, None, "admin")
assert not isinstance(created, AgentIdentityFailure)
assert created["retired_identities"]["connectOrCreate"]["create"]["client_id"] == CLIENT
replacement: Final = managed_write_fields(
{"identity": {**configuration, "client_id": HUMAN}}, managed_agent(), "admin"
)
assert not isinstance(replacement, AgentIdentityFailure)
assert replacement["retired_identities"]["connectOrCreate"]["create"]["client_id"] == HUMAN
def test_unchanged_binding_preserves_revision_and_authentication_evidence() -> None:
configuration: Final = BINDING.model_dump(
exclude={"agent_id", "issuer", "revision", "last_authenticated_at", "active"}
)
assert managed_write_fields({"identity": configuration}, managed_agent(), "admin") == {}
@pytest.mark.parametrize("identity", [None, BINDING.model_copy(update={"active": False})])
def test_enabling_unbound_or_inactive_identity_requires_rebinding(identity: AgentIdentityBinding | None) -> None:
agent: Final = managed_agent().model_copy(update={"identity": identity, "enabled": False})
result: Final = managed_write_fields({"enabled": True}, agent, "admin")
assert isinstance(result, AgentIdentityFailure)
assert "Bind an identity" in result.message
def test_delegated_identity_requires_a_scope() -> None:
agent: Final = managed_agent().model_copy(update={"identity": BINDING.model_copy(update={"required_scopes": ()})})
result: Final = managed_write_fields({"execution_mode": "delegated"}, agent, "admin")
assert isinstance(result, AgentIdentityFailure)
assert "delegated scope" in result.message
@pytest.mark.parametrize(
"incoming",
[
{"identity": {"provider": "microsoft_entra", "tenant_id": "invalid", "client_id": CLIENT}},
{"execution_mode": "unknown"},
],
)
def test_invalid_identity_configuration_returns_a_public_validation_failure(incoming: dict[str, object]) -> None:
result: Final = managed_write_fields(incoming, None, "admin")
assert isinstance(result, AgentIdentityFailure)
assert result.code == "identity_denied"
assert result.message.startswith("Invalid agent identity configuration:")
@pytest.mark.parametrize("roles", ["Agent.Invoke", [42], None])
def test_malformed_application_roles_are_rejected(roles: object) -> None:
result: Final = classify_agent_subject(BINDING, claims(roles=roles), "autonomous")

View file

@ -0,0 +1,114 @@
from types import SimpleNamespace
from typing import Final
from unittest.mock import AsyncMock
import pytest
from fastapi import HTTPException
from litellm.proxy.management_endpoints.sso.agent_subject_enrollment import (
enroll_microsoft_subject,
microsoft_interactive_subject,
)
TENANT: Final = "11111111-1111-4111-8111-111111111111"
OID: Final = "22222222-2222-4222-8222-222222222222"
def test_enrollment_uses_provider_object_id_and_configured_tenant() -> None:
subject: Final = microsoft_interactive_subject(
TENANT, {"id": OID, "mail": "alias@example.com", "tid": "untrusted"}, {}
)
assert subject is not None
assert subject.oid == OID
assert subject.tenant_id == TENANT
assert subject.issuer == f"https://login.microsoftonline.com/{TENANT}/v2.0"
@pytest.mark.parametrize("tenant", [None, "common", "organizations", "invalid"])
def test_multitenant_sso_does_not_guess_the_subject_tenant(tenant: str | None) -> None:
assert microsoft_interactive_subject(tenant, {"id": OID, "tid": TENANT}, {}) is None
@pytest.mark.parametrize("response", [{"mail": "user@example.com"}, {"id": "user@example.com"}, {"id": 42}])
def test_email_and_configurable_aliases_are_not_human_subject_proof(response: dict[str, object]) -> None:
assert microsoft_interactive_subject(TENANT, response, {}) is None
@pytest.mark.parametrize(
"endpoint", ["MICROSOFT_USERINFO_ENDPOINT", "MICROSOFT_TOKEN_ENDPOINT", "MICROSOFT_AUTHORIZATION_ENDPOINT"]
)
def test_custom_provider_endpoints_do_not_enroll_trusted_microsoft_subjects(endpoint: str) -> None:
assert microsoft_interactive_subject(TENANT, {"id": OID}, {endpoint: "https://custom.example"}) is None
@pytest.mark.asyncio
async def test_interactive_enrollment_preserves_the_canonical_local_user() -> None:
table: Final = AsyncMock()
table.upsert.return_value = SimpleNamespace(kind="human", user_id="canonical", verified_via="sso_interactive")
client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_verifiedsubject=table))
subject: Final = microsoft_interactive_subject(TENANT, {"id": OID}, {})
assert subject is not None
await enroll_microsoft_subject(subject, "canonical", client)
table.upsert.assert_awaited_once_with(
where={"issuer_tenant_id_oid": {"issuer": subject.issuer, "tenant_id": TENANT, "oid": OID}},
data={
"create": {
"issuer": subject.issuer,
"tenant_id": TENANT,
"oid": OID,
"user_id": "canonical",
"verified_via": "sso_interactive",
},
"update": {},
},
)
@pytest.mark.asyncio
@pytest.mark.parametrize("user_id,verified_via", [("another-user", "sso_interactive"), ("canonical", "untrusted")])
async def test_interactive_enrollment_does_not_reassign_an_existing_subject(user_id: str, verified_via: str) -> None:
table: Final = AsyncMock()
table.upsert.return_value = SimpleNamespace(kind="human", user_id=user_id, verified_via=verified_via)
client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_verifiedsubject=table))
with pytest.raises(HTTPException) as failure:
await enroll_microsoft_subject(microsoft_interactive_subject(TENANT, {"id": OID}, {}), "canonical", client)
assert failure.value.status_code == 403
assert table.upsert.call_args.kwargs["data"]["update"] == {}
@pytest.mark.asyncio
async def test_enrollment_storage_failure_is_not_a_successful_login() -> None:
table: Final = AsyncMock()
table.upsert.side_effect = RuntimeError("database unavailable")
client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_verifiedsubject=table))
with pytest.raises(HTTPException) as failure:
await enroll_microsoft_subject(microsoft_interactive_subject(TENANT, {"id": OID}, {}), "canonical", client)
assert failure.value.status_code == 503
@pytest.mark.asyncio
@pytest.mark.parametrize("user_id", [None, "", 42])
async def test_enrollment_requires_a_canonical_local_user(user_id: object) -> None:
table: Final = AsyncMock()
client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_verifiedsubject=table))
await enroll_microsoft_subject(microsoft_interactive_subject(TENANT, {"id": OID}, {}), user_id, client)
table.upsert.assert_not_awaited()
@pytest.mark.asyncio
async def test_untrusted_metadata_cannot_enroll_a_human() -> None:
table: Final = AsyncMock()
client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_verifiedsubject=table))
await enroll_microsoft_subject({"issuer": "forged", "tenant_id": TENANT, "oid": OID}, "canonical", client)
table.upsert.assert_not_awaited()
@pytest.mark.asyncio
async def test_scim_agent_subject_cannot_be_enrolled_as_a_human() -> None:
table: Final = AsyncMock()
table.upsert.return_value = SimpleNamespace(kind="agent_user", user_id=None, verified_via="scim")
client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_verifiedsubject=table))
with pytest.raises(HTTPException) as failure:
await enroll_microsoft_subject(microsoft_interactive_subject(TENANT, {"id": OID}, {}), "canonical", client)
assert failure.value.status_code == 403
assert table.upsert.call_args.kwargs["data"]["update"] == {}

View file

@ -206,6 +206,7 @@ def test_microsoft_sso_handler_openid_from_response_with_custom_attributes():
def test_get_microsoft_callback_response():
# Arrange
mock_request = MagicMock(spec=Request)
mock_request.scope = {}
mock_response = {
"mail": "microsoft_user@example.com",
"displayName": "Microsoft User",
@ -8751,6 +8752,7 @@ async def test_redirect_from_openid_persists_assertion_under_canonical_user_id()
assertion = assertion_from_sso_login(_ema_id_token(), "rt_1")
assert assertion is not None
mock_request = MagicMock(spec=Request)
mock_request.scope = {}
mock_request.base_url = "http://localhost:4000/"
mock_request.cookies = {}
@ -8989,6 +8991,7 @@ async def test_browser_funnel_reports_an_uncaptured_assertion(monkeypatch, caplo
"""Wiring: the browser login path must reach the diagnostic, not just define it."""
monkeypatch.setenv("GOOGLE_CLIENT_ID", "cid")
mock_request = MagicMock(spec=Request)
mock_request.scope = {}
mock_request.base_url = "http://localhost:4000/"
mock_request.cookies = {}