Refactor model installation code to support both APIs
This commit is contained in:
@@ -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;
|
||||
}
|
||||
Reference in New Issue
Block a user