Refactor model installation code to support both APIs

This commit is contained in:
Arslan Ablikim
2024-02-06 19:37:27 +08:00
parent 4598d5b2ca
commit 6ca34f4055
5 changed files with 178 additions and 92 deletions
@@ -18,45 +18,18 @@ import ModelCard from "./ModelCard";
import InstallModelSearchBar from "./InstallModelSearchBar";
import ChooseFolder from "./ChooseFolder";
import InstallProgress from "./InstallProgress";
import { indexdb } from "../../db-tables/indexdb";
import AddApiKeyPopover from "./AddApiKeyPopover";
import { getCivitApiKey } from "../../utils/civitUtils";
import { useStateRef } from "../../customHooks/useStateRef";
import {
SearchHit,
SearchRequestBody,
SearchResponse,
SearchModelVersion,
} from "../civitSearchTypes";
const ALL_MODEL_TYPES = [
"Checkpoint",
"TextualInversion",
"Hypernetwork",
"LORA",
"Controlnet",
"Upscaler",
"VAE",
// "Poses",
// "MotionModule",
// "LoCon",
// "AestheticGradient",
// "Wildcards",
] as const; // `as const` makes the array readonly and its elements literal types
const CACHE_EXPIRY_DAYS = 2;
// Infer MODEL_TYPE from the ALL_MODEL_TYPES array
type MODEL_TYPE = (typeof ALL_MODEL_TYPES)[number];
const MODEL_TYPE_TO_FOLDER_MAPPING: Record<MODEL_TYPE, string> = {
Checkpoint: "checkpoints",
TextualInversion: "embeddings",
Hypernetwork: "hypernetworks",
LORA: "loras",
Controlnet: "controlnet",
Upscaler: "upscale_models",
VAE: "vae",
};
ALL_MODEL_TYPES,
FileEssential,
MODEL_TYPE,
MODEL_TYPE_TO_FOLDER_MAPPING,
apiResponse,
} from "./util/modelTypes";
import { getModelFromCivitAPi } from "./util/getModelFromCivitAPI";
import { getModelFromSearch } from "./util/getModelFromSearch";
interface Props {
onclose: () => void;
@@ -68,63 +41,22 @@ export default function InatallModelsModal({
searchQuery: searchQueryProp = "",
modelType: modelTypeProp,
}: Props) {
const [models, setModels] = useState<SearchHit[]>([]);
const [models, setModels] = useState<apiResponse[]>([]);
const [loading, setLoading] = useState(false);
const [modelType, setModelType] = useState(modelTypeProp);
const toast = useToast();
const [installing, setInstalling] = useState<string[]>([]);
const [searchQuery, setSearchQuery] = useState(searchQueryProp);
const { isOpen, onOpen, onClose: onCloseChooseFolderModal } = useDisclosure();
const [fileState, setFile, file] = useStateRef<SearchModelVersion>();
const [fileState, setFile, file] = useStateRef<FileEssential>();
const loadData = useCallback(async () => {
setLoading(true);
const params: SearchRequestBody = {
limit: 30,
filter: "nsfw = false ",
};
if (searchQuery !== "") {
params.q = searchQuery;
}
if (modelType != null) {
params.filter += `AND type = ${modelType}`;
}
const body = JSON.stringify(params);
const cacheEntry = await indexdb.cache?.get(body);
if (cacheEntry?.value != null) {
try {
const { data, timestamp } = JSON.parse(cacheEntry?.value);
// Check if cached data is still valid
const ageInDays = (Date.now() - timestamp) / (1000 * 60 * 60 * 24);
if (ageInDays < CACHE_EXPIRY_DAYS) {
setModels(data);
setLoading(false);
return;
}
} catch (e) {
console.error("err fetching cache", e);
}
}
const data = await fetch(import.meta.env.VITE_CMODEL_SEARCH_URL as string, {
headers: {
"Content-Type": "application/json",
Authorization: import.meta.env.VITE_CMODEL_APP_KEY as string,
},
method: "POST",
body,
});
const json: SearchResponse = await data.json();
setModels(json.hits);
// only cache if there is no search query
if (searchQuery === "") {
indexdb.cache.put({
id: body,
value: JSON.stringify({
data: json.hits,
timestamp: Date.now(),
}),
});
if (!searchQuery) {
const models = await getModelFromCivitAPi(modelType);
setModels(models);
} else {
const models = await getModelFromSearch(searchQuery, modelType);
setModels(models);
}
setLoading(false);
}, [searchQuery, modelType]);
@@ -159,14 +91,14 @@ export default function InatallModelsModal({
url += `?token=${apiKey}`;
}
installModelsApi({
file_hash: file.current?.hashes?.[2],
file_hash: file.current?.SHA256,
save_path: folderPath,
url,
});
setFile(undefined);
onCloseChooseFolderModal();
};
const onClickInstallModel = (_file: SearchModelVersion, model: SearchHit) => {
const onClickInstallModel = (_file: FileEssential, model: apiResponse) => {
const folderPath = MODEL_TYPE_TO_FOLDER_MAPPING[model.type as MODEL_TYPE];
setFile(_file);
if (folderPath == null) {
@@ -11,13 +11,19 @@ import {
} from "@chakra-ui/react";
import { useCallback, useState } from "react";
import { IconDownload } from "@tabler/icons-react";
import { SearchHit, SearchModelVersion } from "../civitSearchTypes";
import { findLeastNsfwImage } from "../../utils/findLeastNsfwImage";
import {
FileEssential,
apiResponse,
isCivitModel,
isCivitVersion,
} from "./util/modelTypes";
import { KBtoGB } from "../utils";
const IMAGE_SIZE = 280;
interface ModelCardProps {
model: SearchHit;
onClickInstallModel: (file: SearchModelVersion, model: SearchHit) => void;
model: apiResponse;
onClickInstallModel: (file: FileEssential, model: apiResponse) => void;
installing: string[];
}
export default function ModelCard({
@@ -25,8 +31,10 @@ export default function ModelCard({
onClickInstallModel,
installing,
}: ModelCardProps) {
const modelPhoto = `https://image.civitai.com/xG1nkqKTMzGDvpLrqFT7WA/${findLeastNsfwImage(model.images)?.url}/width=${IMAGE_SIZE}/`;
const versions = model.versions;
const modelPhoto = isCivitModel(model)
? model.modelVersions?.at(0)?.images?.at(0)?.url
: `https://image.civitai.com/xG1nkqKTMzGDvpLrqFT7WA/${findLeastNsfwImage(model.images)?.url}/width=${IMAGE_SIZE}/`;
const versions = isCivitModel(model) ? model.modelVersions : model.versions;
const [selectedFile, setSelectedFile] = useState<string>(
versions?.[0]?.name ?? "",
);
@@ -41,8 +49,18 @@ export default function ModelCard({
console.error("no file is find by name", selectedFile);
return;
}
onClickInstallModel(curFile, model);
let SHA256;
if (isCivitVersion(curFile)) {
SHA256 = curFile.files?.[0].hashes?.SHA256;
} else {
SHA256 = curFile.hashes?.[2];
}
onClickInstallModel({ ...curFile, SHA256 }, model);
}, [selectedFile]);
const sizeKB =
curFile && isCivitVersion(curFile)
? curFile?.files?.[0]?.sizeKB
: undefined;
return (
<Card width={IMAGE_SIZE} justifyContent={"space-between"} mb={2} gap={1}>
<Image
@@ -115,6 +133,13 @@ export default function ModelCard({
);
})}
</Select>
{sizeKB && (
<Tooltip label={KBtoGB(sizeKB)}>
<Text flexShrink={1} noOfLines={1} width={"40%"}>
{KBtoGB(sizeKB)}
</Text>
</Tooltip>
)}
</HStack>
</Stack>
</Card>
@@ -0,0 +1,48 @@
import { indexdb } from "../../../db-tables/indexdb";
import { CivitiModel } from "../../types";
import { CACHE_EXPIRY_DAYS, MODEL_TYPE } from "./modelTypes";
type CivitModelQueryParams = {
types?: MODEL_TYPE;
query?: string;
limit?: string;
nsfw?: "false";
};
export async function getModelFromCivitAPi(
types?: MODEL_TYPE,
): Promise<CivitiModel[]> {
const params: CivitModelQueryParams = {
limit: "30",
nsfw: "false",
types,
};
const queryString = new URLSearchParams(params).toString();
const fullURL = `https://civitai.com/api/v1/models?${queryString}`;
const cacheEntry = await indexdb.cache?.get(fullURL);
if (cacheEntry?.value != null) {
try {
const { data, timestamp } = JSON.parse(cacheEntry?.value);
// Check if cached data is still valid
const ageInDays = (Date.now() - timestamp) / (1000 * 60 * 60 * 24);
if (ageInDays < CACHE_EXPIRY_DAYS) {
return data;
}
} catch (e) {
console.error("err fetching cache", e);
}
}
const data = await fetch(fullURL);
const json = await data.json();
indexdb.cache.put({
id: fullURL,
value: JSON.stringify({
data: json.items,
timestamp: Date.now(),
}),
});
return json.items;
}
@@ -0,0 +1,30 @@
import {
SearchHit,
SearchRequestBody,
SearchResponse,
} from "../../civitSearchTypes";
import { MODEL_TYPE } from "./modelTypes";
export async function getModelFromSearch(
q: string,
type?: MODEL_TYPE,
): Promise<SearchHit[]> {
const params: SearchRequestBody = {
limit: 30,
filter: "nsfw = false ",
q,
};
if (type) {
params.filter += `AND type = ${type}`;
}
const data = await fetch(import.meta.env.VITE_CMODEL_SEARCH_URL as string, {
headers: {
"Content-Type": "application/json",
Authorization: import.meta.env.VITE_CMODEL_APP_KEY as string,
},
method: "POST",
body: JSON.stringify(params),
});
const json: SearchResponse = await data.json();
return json.hits;
}
@@ -0,0 +1,51 @@
import { SearchHit, SearchModelVersion } from "../../civitSearchTypes";
import { CivitiModel, CivitiModelVersion } from "../../types";
export const ALL_MODEL_TYPES = [
"Checkpoint",
"TextualInversion",
"Hypernetwork",
"LORA",
"Controlnet",
"Upscaler",
"VAE",
// "Poses",
// "MotionModule",
// "LoCon",
// "AestheticGradient",
// "Wildcards",
] as const; // `as const` makes the array readonly and its elements literal types
export const CACHE_EXPIRY_DAYS = 2;
// Infer MODEL_TYPE from the ALL_MODEL_TYPES array
export type MODEL_TYPE = (typeof ALL_MODEL_TYPES)[number];
export const MODEL_TYPE_TO_FOLDER_MAPPING: Record<MODEL_TYPE, string> = {
Checkpoint: "checkpoints",
TextualInversion: "embeddings",
Hypernetwork: "hypernetworks",
LORA: "loras",
Controlnet: "controlnet",
Upscaler: "upscale_models",
VAE: "vae",
};
export type FileEssential = {
id: number;
name?: string;
SHA256?: string;
};
export type apiResponse = CivitiModel | SearchHit;
export function isCivitModel(
model: CivitiModel | SearchHit,
): model is CivitiModel {
return (model as CivitiModel).modelVersions !== undefined;
}
export function isCivitVersion(
version: CivitiModelVersion | SearchModelVersion,
): version is CivitiModelVersion {
return (version as CivitiModelVersion).files?.[0] !== undefined;
}