fix(agents): preserve configuration during identity updates

This commit is contained in:
Joshua Valluru 2026-09-30 13:43:36 -07:00
parent a4492b9b3c
commit 1cdc43665c
7 changed files with 99 additions and 4 deletions

View file

@ -793,6 +793,13 @@ class AgentRegistry:
"""
Update an agent in the database
"""
if "agent_card_params" not in agent:
return await self.patch_agent_in_db(
agent_id=agent_id,
agent=PatchAgentRequest(**agent),
prisma_client=prisma_client,
updated_by=updated_by,
)
try:
agent_name: Final = agent.get("agent_name")

View file

@ -1236,7 +1236,7 @@ def _stored_agent_row(values: Mapping[str, object] | SimpleNamespace) -> LiteLLM
"execution_mode": "autonomous",
**{
key: json.dumps(value)
if key in ("litellm_params", "agent_card_params", "kill_switch") and not isinstance(value, str)
if key in ("litellm_params", "agent_card_params", "kill_switch", "static_headers") and not isinstance(value, str)
else value
for key, value in fields.items()
},

View file

@ -1,4 +1,5 @@
import json
from collections.abc import Mapping
from datetime import datetime, timezone
from types import SimpleNamespace
@ -9,6 +10,7 @@ import httpx
import pytest
from fastapi import FastAPI, HTTPException
from fastapi.testclient import TestClient
from prisma.models import LiteLLM_AgentsTable
from litellm.constants import REDACTED_BY_LITELM_STRING
from litellm.proxy._types import LiteLLM_AuditLogs, LitellmTableNames, LitellmUserRoles, UserAPIKeyAuth
@ -140,6 +142,61 @@ def test_update_agent_not_found(
assert "Agent with ID missing-agent not found" in response.json()["detail"]
class _AgentPersistence:
def __init__(self, row: LiteLLM_AgentsTable) -> None:
self.row = row
async def find_unique(self, **kwargs: object) -> LiteLLM_AgentsTable:
return self.row
async def update(self, *, data: Mapping[str, object], **kwargs: object) -> LiteLLM_AgentsTable:
from tests.test_litellm.proxy.agent_endpoints.test_agent_registry import _stored_agent_row
self.row = _stored_agent_row({**self.row.model_dump(), **data})
return self.row
@pytest.mark.parametrize("method", ["PUT", "PATCH"])
@pytest.mark.parametrize("cardless", [False, True])
def test_identity_settings_edit_preserves_runtime_configuration_on_readback(
monkeypatch: pytest.MonkeyPatch, method: str, cardless: bool
) -> None:
from litellm.proxy import proxy_server
from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry
from tests.test_litellm.proxy.agent_endpoints.test_agent_registry import _stored_agent_row
runtime: Final = {
"agent_card_params": {} if cardless else _sample_agent_card_params(),
"litellm_params": {"make_public": False, "model": "a2a/runtime"},
"static_headers": {"X-Runtime": "configured"},
"extra_headers": ["X-Trace"],
"access_group_ids": ["runtime-group"],
"kill_switch": {"url": "https://runtime.example/stop", "method": "POST"},
}
row: Final = _stored_agent_row(runtime)
table: Final = _AgentPersistence(row)
database: Final = SimpleNamespace(
litellm_agentstable=table,
litellm_verificationtoken=SimpleNamespace(find_many=AsyncMock(return_value=[])),
)
monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=database, writer_db=database))
monkeypatch.setattr(agent_endpoints, "AGENT_REGISTRY", AgentRegistry())
response: Final = client.request(
method, "/v1/agents/agent-123", json={"agent_name": "Renamed agent", "enabled": False}
)
assert response.status_code == 200, response.text
readback: Final = client.get("/v1/agents/agent-123")
assert readback.status_code == 200, readback.text
stored: Final = AgentResponse.model_validate(table.row.model_dump())
expected: Final = AgentResponse.model_validate(row.model_dump()).model_copy(
update={"agent_name": "Renamed agent", "enabled": False}
)
preserved: Final = {*runtime, "agent_name", "enabled", "agent_id"}
assert stored.model_dump(include=preserved) == expected.model_dump(include=preserved)
assert {key: readback.json()[key] for key in preserved} == expected.model_dump(mode="json", include=preserved)
def test_get_agent_by_id_not_found(
mock_prisma_client, mock_user_api_key_auth, monkeypatch
):

View file

@ -33,13 +33,15 @@ export const AgentIdentityFields = ({ accessToken }: { accessToken: string | nul
apiClient
.get<string[]>("/v1/agents/identity/providers", { accessToken })
.then((issuers) => {
if (active)
if (active) {
setError(null);
setTenants(
issuers.flatMap((issuer) => {
const tenant = entraTenantFromIssuer(issuer);
return tenant ? [tenant] : [];
}),
);
}
})
.catch(() => {
if (active) setError("Could not load the gateway's trusted identity providers");

View file

@ -96,6 +96,28 @@ describe("AddAgentForm submit payload", () => {
.mockResolvedValue({} as never);
});
it("clears the provider error when reselecting Entra successfully loads trusted tenants", async () => {
const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never });
const tenant = "11111111-1111-4111-8111-111111111111";
vi.mocked(networking.apiClient.get)
.mockReset()
.mockRejectedValueOnce(new Error("temporarily unavailable"))
.mockResolvedValue([`https://login.microsoftonline.com/${tenant}/v2.0`]);
renderForm();
await user.click(await screen.findByLabelText("Identity Provider"));
await user.click(await screen.findByRole("option", { name: "Microsoft Entra ID" }));
expect(await screen.findByRole("alert")).toHaveTextContent(
"Could not load the gateway's trusted identity providers",
);
await user.click(screen.getByLabelText("Identity Provider"));
await user.click(await screen.findByRole("option", { name: "No explicit identity binding" }));
await user.click(screen.getByLabelText("Identity Provider"));
await user.click(await screen.findByRole("option", { name: "Microsoft Entra ID" }));
await user.click(screen.getByLabelText("Trusted Entra Tenant"));
expect(await screen.findByRole("option", { name: tenant })).toBeInTheDocument();
expect(screen.queryByRole("alert")).not.toBeInTheDocument();
});
it("registers a readable agent with an explicit Entra identity and no virtual key", async () => {
const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never });
const tenant = "11111111-1111-4111-8111-111111111111";

View file

@ -3,7 +3,10 @@ import type { components } from "@/lib/http/schema";
import type { AgentFormValues, AgentRequestPayload } from "./AgentFormKit";
export type EntraAgentIdentity = components["schemas"]["EntraIdentityConfig"];
type AgentIdentityState = Pick<components["schemas"]["AgentResponse"], "identity" | "enabled" | "execution_mode">;
type AgentIdentityState = Pick<
components["schemas"]["AgentResponse"],
"identity" | "enabled" | "execution_mode" | "agent_card_params"
>;
export const IDENTITY_UUID_PATTERN = /^[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}$/i;
@ -90,10 +93,13 @@ export const withAgentIdentity = (
values: AgentFormValues,
existing?: Partial<AgentIdentityState>,
): AgentRequestPayload => {
const { agent_card_params, ...settings } = payload;
const hasCard = !existing || Object.keys(existing.agent_card_params ?? {}).length > 0;
const identityFields = buildIdentityParams(values, existing?.identity);
const managed = values.identity_provider === "microsoft_entra" || Boolean(readAgentIdentity(existing?.identity));
return {
...payload,
...settings,
...(hasCard && agent_card_params ? { agent_card_params } : {}),
...identityFields,
...(managed && values.execution_mode !== undefined ? { execution_mode: values.execution_mode } : {}),
...(managed && values.enabled !== undefined ? { enabled: values.enabled } : {}),

View file

@ -191,6 +191,7 @@ describe("AgentInfoView update payload", () => {
await save(user);
expect(patchedPayload().agent_name).toBe("Renamed agent");
expect(patchedPayload()).not.toHaveProperty("litellm_params");
expect(patchedPayload().agent_card_params === undefined).toBe(card === "empty");
expect(patchedPayload().identity).toMatchObject(identity);
expect(patchedPayload().access_group_ids).toEqual(["ag-entra"]);
expect(networking.patchAgentCall).toHaveBeenCalledWith(