Migrate some queries to use react query, tests pending

This commit is contained in:
yuneng-jiang
2025-11-25 21:00:41 -08:00
parent 577f40bc60
commit cb1809987d
6 changed files with 131 additions and 42 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,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}
@@ -1,9 +1,8 @@
import { KeyIcon, TrashIcon } from "@heroicons/react/outline";
import { ColumnDef } from "@tanstack/react-table";
import { Button, Badge, Icon } from "@tremor/react";
import { Badge, Button, Icon } from "@tremor/react";
import { Tooltip } from "antd";
import { getProviderLogoAndName } from "../../provider_info_helpers";
import { ModelData } from "../../model_dashboard/types";
import { TrashIcon, KeyIcon } from "@heroicons/react/outline";
import { ProviderLogo } from "./ProviderLogo";
export const columns = (
@@ -1,6 +1,7 @@
// fetch_models.ts
import { modelHubCall } from "../../networking";
import { useModelHub } from "@/app/(dashboard)/hooks/models/useModels";
import { useMemo } from "react";
export interface ModelGroup {
model_group: string;
@@ -8,26 +9,34 @@ export interface ModelGroup {
}
/**
* Fetches available models using modelHubCall and formats them for the selection dropdown.
* Hook that fetches available models using modelHubCall and formats them for the selection dropdown.
*/
export const fetchAvailableModels = async (accessToken: string): Promise<ModelGroup[]> => {
try {
const fetchedModels = await modelHubCall(accessToken);
console.log("model_info:", fetchedModels);
export const useAvailableModels = (
accessToken: string | null,
): {
models: ModelGroup[];
isLoading: boolean;
error: Error | null;
} => {
const { data: fetchedModels, isLoading, error } = useModelHub(accessToken);
if (fetchedModels?.data.length > 0) {
const models: ModelGroup[] = fetchedModels.data.map((item: any) => ({
model_group: item.model_group, // Display the model_group to the user
mode: item?.mode, // Save the mode for auto-selection of endpoint type
}));
// Sort models alphabetically by label
models.sort((a, b) => a.model_group.localeCompare(b.model_group));
return models;
const models = useMemo(() => {
if (!fetchedModels?.data || fetchedModels.data.length === 0) {
return [];
}
return [];
} catch (error) {
console.error("Error fetching model info:", error);
throw error;
}
const formattedModels: ModelGroup[] = fetchedModels.data.map((item: any) => ({
model_group: item.model_group,
mode: item?.mode,
}));
formattedModels.sort((a, b) => a.model_group.localeCompare(b.model_group));
return formattedModels;
}, [fetchedModels]);
return {
models,
isLoading,
error: error as Error | null,
};
};