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