diff --git a/service/model_manager/model_list.py b/service/model_manager/model_list.py index 01c40ae..a8c763c 100644 --- a/service/model_manager/model_list.py +++ b/service/model_manager/model_list.py @@ -9,12 +9,15 @@ import threading import os # The path to the file where file_hash_dict will be saved -FILE_HASH_DICT_PATH = os.path.join(os.path.dirname(__file__),"../../hash/file_hash_dict") -if not os.path.exists(FILE_HASH_DICT_PATH): - os.makedirs(FILE_HASH_DICT_PATH) +FILE_HASH_DICT_FOLDER_PATH = os.path.join(os.path.dirname(__file__),"../../hash") +FILE_HASH_DICT_PATH = os.path.join(FILE_HASH_DICT_FOLDER_PATH,"file_hash_dict") +if not os.path.exists(FILE_HASH_DICT_FOLDER_PATH): + os.makedirs(FILE_HASH_DICT_FOLDER_PATH) # Use shelve to store file_hash_dict file_hash_dict = shelve.open(FILE_HASH_DICT_PATH) + file_list = [] +file_list_lock = threading.Lock() # Add a global variable to track if populate_file_hash_dict is done populate_started = False @@ -30,46 +33,55 @@ def compute_hash(file_path): return sha256_hash.hexdigest() def process_file(folder, file): - global file_list file_path = folder_paths.get_full_path(folder, file) # check if the file hash is already calculated if (file_path in file_hash_dict): file_hash = file_hash_dict[file_path] else: - file_hash = compute_hash(file_path) - file_hash_dict[file_path] = file_hash - model_name = os.path.splitext(file)[0] - file_list.append({"model_name": model_name, "model_type": folder, "model_extension": os.path.splitext(file)[1], "file_hash": file_hash}) - send_ws('model_list', file_list) + file_hash = None # placeholder for hash + [model_name, model_extension] = os.path.splitext(file) + return {"model_name": model_name, "model_type": folder, "model_extension": model_extension, "file_hash": file_hash} def populate_file_hash_dict(): - global populate_done + global populate_done, file_list + local_file_list = [] for folder in folder_paths.folder_names_and_paths: if (folder == "configs" or folder == "custom_nodes"): continue files = folder_paths.get_filename_list(folder) for file in files: - process_file(folder, file) + local_file_list.append(process_file(folder, file)) + with file_list_lock: + file_list = local_file_list + send_ws('model_list', file_list) # send the list without hash to the client + # After all files are added, calculate hashes for those that don't have it yet + for file in file_list: + if file["file_hash"] is None: + file_path = folder_paths.get_full_path(file["model_type"], file["model_name"] + file["model_extension"]) + file_hash = compute_hash(file_path) + file_hash_dict[file_path] = file_hash + with file_list_lock: + file["file_hash"] = file_hash + send_ws('model_list', file_list) # update the list with the new hash # Set populate_done to True when done populate_done = True - send_ws('model_list', "done") def start_populate_file_hash_dict(): - global populate_thread, populate_started, populate_done, file_list + global populate_thread, populate_started, populate_done populate_running = populate_started == True and populate_done == False if (populate_running): return - file_list = [] # reset file_list populate_started = True populate_thread = threading.Thread(target=populate_file_hash_dict) populate_thread.daemon = True populate_thread.start() +start_populate_file_hash_dict() + @server.PromptServer.instance.routes.get("/model_manager/get_model_list") def get_model_list(request): - if (populate_started == False): - start_populate_file_hash_dict() - return web.json_response({"file_list": file_list, "populate_done": populate_done}, content_type='application/json') + with file_list_lock: + return web.json_response(file_list, content_type='application/json') loop = asyncio.get_event_loop() def send_ws(event, data): diff --git a/ui/src/Api.ts b/ui/src/Api.ts index f37339e..adc9aed 100644 --- a/ui/src/Api.ts +++ b/ui/src/Api.ts @@ -171,10 +171,7 @@ export async function getAllModelsList() { }, }); const result = await response.json(); - return result as { - file_list: ModelsListRespItem[]; - populate_done: boolean; - }; + return result as ModelsListRespItem[]; } catch (error) { console.error("Error get all models list:", error); } diff --git a/ui/src/App.tsx b/ui/src/App.tsx index a7f11fb..7002116 100644 --- a/ui/src/App.tsx +++ b/ui/src/App.tsx @@ -39,6 +39,14 @@ const RecentFilesDrawer = React.lazy( const GalleryModal = React.lazy(() => import("./gallery/GalleryModal")); import { scanLocalNewFiles } from "./Api"; +const usedWsEvents = [ + // InstallProgress.tsx + "download_progress", + "download_error", + // useUpdateModels.ts + "model_list", +]; + export default function App() { const [curFlowName, setCurFlowName] = useState(null); const [route, setRoute] = useState("root"); @@ -112,6 +120,7 @@ export default function App() { // }, // }; // app.registerExtension(ext); + subsribeToWsToStopWarning(); localStorage.removeItem("workspace"); localStorage.removeItem("comfyspace"); try { @@ -159,6 +168,12 @@ export default function App() { } }; + const subsribeToWsToStopWarning = () => { + usedWsEvents.forEach((event) => { + api.addEventListener(event, () => null); + }); + }; + const checkIsDirty = () => { if (curFlowID.current != null) { const graphJson = app.graph.serialize() ?? {}; diff --git a/ui/src/model-manager/hooks/useUpdateModels.ts b/ui/src/model-manager/hooks/useUpdateModels.ts index 68f0be7..69e8204 100644 --- a/ui/src/model-manager/hooks/useUpdateModels.ts +++ b/ui/src/model-manager/hooks/useUpdateModels.ts @@ -6,7 +6,7 @@ import { api } from "/scripts/api.js"; export const useUpdateModels = () => { // all model types - const [modelTypeList, setModelTypeList] = useState([]); + const [modelTypeList, setModelTypeList] = useState(["checkpoints"]); // all models const [modelsList, setModelsList] = useState([]); @@ -16,33 +16,28 @@ export const useUpdateModels = () => { useEffect(() => { initData(); - api.addEventListener("model_list", (e: { detail: ModelsListRespItem[] | "done" }) => { - if (e.detail === "done") { - setLoading(false); - } else { + api.addEventListener("model_list", (e: { detail: ModelsListRespItem[] }) => { updateModels(e.detail); - } }); }, []); const initData = async () => { - const res = await getAllModelsList(); - if (!res) return; - const { file_list, populate_done } = res; - if (populate_done) setLoading(false); + const file_list = await getAllModelsList(); updateModels(file_list); }; - const updateModels = async (file_list: ModelsListRespItem[]) => { + const updateModels = async (file_list?: ModelsListRespItem[]) => { + if (!file_list) return; + setLoading(false); const modelTypeList = Array.from( new Set(file_list.map((item) => item.model_type)) ); // checkpoints must be in first const index = modelTypeList.indexOf("checkpoints"); - if (index > 0) { + if (index >= 0) { modelTypeList.splice(index, 1); - modelTypeList.unshift("checkpoints"); } + modelTypeList.unshift("checkpoints"); setModelTypeList(modelTypeList); setModelsList(file_list); }; diff --git a/ui/src/model-manager/models-list-drawer/ModelItem.tsx b/ui/src/model-manager/models-list-drawer/ModelItem.tsx index e943e43..32a46d6 100644 --- a/ui/src/model-manager/models-list-drawer/ModelItem.tsx +++ b/ui/src/model-manager/models-list-drawer/ModelItem.tsx @@ -1,4 +1,4 @@ -import { Box, Image, Text } from "@chakra-ui/react"; +import { Box, Flex, Image, Spinner, Text } from "@chakra-ui/react"; import { ModelsListRespItem } from "../types"; import { useEffect, useState } from "react"; @@ -10,10 +10,12 @@ export function ModelItem({ data }: Props) { const [url, setUrl] = useState( "https://image.civitai.com/xG1nkqKTMzGDvpLrqFT7WA/27fd7433-cb0a-4a87-88c1-21ccb2b1a842/width=450/00060-881622046.jpeg" ); + const [hashing, setHashing] = useState(!data.file_hash); useEffect(() => { + setHashing(!data.file_hash); getThumbnail(); - }, []); + }, [data.file_hash]); const getThumbnail = async () => { // if local storage has the thumbnail, use it @@ -31,7 +33,7 @@ export function ModelItem({ data }: Props) { if (thumbnail && valid_url) { setUrl(thumbnail); - } else { + } else if (!hashing) { try { const url = `https://civitai.com/api/v1/model-versions/by-hash/${data.file_hash}`; const resp = await fetch(url); @@ -45,10 +47,40 @@ export function ModelItem({ data }: Props) { if (image_url) { localStorage.setItem(key, image_url); } - } catch (e) {} + } catch (e) { } } }; + if (hashing) { + return ( + + + + {data.model_name} + + + ); + } + return (