mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
fix(ui): keep MCP permissions visible after key, team and MCP server saves (#43810)
* fix(ui): keep MCP permissions visible after key, team and MCP server saves Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): type the object_permission include as a prisma TypedDict Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(ui): do not block key save confirmation on cache refetch Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: ryan <ryan@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
7d9cc28dce
commit
b71f02dbcf
13 changed files with 133 additions and 17 deletions
|
|
@ -2442,11 +2442,13 @@ async def _update_key_row_with_soft_budget(
|
|||
existing_key_row=existing_key_row,
|
||||
changed_by=changed_by,
|
||||
)
|
||||
include_object_permission: Final[prisma.types.LiteLLM_VerificationTokenInclude] = {"object_permission": True}
|
||||
updated_row: Final = await tx.litellm_verificationtoken.update(
|
||||
where=key_where,
|
||||
data=with_settings_updated_at(
|
||||
prisma_client.jsonify_object(MappingProxyType({**update_values, "token": hashed_token}))
|
||||
),
|
||||
include=include_object_permission,
|
||||
)
|
||||
updated_data: Final[Mapping[str, object]] = (
|
||||
updated_row.model_dump() if updated_row is not None else MappingProxyType({})
|
||||
|
|
|
|||
|
|
@ -258,7 +258,7 @@ if TYPE_CHECKING:
|
|||
from prisma.actions import LiteLLM_DeprecatedVerificationTokenActions
|
||||
from prisma.client import TransactionManager
|
||||
from prisma.models import LiteLLM_DeprecatedVerificationToken
|
||||
from prisma.types import HttpConfig
|
||||
from prisma.types import HttpConfig, LiteLLM_VerificationTokenInclude
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
|
||||
|
|
@ -5439,9 +5439,11 @@ class PrismaClient:
|
|||
# check if plain text or hash
|
||||
token = _hash_token_if_needed(token=token)
|
||||
db_data["token"] = token
|
||||
include_object_permission: Final[LiteLLM_VerificationTokenInclude] = {"object_permission": True}
|
||||
response: Final = await VerificationTokenRepository(self).table.update(
|
||||
where={"token": token},
|
||||
data=with_settings_updated_at(db_data),
|
||||
include=include_object_permission,
|
||||
)
|
||||
verbose_proxy_logger.debug("\033[91m" + f"DB Token Table update succeeded {response}" + "\033[0m")
|
||||
_data: dict = {}
|
||||
|
|
|
|||
|
|
@ -20291,7 +20291,11 @@ async def test_update_key_row_with_soft_budget_updates_budget_and_key_in_transac
|
|||
existing_key = LiteLLM_VerificationToken(token="test-token", budget_id=None)
|
||||
created_row = MagicMock(budget_id="budget-new")
|
||||
updated_row = MagicMock()
|
||||
updated_row.model_dump.return_value = {"token": "hashed", "budget_id": "budget-new"}
|
||||
updated_row.model_dump.return_value = {
|
||||
"token": "hashed",
|
||||
"budget_id": "budget-new",
|
||||
"object_permission": {"mcp_servers": ["srv-1"], "mcp_tool_permissions": {"srv-1": ["read"]}},
|
||||
}
|
||||
tx = MagicMock()
|
||||
tx.litellm_budgettable.create = AsyncMock(return_value=created_row)
|
||||
tx.litellm_verificationtoken.update = AsyncMock(return_value=updated_row)
|
||||
|
|
@ -20312,10 +20316,15 @@ async def test_update_key_row_with_soft_budget_updates_budget_and_key_in_transac
|
|||
)
|
||||
|
||||
assert set(result) == {"token", "data"}
|
||||
assert result["data"] == {"token": "hashed", "budget_id": "budget-new"}
|
||||
assert result["data"] == {
|
||||
"token": "hashed",
|
||||
"budget_id": "budget-new",
|
||||
"object_permission": {"mcp_servers": ["srv-1"], "mcp_tool_permissions": {"srv-1": ["read"]}},
|
||||
}
|
||||
tx.litellm_verificationtoken.update.assert_awaited_once()
|
||||
update_call = tx.litellm_verificationtoken.update.await_args
|
||||
assert update_call.kwargs["where"] == {"token": result["token"]}
|
||||
assert update_call.kwargs["include"] == {"object_permission": True}
|
||||
assert update_call.kwargs["data"]["budget_id"] == "budget-new"
|
||||
assert "soft_budget" not in update_call.kwargs["data"]
|
||||
|
||||
|
|
|
|||
|
|
@ -155,6 +155,7 @@ async def test_update_data_token_hashes_and_updates(
|
|||
"token": hashlib.sha256(token.encode()).hexdigest(),
|
||||
"spend": 1.0,
|
||||
"user_id": "u1",
|
||||
"object_permission": {"mcp_servers": ["srv-1"]},
|
||||
},
|
||||
)
|
||||
prisma_client.db.litellm_verificationtoken.update = AsyncMock(return_value=response)
|
||||
|
|
@ -167,15 +168,22 @@ async def test_update_data_token_hashes_and_updates(
|
|||
actual = {
|
||||
"result": result,
|
||||
"where": update_kwargs["where"],
|
||||
"include": update_kwargs["include"],
|
||||
"data_token": update_kwargs["data"]["token"],
|
||||
"data_spend": update_kwargs["data"]["spend"],
|
||||
}
|
||||
assert actual == {
|
||||
"result": {
|
||||
"token": hashed,
|
||||
"data": {"token": hashed, "spend": 1.0, "user_id": "u1"},
|
||||
"data": {
|
||||
"token": hashed,
|
||||
"spend": 1.0,
|
||||
"user_id": "u1",
|
||||
"object_permission": {"mcp_servers": ["srv-1"]},
|
||||
},
|
||||
},
|
||||
"where": {"token": hashed},
|
||||
"include": {"object_permission": True},
|
||||
"data_token": hashed,
|
||||
"data_spend": 1.0,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ import { fetchMCPServers } from "@/components/networking";
|
|||
import { MCPServer } from "@/components/mcp_tools/types";
|
||||
import useAuthorized from "../useAuthorized";
|
||||
|
||||
const mcpServersKeys = createQueryKeys("mcpServers");
|
||||
export const mcpServersKeys = createQueryKeys("mcpServers");
|
||||
|
||||
export const useMCPServers = (teamId?: string | null) => {
|
||||
const { accessToken } = useAuthorized();
|
||||
|
|
|
|||
|
|
@ -1,4 +1,11 @@
|
|||
import { keepPreviousData, useInfiniteQuery, useQuery, useQueryClient, UseQueryResult } from "@tanstack/react-query";
|
||||
import {
|
||||
keepPreviousData,
|
||||
QueryClient,
|
||||
useInfiniteQuery,
|
||||
useQuery,
|
||||
useQueryClient,
|
||||
UseQueryResult,
|
||||
} from "@tanstack/react-query";
|
||||
import { Team } from "@/components/key_team_helpers/key_list";
|
||||
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
||||
import { fetchTeams } from "@/app/(dashboard)/networking";
|
||||
|
|
@ -110,7 +117,7 @@ export const useTeamsTable = (
|
|||
});
|
||||
};
|
||||
|
||||
const teamKeys = createQueryKeys("teams");
|
||||
export const teamKeys = createQueryKeys("teams");
|
||||
export const useTeams = (): UseQueryResult<Team[]> => {
|
||||
const { accessToken, userId, userRole } = useAuthorized();
|
||||
return useQuery<Team[]>({
|
||||
|
|
@ -179,6 +186,11 @@ export const useTeam = (teamId?: string) => {
|
|||
|
||||
const infiniteTeamKeys = createQueryKeys("infiniteTeams");
|
||||
|
||||
export const invalidateTeamQueries = (queryClient: QueryClient) =>
|
||||
Promise.all(
|
||||
[teamsTableKeys, teamKeys, infiniteTeamKeys].map((keys) => queryClient.invalidateQueries({ queryKey: keys.all })),
|
||||
);
|
||||
|
||||
export const useInfiniteTeams = (pageSize: number = 50, search?: string, organizationId?: string | null) => {
|
||||
const { accessToken, userId, userRole } = useAuthorized();
|
||||
const isAdmin = userRole === "Admin" || userRole === "Admin Viewer";
|
||||
|
|
|
|||
|
|
@ -7,13 +7,21 @@ import * as networking from "@/components/networking";
|
|||
import { setSecureItem } from "@/utils/secureStorage";
|
||||
import { EDIT_OAUTH_UI_STATE_KEY } from "./mcp_server_edit";
|
||||
import type { MCPServer } from "@/components/mcp_tools/types";
|
||||
import { mcpServersKeys } from "@/app/(dashboard)/hooks/mcpServers/useMCPServers";
|
||||
|
||||
vi.mock(".", () => ({
|
||||
MCPToolsViewer: () => <div>tools viewer</div>,
|
||||
}));
|
||||
|
||||
vi.mock("./mcp_server_edit", () => ({
|
||||
default: () => <div>edit form</div>,
|
||||
default: ({ mcpServer, onSuccess }: { mcpServer: MCPServer; onSuccess: (server: MCPServer) => void }) => (
|
||||
<div>
|
||||
edit form
|
||||
<button type="button" onClick={() => onSuccess({ ...mcpServer, alias: "renamed" })}>
|
||||
save edit
|
||||
</button>
|
||||
</div>
|
||||
),
|
||||
EDIT_OAUTH_UI_STATE_KEY: "litellm-mcp-oauth-edit-state",
|
||||
}));
|
||||
|
||||
|
|
@ -33,9 +41,15 @@ const baseServer = {
|
|||
auth_type: "api_key",
|
||||
} as MCPServer;
|
||||
|
||||
const renderView = (overrides: Partial<MCPServer> = {}, props: Record<string, unknown> = {}) =>
|
||||
const newQueryClient = () => new QueryClient({ defaultOptions: { queries: { retry: false, gcTime: 0 } } });
|
||||
|
||||
const renderView = (
|
||||
overrides: Partial<MCPServer> = {},
|
||||
props: Record<string, unknown> = {},
|
||||
queryClient: QueryClient = newQueryClient(),
|
||||
) =>
|
||||
render(
|
||||
<QueryClientProvider client={new QueryClient({ defaultOptions: { queries: { retry: false, gcTime: 0 } } })}>
|
||||
<QueryClientProvider client={queryClient}>
|
||||
<MCPServerView
|
||||
mcpServer={{ ...baseServer, ...overrides } as MCPServer}
|
||||
onBack={vi.fn()}
|
||||
|
|
@ -147,6 +161,27 @@ describe("MCPServerView", () => {
|
|||
expect(await screen.findByText("edit form")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("drops the cached server list and tool catalog once the edit form saves", async () => {
|
||||
const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false, gcTime: Infinity } } });
|
||||
const serversKey = mcpServersKeys.list();
|
||||
const toolsKey = ["mcpTools", "srv-1", {}, null];
|
||||
const otherToolsKey = ["mcpTools", "srv-2", {}, null];
|
||||
queryClient.setQueryData(serversKey, [baseServer]);
|
||||
queryClient.setQueryData(toolsKey, { tools: [] });
|
||||
queryClient.setQueryData(otherToolsKey, { tools: [] });
|
||||
const onBack = vi.fn();
|
||||
renderView({}, { onBack }, queryClient);
|
||||
|
||||
await userEvent.click(screen.getByRole("tab", { name: "Settings" }));
|
||||
await userEvent.click(await screen.findByRole("button", { name: "Edit Settings" }));
|
||||
await userEvent.click(await screen.findByRole("button", { name: "save edit" }));
|
||||
|
||||
expect(queryClient.getQueryState(serversKey)?.isInvalidated).toBe(true);
|
||||
expect(queryClient.getQueryState(toolsKey)?.isInvalidated).toBe(true);
|
||||
expect(queryClient.getQueryState(otherToolsKey)?.isInvalidated).toBe(false);
|
||||
expect(onBack).toHaveBeenCalledTimes(1);
|
||||
});
|
||||
|
||||
it("opens straight into the edit form when isEditing is set", async () => {
|
||||
renderView({}, { isEditing: true });
|
||||
|
||||
|
|
|
|||
|
|
@ -1,4 +1,6 @@
|
|||
import React, { useState } from "react";
|
||||
import { useQueryClient } from "@tanstack/react-query";
|
||||
import { mcpServersKeys } from "@/app/(dashboard)/hooks/mcpServers/useMCPServers";
|
||||
import { ArrowLeft, Eye, EyeOff } from "lucide-react";
|
||||
import { Badge } from "@/components/ui/badge";
|
||||
import { Button } from "@/components/ui/button";
|
||||
|
|
@ -64,6 +66,7 @@ export const MCPServerView: React.FC<MCPServerViewProps> = ({
|
|||
}) => {
|
||||
// Open the editing Settings tab on first render when returning from the edit OAuth
|
||||
// redirect, so the "token fetched" feedback shows where the user left off (Settings=2).
|
||||
const queryClient = useQueryClient();
|
||||
const canEdit = isProxyAdmin && !isViewOnly && !mcpServer.is_config;
|
||||
const returningFromEditOAuth = isReturningFromEditOAuth(canEdit, mcpServer.server_id);
|
||||
const [editing, setEditing] = useState(isEditing || returningFromEditOAuth);
|
||||
|
|
@ -75,6 +78,8 @@ export const MCPServerView: React.FC<MCPServerViewProps> = ({
|
|||
const canRevokeUserCredentials = userRole !== null && isProxyAdminRole(userRole) && !isViewOnly;
|
||||
|
||||
const handleSuccess = (updated: MCPServer) => {
|
||||
void queryClient.invalidateQueries({ queryKey: mcpServersKeys.all });
|
||||
void queryClient.invalidateQueries({ queryKey: ["mcpTools", updated.server_id] });
|
||||
setEditing(false);
|
||||
onBack();
|
||||
};
|
||||
|
|
|
|||
|
|
@ -77,7 +77,8 @@ vi.mock("@/app/(dashboard)/hooks/models/useModels", () => ({
|
|||
useAllProxyModels: vi.fn(),
|
||||
}));
|
||||
|
||||
vi.mock("@/app/(dashboard)/hooks/teams/useTeams", () => ({
|
||||
vi.mock("@/app/(dashboard)/hooks/teams/useTeams", async (importOriginal) => ({
|
||||
...(await importOriginal<typeof import("@/app/(dashboard)/hooks/teams/useTeams")>()),
|
||||
useTeam: vi.fn(),
|
||||
}));
|
||||
|
||||
|
|
|
|||
|
|
@ -83,7 +83,8 @@ vi.mock("@/app/(dashboard)/hooks/models/useModels", () => ({
|
|||
useAllProxyModels: vi.fn(),
|
||||
}));
|
||||
|
||||
vi.mock("@/app/(dashboard)/hooks/teams/useTeams", () => ({
|
||||
vi.mock("@/app/(dashboard)/hooks/teams/useTeams", async (importOriginal) => ({
|
||||
...(await importOriginal<typeof import("@/app/(dashboard)/hooks/teams/useTeams")>()),
|
||||
useTeam: vi.fn(),
|
||||
}));
|
||||
|
||||
|
|
@ -233,7 +234,7 @@ vi.mock("../key_team_helpers/filter_helpers", () => ({
|
|||
import { useAllProxyModels } from "@/app/(dashboard)/hooks/models/useModels";
|
||||
import { useKeys } from "@/app/(dashboard)/hooks/keys/useKeys";
|
||||
import { useOrganization } from "@/app/(dashboard)/hooks/organizations/useOrganizations";
|
||||
import { useTeam } from "@/app/(dashboard)/hooks/teams/useTeams";
|
||||
import { teamKeys, teamsTableKeys, useTeam } from "@/app/(dashboard)/hooks/teams/useTeams";
|
||||
import { useCurrentUser } from "@/app/(dashboard)/hooks/users/useCurrentUser";
|
||||
import { useMCPServers } from "@/app/(dashboard)/hooks/mcpServers/useMCPServers";
|
||||
import { useMCPToolsets } from "@/app/(dashboard)/hooks/mcpServers/useMCPToolsets";
|
||||
|
|
@ -1146,6 +1147,26 @@ describe("TeamInfoView", () => {
|
|||
});
|
||||
});
|
||||
|
||||
it("invalidates the cached team list and team detail queries after saving team settings", async () => {
|
||||
const user = userEvent.setup({ delay: null });
|
||||
vi.mocked(networking.teamInfoCall).mockResolvedValue(createMockTeamData({ models: ["gpt-4"] }));
|
||||
vi.mocked(networking.teamUpdateCall).mockResolvedValue({ data: {}, team_id: "123" } as any);
|
||||
const tableKey = teamsTableKeys.list({ page: 1, limit: 10 });
|
||||
const detailKey = teamKeys.detail("123");
|
||||
testQueryClient.setQueryData(tableKey, { teams: [], total: 0 });
|
||||
testQueryClient.setQueryData(detailKey, { team_id: "123" });
|
||||
|
||||
renderWithProviders(<TeamInfoView {...defaultProps} />);
|
||||
|
||||
await user.click(await screen.findByRole("tab", { name: "Settings" }));
|
||||
await user.click(await screen.findByRole("button", { name: /edit settings/i }));
|
||||
await screen.findByLabelText("Team Name");
|
||||
await user.click(screen.getByRole("button", { name: /save changes/i }));
|
||||
|
||||
await waitFor(() => expect(testQueryClient.getQueryState(tableKey)?.isInvalidated).toBe(true));
|
||||
expect(testQueryClient.getQueryState(detailKey)?.isInvalidated).toBe(true);
|
||||
});
|
||||
|
||||
const openSettingsEditorForTeam = async (
|
||||
user: ReturnType<typeof userEvent.setup>,
|
||||
teamOverrides: Record<string, unknown>,
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
|||
import type { components } from "@/lib/http/schema";
|
||||
import useCan from "@/app/(dashboard)/hooks/useCan";
|
||||
import { organizationKeys, useOrganizations } from "@/app/(dashboard)/hooks/organizations/useOrganizations";
|
||||
import { invalidateTeamQueries } from "@/app/(dashboard)/hooks/teams/useTeams";
|
||||
import { useQueryClient } from "@tanstack/react-query";
|
||||
import UserSearchModal from "@/components/common_components/user_search_modal";
|
||||
import {
|
||||
|
|
@ -915,7 +916,8 @@ const TeamInfoView: React.FC<TeamInfoProps> = ({
|
|||
|
||||
const persistTeamUpdate = async (token: string, updateData: Record<string, unknown>) => {
|
||||
await teamUpdateCall(token, updateData);
|
||||
queryClient.invalidateQueries({ queryKey: organizationKeys.all });
|
||||
void queryClient.invalidateQueries({ queryKey: organizationKeys.all });
|
||||
void invalidateTeamQueries(queryClient);
|
||||
setIsEditing(false);
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -7,10 +7,11 @@ vi.mock("@/lib/toast", () => ({
|
|||
}));
|
||||
|
||||
// ---- Hoisted shared mocks (safe to use inside vi.mock factories) ----
|
||||
const { keyUpdateCallMock, keyDeleteCallMock, mockUseAuthorized } = vi.hoisted(() => {
|
||||
const { keyUpdateCallMock, keyDeleteCallMock, invalidateQueriesMock, mockUseAuthorized } = vi.hoisted(() => {
|
||||
return {
|
||||
keyUpdateCallMock: vi.fn().mockResolvedValue({}),
|
||||
keyDeleteCallMock: vi.fn().mockResolvedValue({}),
|
||||
invalidateQueriesMock: vi.fn().mockResolvedValue(undefined),
|
||||
mockUseAuthorized: vi.fn(),
|
||||
};
|
||||
});
|
||||
|
|
@ -170,7 +171,7 @@ vi.mock("@tanstack/react-query", async (importOriginal) => {
|
|||
const actual = await importOriginal<typeof import("@tanstack/react-query")>();
|
||||
return {
|
||||
...actual,
|
||||
useQueryClient: () => ({ invalidateQueries: vi.fn() }),
|
||||
useQueryClient: () => ({ invalidateQueries: invalidateQueriesMock }),
|
||||
};
|
||||
});
|
||||
|
||||
|
|
@ -366,6 +367,24 @@ describe("KeyInfoView handleKeyUpdate mcp_toolsets", () => {
|
|||
});
|
||||
});
|
||||
|
||||
describe("KeyInfoView handleKeyUpdate cache sync", () => {
|
||||
it("should invalidate every cached key query so the list and detail views re-read the saved key", async () => {
|
||||
keyUpdateCallMock.mockResolvedValueOnce({
|
||||
object_permission: { mcp_servers: ["srv-1"], mcp_tool_permissions: { "srv-1": ["read_wiki"] } },
|
||||
});
|
||||
renderView(true);
|
||||
|
||||
fireEvent.click(screen.getByText("Settings"));
|
||||
fireEvent.click(screen.getByText("Edit Settings"));
|
||||
(globalThis as any).__TEST_FORM_VALUES = { token: "tok_123", metadata: {} };
|
||||
|
||||
fireEvent.click(screen.getByText("Mock Submit"));
|
||||
|
||||
await waitFor(() => expect(toast.success).toHaveBeenCalledWith("Key updated successfully"));
|
||||
expect(invalidateQueriesMock).toHaveBeenCalledWith({ queryKey: ["keys"] });
|
||||
});
|
||||
});
|
||||
|
||||
describe("KeyInfoView handleKeyUpdate skills", () => {
|
||||
it("should forward the skills the edit form supplies into object_permission and drop the form key", async () => {
|
||||
renderView(true);
|
||||
|
|
|
|||
|
|
@ -381,8 +381,8 @@ export default function KeyInfoView({
|
|||
|
||||
const newKeyValues = await keyUpdateCall(accessToken, formValues);
|
||||
|
||||
// Update local state
|
||||
setCurrentKeyData((prevData) => (prevData ? { ...prevData, ...newKeyValues } : undefined));
|
||||
void queryClient.invalidateQueries({ queryKey: keyKeys.all });
|
||||
|
||||
if (onKeyDataUpdate) {
|
||||
onKeyDataUpdate(newKeyValues);
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue