mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(agents): preserve retired identity ownership
This commit is contained in:
parent
f7b84953b5
commit
a4492b9b3c
4 changed files with 129 additions and 33 deletions
|
|
@ -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},
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue