Merge branch 'dev' into aller

This commit is contained in:
yanle
2024-01-21 18:06:06 +08:00
5 changed files with 89 additions and 38 deletions
+29 -17
View File
@@ -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):
+1 -4
View File
@@ -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);
}
+15
View File
@@ -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<string | null>(null);
const [route, setRoute] = useState<Route>("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() ?? {};
+8 -13
View File
@@ -6,7 +6,7 @@ import { api } from "/scripts/api.js";
export const useUpdateModels = () => {
// all model types
const [modelTypeList, setModelTypeList] = useState<string[]>([]);
const [modelTypeList, setModelTypeList] = useState<string[]>(["checkpoints"]);
// all models
const [modelsList, setModelsList] = useState<ModelsListRespItem[]>([]);
@@ -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);
};
@@ -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 (
<Flex
position="relative"
borderRadius={4}
bg="rgba(0, 0, 0, 0.5)"
height={178}
justifyContent="center"
alignItems="center"
>
<Spinner />
<Text
position="absolute"
bottom="0"
left="0"
right="0"
bg="rgba(0, 0, 0, 0.5)"
color="white"
textAlign="center"
p="0"
fontSize={12}
borderBottomRightRadius={4}
borderBottomLeftRadius={4}
>
{data.model_name}
</Text>
</Flex>
);
}
return (
<Box position="relative" borderRadius={4}>
<Image