Merge pull request #17123 from BerriAI/litellm_ui_model_page_perf

[Infra] UI - Migrate modelInfoCall to ReactQuery
This commit is contained in:
yuneng-jiang
2025-11-26 10:12:09 -08:00
committed by GitHub
8 changed files with 188 additions and 60 deletions
@@ -0,0 +1,35 @@
// Query keys factory
type ListParams = {
page?: number;
limit?: number;
filters?: Record<string, string | number>;
};
/**
* Generates a query keys factory for a given resource.
*
* @param resource - The name of the resource (e.g., "books", "users", "keys")
* @returns An object with query key generators following the standard pattern
*
* @example
* ```ts
* const bookKeys = createQueryKeys("books");
* // bookKeys.all -> ["books"]
* // bookKeys.lists() -> ["books", "list"]
* // bookKeys.list({ page: 1 }) -> ["books", "list", { params: { page: 1 } }]
* // bookKeys.details() -> ["books", "detail"]
* // bookKeys.detail("123") -> ["books", "detail", "123"]
* ```
*/
export function createQueryKeys<T extends string>(resource: T) {
const all = [resource] as const;
return {
all,
lists: () => [...all, "list"] as const,
list: (params?: ListParams) => [...all, "list", { params }] as const,
details: () => [...all, "detail"] as const,
detail: (uid: string) => [...all, "detail", uid] as const,
};
}
@@ -0,0 +1,27 @@
import { useQuery } from "@tanstack/react-query";
import { createQueryKeys } from "../common/queryKeysFactory";
import { modelInfoCall, modelHubCall } from "@/components/networking";
const modelKeys = createQueryKeys("models");
const modelHubKeys = createQueryKeys("modelHub");
export const useModelsInfo = (accessToken: string | null, userID: string | null, userRole: string | null) => {
return useQuery({
queryKey: modelKeys.list({
filters: {
...(userID && { userID }),
...(userRole && { userRole }),
},
}),
queryFn: async () => await modelInfoCall(accessToken!, userID!, userRole!),
enabled: Boolean(accessToken && userID && userRole),
});
};
export const useModelHub = (accessToken: string | null) => {
return useQuery({
queryKey: modelHubKeys.list({}),
queryFn: async () => await modelHubCall(accessToken!),
enabled: Boolean(accessToken),
});
};
@@ -67,7 +67,7 @@ const useAuthorized = () => {
userRole: formatUserRole(decoded?.user_role ?? null),
premiumUser: decoded?.premium_user ?? null,
disabledPersonalKeyCreation: decoded?.disabled_non_admin_personal_key_creation ?? null,
showSSOBanner: decoded?.login_method === "username_password" ?? false,
showSSOBanner: decoded?.login_method === "username_password",
};
};
@@ -1,6 +1,7 @@
/* @vitest-environment jsdom */
import { render } from "@testing-library/react";
import { describe, it, expect, vi } from "vitest";
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import ModelsAndEndpointsView from "./ModelsAndEndpointsView";
// Minimal stubs to avoid Next.js router and network usage during render
@@ -56,19 +57,24 @@ vi.mock("@/app/(dashboard)/hooks/useTeams", () => ({
}),
}));
const createQueryClient = () =>
new QueryClient({
defaultOptions: { queries: { retry: false, gcTime: 0 } },
});
describe("ModelsAndEndpointsView", () => {
it(
"should render the models and endpoints view",
async () => {
// JSDOM polyfill for libraries expecting ResizeObserver (e.g., recharts)
// Note: ResizeObserver is now globally mocked in setupTests.ts, but keeping this for backwards compatibility
// eslint-disable-next-line @typescript-eslint/no-explicit-any
(global as any).ResizeObserver = class {
observe() {}
unobserve() {}
disconnect() {}
};
const { findByText } = render(
it("should render the models and endpoints view", async () => {
// JSDOM polyfill for libraries expecting ResizeObserver (e.g., recharts)
// Note: ResizeObserver is now globally mocked in setupTests.ts, but keeping this for backwards compatibility
// eslint-disable-next-line @typescript-eslint/no-explicit-any
(global as any).ResizeObserver = class {
observe() {}
unobserve() {}
disconnect() {}
};
const queryClient = createQueryClient();
const { findByText } = render(
<QueryClientProvider client={queryClient}>
<ModelsAndEndpointsView
accessToken="123"
token="123"
@@ -79,10 +85,9 @@ describe("ModelsAndEndpointsView", () => {
setModelData={() => {}}
premiumUser={false}
teams={[]}
/>,
);
expect(await findByText("Model Management", {}, { timeout: 10000 })).toBeInTheDocument();
},
15000,
);
/>
</QueryClientProvider>,
);
expect(await findByText("Model Management", {}, { timeout: 10000 })).toBeInTheDocument();
}, 15000);
});
@@ -1,5 +1,6 @@
import React, { useState, useEffect, useRef } from "react";
import { Text, Grid, Col } from "@tremor/react";
import { useQueryClient } from "@tanstack/react-query";
import { CredentialItem, credentialListCall, CredentialsResponse } from "@/components/networking";
import { handleAddModelSubmit } from "@/components/add_model/handle_add_model_submit";
@@ -9,7 +10,6 @@ import { getDisplayModelName } from "@/components/view_model/model_name_display"
import { TabPanel, TabPanels, TabGroup, TabList, Tab, Icon } from "@tremor/react";
import { DateRangePickerValue } from "@tremor/react";
import {
modelInfoCall,
modelCostMap,
modelMetricsCall,
streamingModelMetricsCall,
@@ -22,6 +22,7 @@ import {
adminGlobalActivityExceptionsPerDeployment,
allEndUsersCall,
} from "@/components/networking";
import { useModelsInfo } from "@/app/(dashboard)/hooks/models/useModels";
import { Form } from "antd";
import { Typography } from "antd";
import { RefreshIcon } from "@heroicons/react/outline";
@@ -152,6 +153,14 @@ const ModelsAndEndpointsView: React.FC<ModelDashboardProps> = ({
const dropdownRef = useRef<HTMLDivElement>(null);
const [selectedTabIndex, setSelectedTabIndex] = useState(0);
const queryClient = useQueryClient();
const {
data: modelDataResponse,
isLoading: isLoadingModels,
refetch: refetchModels,
} = useModelsInfo(accessToken, userID, userRole);
const setProviderModelsFn = (provider: Providers) => {
const _providerModels = getProviderModels(provider, modelMap);
setProviderModels(_providerModels);
@@ -180,6 +189,7 @@ const ModelsAndEndpointsView: React.FC<ModelDashboardProps> = ({
const uploadProps: UploadProps = {
name: "file",
accept: ".json",
pastable: false,
beforeUpload: (file) => {
if (file.type === "application/json") {
const reader = new FileReader();
@@ -207,6 +217,9 @@ const ModelsAndEndpointsView: React.FC<ModelDashboardProps> = ({
// Update the 'lastRefreshed' state to the current date and time
const currentDate = new Date();
setLastRefreshed(currentDate.toLocaleString());
// Invalidate and refetch models data using React Query
queryClient.invalidateQueries({ queryKey: ["models", "list"] });
refetchModels();
};
const handleSaveRetrySettings = async () => {
@@ -239,13 +252,11 @@ const ModelsAndEndpointsView: React.FC<ModelDashboardProps> = ({
};
useEffect(() => {
if (!accessToken || !token || !userRole || !userID) {
if (!accessToken || !token || !userRole || !userID || !modelDataResponse) {
return;
}
const fetchData = async () => {
try {
// Replace with your actual API call for model data
const modelDataResponse = await modelInfoCall(accessToken, userID, userRole);
setModelData(modelDataResponse);
const _providerSettings = await modelSettingsCall(accessToken);
if (_providerSettings) {
@@ -372,7 +383,7 @@ const ModelsAndEndpointsView: React.FC<ModelDashboardProps> = ({
}
};
if (accessToken && token && userRole && userID) {
if (accessToken && token && userRole && userID && modelDataResponse) {
fetchData();
}
@@ -383,11 +394,9 @@ const ModelsAndEndpointsView: React.FC<ModelDashboardProps> = ({
if (modelMap == null) {
fetchModelMap();
}
}, [accessToken, token, userRole, userID, modelDataResponse]);
handleRefreshClick();
}, [accessToken, token, userRole, userID, modelMap, lastRefreshed, selectedTeam]);
if (!modelData) {
if (!modelData || isLoadingModels) {
return <div>Loading...</div>;
}
@@ -597,15 +606,25 @@ const ModelsAndEndpointsView: React.FC<ModelDashboardProps> = ({
setEditModalVisible={setEditModalVisible}
setSelectedModel={setSelectedModel}
onModelUpdate={(updatedModel) => {
// Update the model in the modelData.data array
const updatedModelData = {
...modelData,
data: modelData.data.map((model: any) =>
model.model_info.id === updatedModel.model_info.id ? updatedModel : model,
),
};
setModelData(updatedModelData);
// Trigger a refresh to update UI
// Handle model deletion
if (updatedModel.deleted) {
const updatedModelData = {
...modelData,
data: modelData.data.filter((model: any) => model.model_info.id !== updatedModel.model_info.id),
};
setModelData(updatedModelData);
} else {
// Update the model in the modelData.data array
const updatedModelData = {
...modelData,
data: modelData.data.map((model: any) =>
model.model_info.id === updatedModel.model_info.id ? updatedModel : model,
),
};
setModelData(updatedModelData);
}
// Invalidate cache and trigger a refresh to update UI
queryClient.invalidateQueries({ queryKey: ["models", "list"] });
handleRefreshClick();
}}
modelAccessGroups={availableModelAccessGroups}
@@ -0,0 +1,39 @@
import React from "react";
import { describe, it, expect, vi, beforeEach } from "vitest";
import { screen, render, fireEvent, act } from "@testing-library/react";
import { ProviderLogo } from "./ProviderLogo";
import * as providerInfoHelpers from "../../provider_info_helpers";
vi.mock("../../provider_info_helpers");
describe("ProviderLogo", () => {
beforeEach(() => {
vi.resetAllMocks();
});
it("should render the component", () => {
vi.mocked(providerInfoHelpers.getProviderLogoAndName).mockReturnValue({
logo: "",
displayName: "Test Provider",
});
render(<ProviderLogo provider="test" />);
expect(screen.getByText("t")).toBeInTheDocument();
});
it("should show fallback when image fails to load", () => {
vi.mocked(providerInfoHelpers.getProviderLogoAndName).mockReturnValue({
logo: "/path/to/logo.png",
displayName: "Test Provider",
});
render(<ProviderLogo provider="test" />);
const img = screen.getByRole("img", { name: "test logo" });
expect(img).toBeInTheDocument();
act(() => {
fireEvent.error(img);
});
expect(screen.getByText("t")).toBeInTheDocument();
expect(screen.queryByRole("img")).not.toBeInTheDocument();
});
});
@@ -0,0 +1,24 @@
import React, { useState } from "react";
import { getProviderLogoAndName } from "../../provider_info_helpers";
interface ProviderLogoProps {
provider: string;
className?: string;
}
export const ProviderLogo: React.FC<ProviderLogoProps> = ({ provider, className = "w-4 h-4" }) => {
const [hasError, setHasError] = useState(false);
const { logo } = getProviderLogoAndName(provider);
const showFallback = hasError || !logo;
if (showFallback) {
return (
<div className={`${className} rounded-full bg-gray-200 flex items-center justify-center text-xs`}>
{provider?.charAt(0) || "-"}
</div>
);
}
return <img src={logo} alt={`${provider} logo`} className={className} onError={() => setHasError(true)} />;
};
@@ -3,7 +3,7 @@ import { ColumnDef } from "@tanstack/react-table";
import { Badge, Button, Icon } from "@tremor/react";
import { Tooltip } from "antd";
import { ModelData } from "../../model_dashboard/types";
import { getProviderLogoAndName } from "../../provider_info_helpers";
import { ProviderLogo } from "./ProviderLogo";
export const columns = (
userRole: string,
@@ -62,28 +62,7 @@ export const columns = (
{/* Provider Icon */}
<div className="flex-shrink-0 mt-0.5">
{model.provider ? (
<img
src={getProviderLogoAndName(model.provider).logo}
alt={`${model.provider} logo`}
className="w-4 h-4"
onError={(e) => {
const target = e.currentTarget as HTMLImageElement;
const parent = target.parentElement;
if (!parent || !parent.contains(target)) {
return;
}
try {
const fallbackDiv = document.createElement("div");
fallbackDiv.className =
"w-4 h-4 rounded-full bg-gray-200 flex items-center justify-center text-xs";
fallbackDiv.textContent = model.provider?.charAt(0) || "-";
parent.replaceChild(fallbackDiv, target);
} catch (error) {
console.error("Failed to replace provider logo fallback:", error);
}
}}
/>
<ProviderLogo provider={model.provider} />
) : (
<div className="w-4 h-4 rounded-full bg-gray-200 flex items-center justify-center text-xs">-</div>
)}