From a4492b9b3c85cf201d39c0a01435f7b5c9c944e8 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 19:26:41 -0700 Subject: [PATCH] fix(agents): preserve retired identity ownership --- .../proxy/agent_endpoints/agent_registry.py | 37 ++++++-- .../proxy/agent_endpoints/managed_identity.py | 30 ++---- .../agent_endpoints/test_agent_registry.py | 91 ++++++++++++++++++- .../agent_endpoints/test_managed_identity.py | 4 +- 4 files changed, 129 insertions(+), 33 deletions(-) diff --git a/litellm/proxy/agent_endpoints/agent_registry.py b/litellm/proxy/agent_endpoints/agent_registry.py index ef7128a0e8e..2cfd468cf3d 100644 --- a/litellm/proxy/agent_endpoints/agent_registry.py +++ b/litellm/proxy/agent_endpoints/agent_registry.py @@ -22,7 +22,11 @@ from litellm.proxy.management_helpers.object_permission_utils import ( from litellm.proxy.utils import PrismaClient from litellm.repositories.base_repository import is_unique_violation from litellm.repositories.prisma_protocols import TableActions -from litellm.repositories.table_repositories import AgentsRepository, ObjectPermissionRepository +from litellm.repositories.table_repositories import ( + AgentsRepository, + ObjectPermissionRepository, + RetiredAgentIdentityRepository, +) from litellm.types.agents import AgentConfig, AgentKillSwitchConfig, AgentResponse, PatchAgentRequest from litellm.types.proxy.agent_identity import AgentIdentityFailure @@ -161,15 +165,33 @@ async def _permission_write( return updated -def _managed_fields( +async def _managed_fields( incoming: Mapping[str, object], existing: AgentResponse | None, updated_by: str, + client: PrismaClient, ) -> Mapping[str, object]: result: Final = managed_write_fields(incoming, existing, updated_by) if isinstance(result, AgentIdentityFailure): raise_identity_failure(result, 400) - return result + history: Final = result.get("retired_identities") + if history is None: + return result + entry: Final = history["create"] + prior: Final = await RetiredAgentIdentityRepository(client, use_writer=True).table.find_unique( + where={ + "provider_tenant_id_client_id": { + "provider": entry["provider"], + "tenant_id": entry["tenant_id"], + "client_id": entry["client_id"], + } + } + ) + if prior is None: + return result + if existing is None or prior.agent_id != existing.agent_id: + raise HTTPException(409, "Entra application was already registered to another agent") + return {key: value for key, value in result.items() if key != "retired_identities"} def _dump_agent_params(raw: Mapping[str, object]) -> dict[str, object]: @@ -631,7 +653,7 @@ class AgentRegistry: # Create agent in DB created_agent: Final = await agents_table(prisma_client).create( - data={**create_data, **_managed_fields(agent, None, created_by)}, + data={**create_data, **await _managed_fields(agent, None, created_by, prisma_client)}, include={"object_permission": True, "identity": True}, ) @@ -743,7 +765,9 @@ class AgentRegistry: where={"agent_id": agent_id}, data={ **update_data, - **_managed_fields(agent, AgentResponse.model_validate(existing_record.model_dump()), updated_by), + **await _managed_fields( + agent, AgentResponse.model_validate(existing_record.model_dump()), updated_by, prisma_client + ), "updated_by": updated_by, "updated_at": datetime.now(timezone.utc), }, @@ -839,10 +863,11 @@ class AgentRegistry: where={"agent_id": agent_id}, data={ **update_data, - **_managed_fields( + **await _managed_fields( agent, AgentResponse.model_validate(existing_row.model_dump()) if existing_row else None, updated_by, + prisma_client, ), }, include={"object_permission": True, "identity": True}, diff --git a/litellm/proxy/agent_endpoints/managed_identity.py b/litellm/proxy/agent_endpoints/managed_identity.py index 260b74fcbd1..abab21901ee 100644 --- a/litellm/proxy/agent_endpoints/managed_identity.py +++ b/litellm/proxy/agent_endpoints/managed_identity.py @@ -49,21 +49,12 @@ class IdentityHistoryKey(TypedDict): 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] + create: ReadOnly[IdentityHistoryEntry] class ManagedWriteFields(TypedDict, total=False): @@ -161,20 +152,11 @@ def _identity_write(identity: EntraIdentityConfig | None, existing: AgentRespons } 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, - }, + "create": { + "provider": identity.provider, + "issuer": identity.issuer, + "tenant_id": identity.tenant_id, + "client_id": identity.client_id, } }, "identity_managed": True, diff --git a/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py b/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py index ba160b8e1b4..370376c8389 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py @@ -1527,7 +1527,15 @@ async def test_duplicate_agent_binding_returns_conflict_for_every_write(operatio registry: Final = AgentRegistry() client: Final = MagicMock() client.db.litellm_agentstable.find_unique = AsyncMock(return_value=_stored_agent_row({"agent_id": "agent-123"})) - failure: Final = UniqueViolationError({"user_facing_error": {"message": "Unique constraint failed", "meta": {"target": ["client_id"]}, "error_code": "P2002"}}) + failure: Final = UniqueViolationError( + { + "user_facing_error": { + "message": "Unique constraint failed", + "meta": {"target": ["client_id"]}, + "error_code": "P2002", + } + } + ) client.db.litellm_agentstable.create = AsyncMock(side_effect=failure) client.db.litellm_agentstable.update = AsyncMock(side_effect=failure) incoming: Final = {"agent_name": "Agent", "agent_card_params": {}} @@ -1542,3 +1550,84 @@ async def test_duplicate_agent_binding_returns_conflict_for_every_write(operatio await write assert denied.value.status_code == 409 assert denied.value.detail == "Agent name or Entra application is already registered" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["create", "patch", "put"]) +@pytest.mark.parametrize("owner", ["previous-agent", None]) +async def test_retired_application_cannot_transfer_to_another_agent(operation: str, owner: str | None) -> None: + from fastapi import HTTPException + + registry: Final = AgentRegistry() + client: Final = MagicMock() + row: Final = _stored_agent_row({"agent_id": "agent-123"}) + client.db.litellm_agentstable.find_unique = AsyncMock(return_value=row) + client.db.litellm_agentstable.create = AsyncMock(return_value=row) + client.db.litellm_agentstable.update = AsyncMock(return_value=row) + client.writer_db.litellm_retiredagentidentity.find_unique = AsyncMock(return_value=SimpleNamespace(agent_id=owner)) + incoming: Final = { + "agent_name": "Agent", + "agent_card_params": {}, + "identity": { + "provider": "microsoft_entra", + "tenant_id": "11111111-1111-4111-8111-111111111111", + "client_id": "22222222-2222-4222-8222-222222222222", + "service_principal_id": "33333333-3333-4333-8333-333333333333", + }, + } + write: Final = ( + registry.add_agent_to_db(incoming, client, created_by="admin") + if operation == "create" + else (registry.patch_agent_in_db if operation == "patch" else registry.update_agent_in_db)( + "agent-123", incoming, client, updated_by="admin" + ) + ) + with pytest.raises(HTTPException) as denied: + await write + assert denied.value.status_code == 409 + client.db.litellm_agentstable.create.assert_not_awaited() + client.db.litellm_agentstable.update.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["create", "patch", "put"]) +@pytest.mark.parametrize("prior_owner", [False, True]) +async def test_application_registration_preserves_its_existing_owner(operation: str, prior_owner: bool) -> None: + registry: Final = AgentRegistry() + client: Final = MagicMock() + row: Final = _stored_agent_row({"agent_id": "agent-123"}) + client.db.litellm_agentstable.find_unique = AsyncMock(return_value=row) + client.db.litellm_agentstable.create = AsyncMock(return_value=row) + client.db.litellm_agentstable.update = AsyncMock(return_value=row) + client.writer_db.litellm_retiredagentidentity.find_unique = AsyncMock( + return_value=SimpleNamespace(agent_id="agent-123") if prior_owner and operation != "create" else None + ) + incoming: Final = { + "agent_name": "Agent", + "agent_card_params": {}, + "identity": { + "provider": "microsoft_entra", + "tenant_id": "11111111-1111-4111-8111-111111111111", + "client_id": "22222222-2222-4222-8222-222222222222", + "service_principal_id": "33333333-3333-4333-8333-333333333333", + }, + } + if operation == "create": + result: Final = await registry.add_agent_to_db(incoming, client, created_by="admin") + else: + update: Final = registry.patch_agent_in_db if operation == "patch" else registry.update_agent_in_db + result = await update("agent-123", incoming, client, updated_by="admin") + assert result.agent_id == "agent-123" + write: Final = ( + client.db.litellm_agentstable.create if operation == "create" else client.db.litellm_agentstable.update + ) + data: Final = write.call_args.kwargs["data"] + if prior_owner and operation != "create": + assert "retired_identities" not in data + else: + assert data["retired_identities"] == { + "create": { + **{key: value for key, value in incoming["identity"].items() if key != "service_principal_id"}, + "issuer": "https://login.microsoftonline.com/11111111-1111-4111-8111-111111111111/v2.0", + } + } diff --git a/tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py b/tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py index 45fe4b0655f..17f3cdb52f5 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py @@ -162,12 +162,12 @@ def test_each_application_binding_records_its_history_atomically() -> None: ) created: Final = managed_write_fields({"identity": configuration}, None, "admin") assert not isinstance(created, AgentIdentityFailure) - assert created["retired_identities"]["connectOrCreate"]["create"]["client_id"] == CLIENT + assert created["retired_identities"]["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 + assert replacement["retired_identities"]["create"]["client_id"] == HUMAN def test_unchanged_binding_preserves_revision_and_authentication_evidence() -> None: