mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
feat(agents): validate identity lifecycle and enroll trusted SSO subjects
This commit is contained in:
parent
e0ec32bfe1
commit
0e17665007
6 changed files with 477 additions and 6 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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"] == {}
|
||||
|
|
@ -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 = {}
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue