fix(agents): preserve retired identity ownership

This commit is contained in:
Joshua Valluru 2026-09-28 19:26:41 -07:00
parent f7b84953b5
commit a4492b9b3c
4 changed files with 129 additions and 33 deletions

View file

@ -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},

View file

@ -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,

View file

@ -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",
}
}

View file

@ -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: