Compare commits

...
22 Commits
Author SHA1 Message Date
nick 0582d1d869 merge 2024-08-07 20:43:38 -07:00
nick ce073a86c7 block on bad prompt 2024-08-07 20:42:12 -07:00
Emmanuel Morales 3a85a1edf2 feat(text): create node for external text list (#60)
* feat(text): create node for external text list 

This is to send a list of texts to other nodes

* refactor: remove prints and rename variable

* style: update comment

* refactor: remove unused optional inputs
2024-08-06 21:35:46 -06:00
karrix 369c1456a9 add: node focusing function 2024-08-05 00:59:52 +08:00
bennykok 01e323b7e2 fix: excessive log 2024-08-03 22:22:06 -07:00
bennykok db684d044a fix: not yield 2024-08-03 21:56:16 -07:00
BennyKok 8e12803ea1 Retry logic when calling api (#57)
* fix: retry logic, bypass logfire, clean up log

* fix: max_retries and retry_delay_multiplier, do not throw when pass the retry failed
2024-08-01 20:43:21 -07:00
Nick Kao 7585d5049a Merge pull request #58 from GwonHyeok/main
fix: ExternalLoRA node Make downloaded files reusable
2024-08-01 19:50:59 -07:00
GwonHyeok 772bb09240 fix: ExternalLoRA node Make downloaded files reusable 2024-08-02 10:29:24 +09:00
bennykok 9a7e18e651 fix: fe communication 2024-08-01 10:50:08 -07:00
Hmily a02c8d237f fix: Fix request deploy service interface error (#56) 2024-08-01 10:47:45 -07:00
nick 2ba5a0ff3d external lora 2024-08-01 10:43:24 -07:00
bennykok e0eae1068b fix: make external lora and checkpoint wildcard 2024-07-26 17:39:40 -07:00
bennykok 4f1a80fb64 fix: log issues with websocket 2024-07-22 13:36:39 -07:00
Hmily b4273b1907 fix: update next version and routing parameter errors (#55) 2024-07-22 09:40:23 -07:00
nick 10ba00e3dd update: external video node 2024-07-20 00:16:39 -07:00
nick eb40fddb76 Merge branch 'main' of https://github.com/bennykok/comfyui-deploy 2024-07-20 00:16:27 -07:00
nick 3c9d1865ca video node 2024-07-20 00:15:41 -07:00
bennykok 6fa38e9bb8 fix 2024-07-13 19:17:30 -07:00
bennykok 6e4532078f feat: update plugin js 2024-07-12 12:24:10 -07:00
nick 48d21f8d52 feat: audio output from external video node 2024-07-12 11:20:18 -07:00
BennyKokandnick a2ac1adf01 Streaming support (#52)
* feat: add streaming endpoint

* fix: run issues

* feat(plugin): add dispatchAPIEventData

* fix(plugin): event

* fix: streaming event format

* fix: prompt error

* fix: node_error proxy

* chore(plugin): add log

* custom route

---------

Co-authored-by: nick <[email protected]>
2024-07-11 20:03:41 -07:00
11 changed files with 1943 additions and 1140 deletions
+7 -1
View File
@@ -5,6 +5,12 @@ import torch
import folder_paths import folder_paths
from tqdm import tqdm from tqdm import tqdm
class AnyType(str):
def __ne__(self, __value: object) -> bool:
return False
WILDCARD = AnyType("*")
class ComfyUIDeployExternalCheckpoint: class ComfyUIDeployExternalCheckpoint:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
@@ -20,7 +26,7 @@ class ComfyUIDeployExternalCheckpoint:
} }
} }
RETURN_TYPES = (folder_paths.get_filename_list("checkpoints"),) RETURN_TYPES = (WILDCARD,)
RETURN_NAMES = ("path",) RETURN_NAMES = ("path",)
FUNCTION = "run" FUNCTION = "run"
+25 -6
View File
@@ -5,6 +5,14 @@ import torch
import folder_paths import folder_paths
class AnyType(str):
def __ne__(self, __value: object) -> bool:
return False
WILDCARD = AnyType("*")
class ComfyUIDeployExternalLora: class ComfyUIDeployExternalLora:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
@@ -17,27 +25,38 @@ class ComfyUIDeployExternalLora:
}, },
"optional": { "optional": {
"default_lora_name": (folder_paths.get_filename_list("loras"),), "default_lora_name": (folder_paths.get_filename_list("loras"),),
"lora_save_name": ( # if `default_lora_name` is a link to download a file, we will attempt to save it with this name
"STRING",
{"multiline": False, "default": ""},
),
}, },
} }
RETURN_TYPES = (folder_paths.get_filename_list("loras"),) RETURN_TYPES = (WILDCARD,)
RETURN_NAMES = ("path",) RETURN_NAMES = ("path",)
FUNCTION = "run" FUNCTION = "run"
CATEGORY = "deploy" CATEGORY = "deploy"
def run(self, input_id, default_lora_name=None): def run(self, input_id, default_lora_name=None, lora_save_name=None):
import requests import requests
import os import os
import uuid import uuid
if default_lora_name.startswith("http"): if default_lora_name.startswith("http"):
unique_filename = str(uuid.uuid4()) + ".safetensors" if lora_save_name:
print(unique_filename) existing_loras = folder_paths.get_filename_list("loras")
# Check if lora_save_name exists in the list
if lora_save_name in existing_loras:
print(f"using lora: {lora_save_name}")
return (lora_save_name,)
else:
lora_save_name = str(uuid.uuid4()) + ".safetensors"
print(lora_save_name)
print(folder_paths.folder_names_and_paths["loras"][0][0]) print(folder_paths.folder_names_and_paths["loras"][0][0])
destination_path = os.path.join( destination_path = os.path.join(
folder_paths.folder_names_and_paths["loras"][0][0], unique_filename folder_paths.folder_names_and_paths["loras"][0][0], lora_save_name
) )
print(destination_path) print(destination_path)
print("Downloading external lora - " + input_id + " to " + destination_path) print("Downloading external lora - " + input_id + " to " + destination_path)
@@ -48,7 +67,7 @@ class ComfyUIDeployExternalLora:
) )
with open(destination_path, "wb") as out_file: with open(destination_path, "wb") as out_file:
out_file.write(response.content) out_file.write(response.content)
return (unique_filename,) return (lora_save_name,)
else: else:
print(f"using lora: {default_lora_name}") print(f"using lora: {default_lora_name}")
return (default_lora_name,) return (default_lora_name,)
+43
View File
@@ -0,0 +1,43 @@
import folder_paths
from PIL import Image, ImageOps
import numpy as np
import torch
import json
class ComfyUIDeployExternalTextList:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"input_id": (
"STRING",
{"multiline": False, "default": 'input_text_list'},
),
"text": (
"STRING",
{"multiline": True, "default": "[]"},
),
}
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("text",)
OUTPUT_IS_LIST = (True,)
FUNCTION = "run"
CATEGORY = "text"
def run(self, input_id, text=None):
text_list = []
try:
text_list = json.loads(text) # Assuming text is a JSON array string
except Exception as e:
print(f"Error processing images: {e}")
pass
return [text_list]
NODE_CLASS_MAPPINGS = {"ComfyUIDeployExternalTextList": ComfyUIDeployExternalTextList}
NODE_DISPLAY_NAME_MAPPINGS = {"ComfyUIDeployExternalTextList": "External Text List (ComfyUI Deploy)"}
+325 -64
View File
@@ -1,10 +1,15 @@
# credit goes to https://github.com/Kosinkadink/ComfyUI-VideoHelperSuite and is meant to work with # credit goes to https://github.com/Kosinkadink/ComfyUI-VideoHelperSuite
# Intended to work with https://github.com/NicholasKao1029/ComfyUI-VideoHelperSuite/tree/main
import os import os
import itertools import itertools
import numpy as np import numpy as np
import torch import torch
from typing import Union
from torch import Tensor
import cv2 import cv2
import psutil
from collections.abc import Mapping
import folder_paths import folder_paths
from comfy.utils import common_upscale from comfy.utils import common_upscale
@@ -90,13 +95,25 @@ if gifski_path is None:
gifski_path = shutil.which("gifski") gifski_path = shutil.which("gifski")
def is_safe_path(path):
if "VHS_STRICT_PATHS" not in os.environ:
return True
basedir = os.path.abspath(".")
try:
common_path = os.path.commonpath([basedir, path])
except:
# Different drive on windows
return False
return common_path == basedir
def get_sorted_dir_files_from_directory( def get_sorted_dir_files_from_directory(
directory: str, directory: str,
skip_first_images: int = 0, skip_first_images: int = 0,
select_every_nth: int = 1, select_every_nth: int = 1,
extensions: Iterable = None, extensions: Iterable = None,
): ):
directory = directory.strip() directory = strip_path(directory)
dir_files = os.listdir(directory) dir_files = os.listdir(directory)
dir_files = sorted(dir_files) dir_files = sorted(dir_files)
dir_files = [os.path.join(directory, x) for x in dir_files] dir_files = [os.path.join(directory, x) for x in dir_files]
@@ -177,18 +194,59 @@ def requeue_workflow(requeue_required=(-1, True)):
def get_audio(file, start_time=0, duration=0): def get_audio(file, start_time=0, duration=0):
args = [ffmpeg_path, "-v", "error", "-i", file] args = [ffmpeg_path, "-i", file]
if start_time > 0: if start_time > 0:
args += ["-ss", str(start_time)] args += ["-ss", str(start_time)]
if duration > 0: if duration > 0:
args += ["-t", str(duration)] args += ["-t", str(duration)]
try: try:
# TODO: scan for sample rate and maintain
res = subprocess.run( res = subprocess.run(
args + ["-f", "wav", "-"], stdout=subprocess.PIPE, check=True args + ["-f", "f32le", "-"], capture_output=True, check=True
).stdout )
audio = torch.frombuffer(bytearray(res.stdout), dtype=torch.float32)
match = re.search(", (\\d+) Hz, (\\w+), ", res.stderr.decode("utf-8"))
except subprocess.CalledProcessError as e: except subprocess.CalledProcessError as e:
return False raise Exception(
return res f"VHS failed to extract audio from {file}:\n" + e.stderr.decode("utf-8")
)
if match:
ar = int(match.group(1))
# NOTE: Just throwing an error for other channel types right now
# Will deal with issues if they come
ac = {"mono": 1, "stereo": 2}[match.group(2)]
else:
ar = 44100
ac = 2
audio = audio.reshape((-1, ac)).transpose(0, 1).unsqueeze(0)
return {"waveform": audio, "sample_rate": ar}
class LazyAudioMap(Mapping):
def __init__(self, file, start_time, duration):
self.file = file
self.start_time = start_time
self.duration = duration
self._dict = None
def __getitem__(self, key):
if self._dict is None:
self._dict = get_audio(self.file, self.start_time, self.duration)
return self._dict[key]
def __iter__(self):
if self._dict is None:
self._dict = get_audio(self.file, self.start_time, self.duration)
return iter(self._dict)
def __len__(self):
if self._dict is None:
self._dict = get_audio(self.file, self.start_time, self.duration)
return len(self._dict)
def lazy_get_audio(file, start_time=0, duration=0):
return LazyAudioMap(file, start_time, duration)
def lazy_eval(func): def lazy_eval(func):
@@ -230,6 +288,19 @@ def validate_sequence(path):
return False return False
def strip_path(path):
# This leaves whitespace inside quotes and only a single "
# thus ' ""test"' -> '"test'
# consider path.strip(string.whitespace+"\"")
# or weightier re.fullmatch("[\\s\"]*(.+?)[\\s\"]*", path).group(1)
path = path.strip()
if path.startswith('"'):
path = path[1:]
if path.endswith('"'):
path = path[:-1]
return path
def hash_path(path): def hash_path(path):
if path is None: if path is None:
return "input" return "input"
@@ -286,6 +357,145 @@ def target_size(
return (width, height) return (width, height)
def validate_index(
index: int,
length: int = 0,
is_range: bool = False,
allow_negative=False,
allow_missing=False,
) -> int:
# if part of range, do nothing
if is_range:
return index
# otherwise, validate index
# validate not out of range - only when latent_count is passed in
if length > 0 and index > length - 1 and not allow_missing:
raise IndexError(f"Index '{index}' out of range for {length} item(s).")
# if negative, validate not out of range
if index < 0:
if not allow_negative:
raise IndexError(f"Negative indeces not allowed, but was '{index}'.")
conv_index = length + index
if conv_index < 0 and not allow_missing:
raise IndexError(
f"Index '{index}', converted to '{conv_index}' out of range for {length} item(s)."
)
index = conv_index
return index
def convert_to_index_int(
raw_index: str,
length: int = 0,
is_range: bool = False,
allow_negative=False,
allow_missing=False,
) -> int:
try:
return validate_index(
int(raw_index),
length=length,
is_range=is_range,
allow_negative=allow_negative,
allow_missing=allow_missing,
)
except ValueError as e:
raise ValueError(f"Index '{raw_index}' must be an integer.", e)
def convert_str_to_indexes(
indexes_str: str, length: int = 0, allow_missing=False
) -> list[int]:
if not indexes_str:
return []
int_indexes = list(range(0, length))
allow_negative = length > 0
chosen_indexes = []
# parse string - allow positive ints, negative ints, and ranges separated by ':'
groups = indexes_str.split(",")
groups = [g.strip() for g in groups]
for g in groups:
# parse range of indeces (e.g. 2:16)
if ":" in g:
index_range = g.split(":", 2)
index_range = [r.strip() for r in index_range]
start_index = index_range[0]
if len(start_index) > 0:
start_index = convert_to_index_int(
start_index,
length=length,
is_range=True,
allow_negative=allow_negative,
allow_missing=allow_missing,
)
else:
start_index = 0
end_index = index_range[1]
if len(end_index) > 0:
end_index = convert_to_index_int(
end_index,
length=length,
is_range=True,
allow_negative=allow_negative,
allow_missing=allow_missing,
)
else:
end_index = length
# support step as well, to allow things like reversing, every-other, etc.
step = 1
if len(index_range) > 2:
step = index_range[2]
if len(step) > 0:
step = convert_to_index_int(
step,
length=length,
is_range=True,
allow_negative=True,
allow_missing=True,
)
else:
step = 1
# if latents were passed in, base indeces on known latent count
if len(int_indexes) > 0:
chosen_indexes.extend(int_indexes[start_index:end_index][::step])
# otherwise, assume indeces are valid
else:
chosen_indexes.extend(list(range(start_index, end_index, step)))
# parse individual indeces
else:
chosen_indexes.append(
convert_to_index_int(
g,
length=length,
allow_negative=allow_negative,
allow_missing=allow_missing,
)
)
return chosen_indexes
def select_indexes(input_obj: Union[Tensor, list], idxs: list):
if type(input_obj) == Tensor:
return input_obj[idxs]
else:
return [input_obj[i] for i in idxs]
def select_indexes_from_str(
input_obj: Union[Tensor, list], indexes: str, err_if_missing=True, err_if_empty=True
):
real_idxs = convert_str_to_indexes(
indexes, len(input_obj), allow_missing=not err_if_missing
)
if err_if_empty and len(real_idxs) == 0:
raise Exception(f"Nothing was selected based on indexes found in '{indexes}'.")
return select_indexes(input_obj, real_idxs)
###
def cv_frame_generator( def cv_frame_generator(
video, video,
force_rate, force_rate,
@@ -295,9 +505,10 @@ def cv_frame_generator(
meta_batch=None, meta_batch=None,
unique_id=None, unique_id=None,
): ):
video_cap = cv2.VideoCapture(video) video_cap = cv2.VideoCapture(strip_path(video))
if not video_cap.isOpened(): if not video_cap.isOpened():
raise ValueError(f"{video} could not be loaded with cv.") raise ValueError(f"{video} could not be loaded with cv.")
pbar = None
# extract video metadata # extract video metadata
fps = video_cap.get(cv2.CAP_PROP_FPS) fps = video_cap.get(cv2.CAP_PROP_FPS)
@@ -319,6 +530,8 @@ def cv_frame_generator(
target_frame_time = 1 / force_rate target_frame_time = 1 / force_rate
yield (width, height, fps, duration, total_frames, target_frame_time) yield (width, height, fps, duration, total_frames, target_frame_time)
if meta_batch is not None:
yield min(frame_load_cap, total_frames)
time_offset = target_frame_time - base_frame_time time_offset = target_frame_time - base_frame_time
while video_cap.isOpened(): while video_cap.isOpened():
@@ -349,7 +562,8 @@ def cv_frame_generator(
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
# convert frame to comfyui's expected format # convert frame to comfyui's expected format
# TODO: frame contains no exif information. Check if opencv2 has already applied # TODO: frame contains no exif information. Check if opencv2 has already applied
frame = np.array(frame, dtype=np.float32) / 255.0 frame = np.array(frame, dtype=np.float32)
torch.from_numpy(frame).div_(255)
if prev_frame is not None: if prev_frame is not None:
inp = yield prev_frame inp = yield prev_frame
if inp is not None: if inp is not None:
@@ -357,6 +571,8 @@ def cv_frame_generator(
return return
prev_frame = frame prev_frame = frame
frames_added += 1 frames_added += 1
if pbar is not None:
pbar.update_absolute(frames_added, frame_load_cap)
# if cap exists and we've reached it, stop processing frames # if cap exists and we've reached it, stop processing frames
if frame_load_cap > 0 and frames_added >= frame_load_cap: if frame_load_cap > 0 and frames_added >= frame_load_cap:
break break
@@ -367,6 +583,17 @@ def cv_frame_generator(
yield prev_frame yield prev_frame
def batched(it, n):
while batch := tuple(itertools.islice(it, n)):
yield batch
def batched_vae_encode(images, vae, frames_per_batch):
for batch in batched(images, frames_per_batch):
image_batch = torch.from_numpy(np.array(batch))
yield from vae.encode(image_batch).numpy()
def load_video_cv( def load_video_cv(
video: str, video: str,
force_rate: int, force_rate: int,
@@ -378,6 +605,8 @@ def load_video_cv(
select_every_nth: int, select_every_nth: int,
meta_batch=None, meta_batch=None,
unique_id=None, unique_id=None,
memory_limit_mb=None,
vae=None,
): ):
if meta_batch is None or unique_id not in meta_batch.inputs: if meta_batch is None or unique_id not in meta_batch.inputs:
gen = cv_frame_generator( gen = cv_frame_generator(
@@ -401,30 +630,89 @@ def load_video_cv(
total_frames, total_frames,
target_frame_time, target_frame_time,
) )
meta_batch.total_frames = min(meta_batch.total_frames, next(gen))
else: else:
(gen, width, height, fps, duration, total_frames, target_frame_time) = ( (gen, width, height, fps, duration, total_frames, target_frame_time) = (
meta_batch.inputs[unique_id] meta_batch.inputs[unique_id]
) )
memory_limit = None
if memory_limit_mb is not None:
memory_limit *= 2**20
else:
# TODO: verify if garbage collection should be performed here.
# leaves ~128 MB unreserved for safety
try:
memory_limit = (
psutil.virtual_memory().available + psutil.swap_memory().free
) - 2**27
except:
print(
"Failed to calculate available memory. Memory load limit has been disabled"
)
if memory_limit is not None:
if vae is not None:
# space required to load as f32, exist as latent with wiggle room, decode to f32
max_loadable_frames = int(
memory_limit // (width * height * 3 * (4 + 4 + 1 / 10))
)
else:
# TODO: use better estimate for when vae is not None
# Consider completely ignoring for load_latent case?
max_loadable_frames = int(memory_limit // (width * height * 3 * (0.1)))
if meta_batch is not None: if meta_batch is not None:
if meta_batch.frames_per_batch > max_loadable_frames:
raise RuntimeError(
f"Meta Batch set to {meta_batch.frames_per_batch} frames but only {max_loadable_frames} can fit in memory"
)
gen = itertools.islice(gen, meta_batch.frames_per_batch) gen = itertools.islice(gen, meta_batch.frames_per_batch)
else:
original_gen = gen
gen = itertools.islice(gen, max_loadable_frames)
downscale_ratio = getattr(vae, "downscale_ratio", 8)
frames_per_batch = (1920 * 1080 * 16) // (width * height) or 1
if force_size != "Disabled" or vae is not None:
new_size = target_size(
width, height, force_size, custom_width, custom_height, downscale_ratio
)
if new_size[0] != width or new_size[1] != height:
def rescale(frame):
s = torch.from_numpy(
np.fromiter(frame, np.dtype((np.float32, (height, width, 3))))
)
s = s.movedim(-1, 1)
s = common_upscale(s, new_size[0], new_size[1], "lanczos", "center")
return s.movedim(1, -1).numpy()
gen = itertools.chain.from_iterable(
map(rescale, batched(gen, frames_per_batch))
)
else:
new_size = width, height
if vae is not None:
gen = batched_vae_encode(gen, vae, frames_per_batch)
vw, vh = new_size[0] // downscale_ratio, new_size[1] // downscale_ratio
images = torch.from_numpy(np.fromiter(gen, np.dtype((np.float32, (4, vh, vw)))))
else:
# Some minor wizardry to eliminate a copy and reduce max memory by a factor of ~2 # Some minor wizardry to eliminate a copy and reduce max memory by a factor of ~2
images = torch.from_numpy( images = torch.from_numpy(
np.fromiter(gen, np.dtype((np.float32, (height, width, 3)))) np.fromiter(gen, np.dtype((np.float32, (new_size[1], new_size[0], 3))))
) )
if meta_batch is None and memory_limit is not None:
try:
next(original_gen)
raise RuntimeError(
f"Memory limit hit after loading {len(images)} frames. Stopping execution."
)
except StopIteration:
pass
if len(images) == 0: if len(images) == 0:
raise RuntimeError("No frames generated") raise RuntimeError("No frames generated")
if force_size != "Disabled":
new_size = target_size(width, height, force_size, custom_width, custom_height)
if new_size[0] != width or new_size[1] != height:
s = images.movedim(-1, 1)
s = common_upscale(s, new_size[0], new_size[1], "lanczos", "center")
images = s.movedim(1, -1)
# Setup lambda for lazy audio capture # Setup lambda for lazy audio capture
audio = lambda: get_audio( audio = lazy_get_audio(
video, video,
skip_first_frames * target_frame_time, skip_first_frames * target_frame_time,
frame_load_cap * target_frame_time * select_every_nth, frame_load_cap * target_frame_time * select_every_nth,
@@ -440,13 +728,16 @@ def load_video_cv(
"loaded_fps": 1 / target_frame_time, "loaded_fps": 1 / target_frame_time,
"loaded_frame_count": len(images), "loaded_frame_count": len(images),
"loaded_duration": len(images) * target_frame_time, "loaded_duration": len(images) * target_frame_time,
"loaded_width": images.shape[2], "loaded_width": new_size[0],
"loaded_height": images.shape[1], "loaded_height": new_size[1],
} }
if vae is None:
return (images, len(images), lazy_eval(audio), video_info) return (images, len(images), audio, video_info, None)
else:
return (None, len(images), audio, video_info, {"samples": images})
# modeled after Video upload node
class ComfyUIDeployExternalVideo: class ComfyUIDeployExternalVideo:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
@@ -457,68 +748,38 @@ class ComfyUIDeployExternalVideo:
file_parts = f.split(".") file_parts = f.split(".")
if len(file_parts) > 1 and (file_parts[-1] in video_extensions): if len(file_parts) > 1 and (file_parts[-1] in video_extensions):
files.append(f) files.append(f)
return { return {"required": {
"required": {
"input_id": ( "input_id": (
"STRING", "STRING",
{"multiline": False, "default": "input_video"}, {"multiline": False, "default": "input_video"},
), ),
"force_rate": ("INT", {"default": 0, "min": 0, "max": 60, "step": 1}), "force_rate": ("INT", {"default": 0, "min": 0, "max": 60, "step": 1}),
"force_size": ( "force_size": (["Disabled", "Custom Height", "Custom Width", "Custom", "256x?", "?x256", "256x256", "512x?", "?x512", "512x512"],),
[ "custom_width": ("INT", {"default": 512, "min": 0, "max": DIMMAX, "step": 8}),
"Disabled", "custom_height": ("INT", {"default": 512, "min": 0, "max": DIMMAX, "step": 8}),
"Custom Height", "frame_load_cap": ("INT", {"default": 0, "min": 0, "max": BIGMAX, "step": 1}),
"Custom Width", "skip_first_frames": ("INT", {"default": 0, "min": 0, "max": BIGMAX, "step": 1}),
"Custom", "select_every_nth": ("INT", {"default": 1, "min": 1, "max": BIGMAX, "step": 1}),
"256x?",
"?x256",
"256x256",
"512x?",
"?x512",
"512x512",
],
),
"custom_width": (
"INT",
{"default": 512, "min": 0, "max": DIMMAX, "step": 8},
),
"custom_height": (
"INT",
{"default": 512, "min": 0, "max": DIMMAX, "step": 8},
),
"frame_load_cap": (
"INT",
{"default": 0, "min": 0, "max": BIGMAX, "step": 1},
),
"skip_first_frames": (
"INT",
{"default": 0, "min": 0, "max": BIGMAX, "step": 1},
),
"select_every_nth": (
"INT",
{"default": 1, "min": 1, "max": BIGMAX, "step": 1},
),
}, },
"optional": { "optional": {
"meta_batch": ("VHS_BatchManager",), "meta_batch": ("VHS_BatchManager",),
"vae": ("VAE",),
"default_value": (sorted(files),), "default_value": (sorted(files),),
}, },
"hidden": {"unique_id": "UNIQUE_ID"}, "hidden": {
"unique_id": "UNIQUE_ID"
},
} }
CATEGORY = "Video Helper Suite 🎥🅥🅗🅢" CATEGORY = "Video Helper Suite 🎥🅥🅗🅢"
RETURN_TYPES = ( RETURN_TYPES = ("IMAGE", "INT", "AUDIO", "VHS_VIDEOINFO", "LATENT")
"IMAGE",
"INT",
"VHS_AUDIO",
"VHS_VIDEOINFO",
)
RETURN_NAMES = ( RETURN_NAMES = (
"IMAGE", "IMAGE",
"frame_count", "frame_count",
"audio", "audio",
"video_info", "video_info",
"LATENT",
) )
FUNCTION = "load_video" FUNCTION = "load_video"
+348 -94
View File
@@ -18,13 +18,128 @@ import threading
import hashlib import hashlib
import aiohttp import aiohttp
import aiofiles import aiofiles
from typing import List, Union, Any, Optional from typing import Dict, List, Union, Any, Optional
from PIL import Image from PIL import Image
import copy import copy
import struct import struct
from aiohttp import ClientError
import atexit
# Global session
client_session = None
# def create_client_session():
# global client_session
# if client_session is None:
# client_session = aiohttp.ClientSession()
async def ensure_client_session():
global client_session
if client_session is None:
client_session = aiohttp.ClientSession()
async def cleanup():
global client_session
if client_session:
await client_session.close()
def exit_handler():
print("Exiting the application. Initiating cleanup...")
loop = asyncio.get_event_loop()
loop.run_until_complete(cleanup())
atexit.register(exit_handler)
max_retries = int(os.environ.get('MAX_RETRIES', '3'))
retry_delay_multiplier = float(os.environ.get('RETRY_DELAY_MULTIPLIER', '2'))
print(f"max_retries: {max_retries}, retry_delay_multiplier: {retry_delay_multiplier}")
async def async_request_with_retry(method, url, **kwargs):
global client_session
await ensure_client_session()
retry_delay = 1 # Start with 1 second delay
for attempt in range(max_retries):
try:
async with client_session.request(method, url, **kwargs) as response:
response.raise_for_status()
return response
except ClientError as e:
if attempt == max_retries - 1:
logger.error(f"Request failed after {max_retries} attempts: {e}")
# raise
logger.warning(f"Request failed (attempt {attempt + 1}/{max_retries}): {e}")
await asyncio.sleep(retry_delay)
retry_delay *= retry_delay_multiplier # Exponential backoff
from logging import basicConfig, getLogger
# Check for an environment variable to enable/disable Logfire
use_logfire = os.environ.get('USE_LOGFIRE', 'false').lower() == 'true'
if use_logfire:
try:
import logfire
logfire.configure(
send_to_logfire="if-token-present"
)
logger = logfire
except ImportError:
print("Logfire not installed or disabled. Using standard Python logger.")
use_logfire = False
if not use_logfire:
# Use a standard Python logger when Logfire is disabled or not available
logger = getLogger("comfy-deploy")
basicConfig(level="INFO") # You can adjust the logging level as needed
def log(level, message, **kwargs):
if use_logfire:
getattr(logger, level)(message, **kwargs)
else:
getattr(logger, level)(f"{message} {kwargs}")
# For a span, you might need to create a context manager
from contextlib import contextmanager
@contextmanager
def log_span(name):
if use_logfire:
with logger.span(name):
yield
else:
yield
# logger.info(f"Start: {name}")
# yield
# logger.info(f"End: {name}")
from globals import StreamingPrompt, Status, sockets, SimplePrompt, streaming_prompt_metadata, prompt_metadata from globals import StreamingPrompt, Status, sockets, SimplePrompt, streaming_prompt_metadata, prompt_metadata
class EventEmitter:
def __init__(self):
self.listeners = {}
def on(self, event, listener):
if event not in self.listeners:
self.listeners[event] = []
self.listeners[event].append(listener)
def off(self, event, listener):
if event in self.listeners:
self.listeners[event].remove(listener)
if not self.listeners[event]:
del self.listeners[event]
def emit(self, event, *args, **kwargs):
if event in self.listeners:
for listener in self.listeners[event]:
listener(*args, **kwargs)
# Create a global event emitter instance
event_emitter = EventEmitter()
api = None api = None
api_task = None api_task = None
@@ -32,18 +147,18 @@ cd_enable_log = os.environ.get('CD_ENABLE_LOG', 'false').lower() == 'true'
cd_enable_run_log = os.environ.get('CD_ENABLE_RUN_LOG', 'false').lower() == 'true' cd_enable_run_log = os.environ.get('CD_ENABLE_RUN_LOG', 'false').lower() == 'true'
bypass_upload = os.environ.get('CD_BYPASS_UPLOAD', 'false').lower() == 'true' bypass_upload = os.environ.get('CD_BYPASS_UPLOAD', 'false').lower() == 'true'
print("CD_BYPASS_UPLOAD", bypass_upload) logger.info(f"CD_BYPASS_UPLOAD {bypass_upload}")
def clear_current_prompt(sid): def clear_current_prompt(sid):
prompt_server = server.PromptServer.instance prompt_server = server.PromptServer.instance
to_delete = list(streaming_prompt_metadata[sid].running_prompt_ids) # Convert set to list to_delete = list(streaming_prompt_metadata[sid].running_prompt_ids) # Convert set to list
print("clearning out prompt: ", to_delete) logger.info(f"clearing out prompt: {to_delete}")
for id_to_delete in to_delete: for id_to_delete in to_delete:
delete_func = lambda a: a[1] == id_to_delete delete_func = lambda a: a[1] == id_to_delete
prompt_server.prompt_queue.delete_queue_item(delete_func) prompt_server.prompt_queue.delete_queue_item(delete_func)
print("deleted prompt: ", id_to_delete, prompt_server.prompt_queue.get_tasks_remaining()) logger.info(f"deleted prompt: {id_to_delete}, remaining tasks: {prompt_server.prompt_queue.get_tasks_remaining()}")
streaming_prompt_metadata[sid].running_prompt_ids.clear() streaming_prompt_metadata[sid].running_prompt_ids.clear()
@@ -84,7 +199,7 @@ def post_prompt(json_data):
} }
return response return response
else: else:
print("invalid prompt:", valid[1]) logger.info("invalid prompt:", valid[1])
return {"error": valid[1], "node_errors": valid[3]} return {"error": valid[1], "node_errors": valid[3]}
else: else:
return {"error": "no prompt", "node_errors": []} return {"error": "no prompt", "node_errors": []}
@@ -158,11 +273,11 @@ def send_prompt(sid: str, inputs: StreamingPrompt):
# Random seed # Random seed
apply_random_seed_to_workflow(workflow_api) apply_random_seed_to_workflow(workflow_api)
print("getting inputs" , inputs.inputs) logger.info("getting inputs" , inputs.inputs)
apply_inputs_to_workflow(workflow_api, inputs.inputs, sid=sid) apply_inputs_to_workflow(workflow_api, inputs.inputs, sid=sid)
print(workflow_api) logger.info(workflow_api)
prompt_id = str(uuid.uuid4()) prompt_id = str(uuid.uuid4())
@@ -185,12 +300,11 @@ def send_prompt(sid: str, inputs: StreamingPrompt):
error_type = type(e).__name__ error_type = type(e).__name__
stack_trace_short = traceback.format_exc().strip().split('\n')[-2] stack_trace_short = traceback.format_exc().strip().split('\n')[-2]
stack_trace = traceback.format_exc().strip() stack_trace = traceback.format_exc().strip()
print(f"error: {error_type}, {e}") logger.info(f"error: {error_type}, {e}")
print(f"stack trace: {stack_trace_short}") logger.info(f"stack trace: {stack_trace_short}")
@server.PromptServer.instance.routes.post("/comfyui-deploy/run") @server.PromptServer.instance.routes.post("/comfyui-deploy/run")
async def comfy_deploy_run(request): async def comfy_deploy_run(request):
prompt_server = server.PromptServer.instance
data = await request.json() data = await request.json()
# In older version, we use workflow_api, but this has inputs already swapped in nextjs frontend, which is tricky # In older version, we use workflow_api, but this has inputs already swapped in nextjs frontend, which is tricky
@@ -221,8 +335,8 @@ async def comfy_deploy_run(request):
error_type = type(e).__name__ error_type = type(e).__name__
stack_trace_short = traceback.format_exc().strip().split('\n')[-2] stack_trace_short = traceback.format_exc().strip().split('\n')[-2]
stack_trace = traceback.format_exc().strip() stack_trace = traceback.format_exc().strip()
print(f"error: {error_type}, {e}") logger.info(f"error: {error_type}, {e}")
print(f"stack trace: {stack_trace_short}") logger.info(f"stack trace: {stack_trace_short}")
await update_run_with_output(prompt_id, { await update_run_with_output(prompt_id, {
"error": { "error": {
"error_type": error_type, "error_type": error_type,
@@ -234,15 +348,8 @@ async def comfy_deploy_run(request):
return web.Response(status=500, reason=f"{error_type}: {e}, {stack_trace_short}") return web.Response(status=500, reason=f"{error_type}: {e}, {stack_trace_short}")
status = 200 status = 200
# if "error" in res:
# status = 400
# await update_run_with_output(prompt_id, {
# "error": {
# **res
# }
# })
if "node_errors" in res and res["node_errors"]: if "node_errors" in res and res["node_errors"] is not None:
# Even tho there are node_errors it can still be run # Even tho there are node_errors it can still be run
status = 400 status = 400
await update_run_with_output(prompt_id, { await update_run_with_output(prompt_id, {
@@ -257,24 +364,134 @@ async def comfy_deploy_run(request):
return web.json_response(res, status=status) return web.json_response(res, status=status)
async def stream_prompt(data):
# In older version, we use workflow_api, but this has inputs already swapped in nextjs frontend, which is tricky
workflow_api = data.get("workflow_api_raw")
# The prompt id generated from comfy deploy, can be None
prompt_id = data.get("prompt_id")
inputs = data.get("inputs")
# Now it handles directly in here
apply_random_seed_to_workflow(workflow_api)
apply_inputs_to_workflow(workflow_api, inputs)
prompt = {
"prompt": workflow_api,
"client_id": "comfy_deploy_instance", #api.client_id
"prompt_id": prompt_id
}
prompt_metadata[prompt_id] = SimplePrompt(
status_endpoint=data.get('status_endpoint'),
file_upload_endpoint=data.get('file_upload_endpoint'),
workflow_api=workflow_api
)
# log('info', "Begin prompt", prompt=prompt)
try:
res = post_prompt(prompt)
except Exception as e:
error_type = type(e).__name__
stack_trace_short = traceback.format_exc().strip().split('\n')[-2]
stack_trace = traceback.format_exc().strip()
logger.info(f"error: {error_type}, {e}")
logger.info(f"stack trace: {stack_trace_short}")
await update_run_with_output(prompt_id, {
"error": {
"error_type": error_type,
"stack_trace": stack_trace
}
})
# When there are critical errors, the prompt is actually not run
await update_run(prompt_id, Status.FAILED)
# return web.Response(status=500, reason=f"{error_type}: {e}, {stack_trace_short}")
# raise Exception("Prompt failed")
status = 200
if "node_errors" in res and res["node_errors"] is not None:
# Even tho there are node_errors it can still be run
status = 400
await update_run_with_output(prompt_id, {
"error": {
**res
}
})
# When there are critical errors, the prompt is actually not run
if "error" in res:
await update_run(prompt_id, Status.FAILED)
# raise Exception("Prompt failed")
return res
# return web.json_response(res, status=status)
comfy_message_queues: Dict[str, asyncio.Queue] = {}
@server.PromptServer.instance.routes.post('/comfyui-deploy/run/streaming')
async def stream_response(request):
response = web.StreamResponse(status=200, reason='OK', headers={'Content-Type': 'text/event-stream'})
await response.prepare(request)
pending = True
data = await request.json()
prompt_id = data.get("prompt_id")
comfy_message_queues[prompt_id] = asyncio.Queue()
with log_span('Streaming Run'):
log('info', 'Streaming prompt')
try:
result = await stream_prompt(data=data)
await response.write(f"event: event_update\ndata: {json.dumps(result)}\n\n".encode('utf-8'))
# await response.write(.encode('utf-8'))
await response.drain() # Ensure the buffer is flushed
while pending:
if prompt_id in comfy_message_queues:
if not comfy_message_queues[prompt_id].empty():
data = await comfy_message_queues[prompt_id].get()
# log('info', data["event"], data=json.dumps(data))
# logger.info("listener", data)
await response.write(f"event: event_update\ndata: {json.dumps(data)}\n\n".encode('utf-8'))
await response.drain() # Ensure the buffer is flushed
if data["event"] == "status":
if data["data"]["status"] in (Status.FAILED.value, Status.SUCCESS.value):
pending = False
await asyncio.sleep(0.1) # Adjust the sleep duration as needed
except asyncio.CancelledError:
log('info', "Streaming was cancelled")
raise
except Exception as e:
log('error', "Streaming error", error=e)
finally:
# event_emitter.off("send_json", task)
await response.write_eof()
comfy_message_queues.pop(prompt_id, None)
return response
def get_comfyui_path_from_file_path(file_path): def get_comfyui_path_from_file_path(file_path):
file_path_parts = file_path.split("\\") file_path_parts = file_path.split("\\")
if file_path_parts[0] == "input": if file_path_parts[0] == "input":
print("matching input") logger.info("matching input")
file_path = os.path.join(folder_paths.get_directory_by_type("input"), *file_path_parts[1:]) file_path = os.path.join(folder_paths.get_directory_by_type("input"), *file_path_parts[1:])
elif file_path_parts[0] == "models": elif file_path_parts[0] == "models":
print("matching models") logger.info("matching models")
file_path = folder_paths.get_full_path(file_path_parts[1], os.path.join(*file_path_parts[2:])) file_path = folder_paths.get_full_path(file_path_parts[1], os.path.join(*file_path_parts[2:]))
print(file_path) logger.info(file_path)
return file_path return file_path
# Form ComfyUI Manager # Form ComfyUI Manager
async def compute_sha256_checksum(filepath): async def compute_sha256_checksum(filepath):
print("computing sha256 checksum") logger.info("computing sha256 checksum")
chunk_size = 1024 * 256 # Example: 256KB chunk_size = 1024 * 256 # Example: 256KB
filepath = get_comfyui_path_from_file_path(filepath) filepath = get_comfyui_path_from_file_path(filepath)
"""Compute the SHA256 checksum of a file, in chunks, asynchronously""" """Compute the SHA256 checksum of a file, in chunks, asynchronously"""
@@ -297,7 +514,7 @@ async def get_installed_models(request):
file_list = folder_paths.get_filename_list(key) file_list = folder_paths.get_filename_list(key)
value_json_compatible = (value[0], list(value[1]), file_list) value_json_compatible = (value[0], list(value[1]), file_list)
new_dict[key] = value_json_compatible new_dict[key] = value_json_compatible
# print(new_dict) # logger.info(new_dict)
return web.json_response(new_dict) return web.json_response(new_dict)
# This is start uploading the files to Comfy Deploy # This is start uploading the files to Comfy Deploy
@@ -307,7 +524,7 @@ async def upload_file_endpoint(request):
file_path = data.get("file_path") file_path = data.get("file_path")
print("Original file path", file_path) logger.info("Original file path", file_path)
file_path = get_comfyui_path_from_file_path(file_path) file_path = get_comfyui_path_from_file_path(file_path)
@@ -346,10 +563,9 @@ async def upload_file_endpoint(request):
if get_url: if get_url:
try: try:
async with aiohttp.ClientSession() as session:
headers = {'Authorization': f'Bearer {token}'} headers = {'Authorization': f'Bearer {token}'}
params = {'file_size': file_size, 'type': file_type} params = {'file_size': file_size, 'type': file_type}
async with session.get(get_url, params=params, headers=headers) as response: response = await async_request_with_retry('GET', get_url, params=params, headers=headers)
if response.status == 200: if response.status == 200:
content = await response.json() content = await response.json()
upload_url = content["upload_url"] upload_url = content["upload_url"]
@@ -360,7 +576,7 @@ async def upload_file_endpoint(request):
# "x-amz-acl": "public-read", # "x-amz-acl": "public-read",
"Content-Length": str(file_size) "Content-Length": str(file_size)
} }
async with session.put(upload_url, data=f, headers=headers) as upload_response: upload_response = await async_request_with_retry('PUT', upload_url, data=f, headers=headers)
if upload_response.status == 200: if upload_response.status == 200:
return web.json_response({ return web.json_response({
"message": "File uploaded successfully", "message": "File uploaded successfully",
@@ -429,7 +645,7 @@ async def get_file_hash(request):
file_hash = await compute_sha256_checksum(full_file_path) file_hash = await compute_sha256_checksum(full_file_path)
end_time = time.time() end_time = time.time()
elapsed_time = end_time - start_time elapsed_time = end_time - start_time
print(f"Cache miss -> Execution time: {elapsed_time} seconds") logger.info(f"Cache miss -> Execution time: {elapsed_time} seconds")
# Update the in-memory cache # Update the in-memory cache
file_hash_cache[full_file_path] = file_hash file_hash_cache[full_file_path] = file_hash
@@ -449,10 +665,10 @@ async def update_realtime_run_status(realtime_id: str, status_endpoint: str, sta
"run_id": realtime_id, "run_id": realtime_id,
"status": status.value, "status": status.value,
} }
if (status_endpoint is None):
return
# requests.post(status_endpoint, json=body) # requests.post(status_endpoint, json=body)
async with aiohttp.ClientSession() as session: await async_request_with_retry('POST', status_endpoint, json=body)
async with session.post(status_endpoint, json=body) as response:
pass
@server.PromptServer.instance.routes.get('/comfyui-deploy/ws') @server.PromptServer.instance.routes.get('/comfyui-deploy/ws')
async def websocket_handler(request): async def websocket_handler(request):
@@ -473,13 +689,12 @@ async def websocket_handler(request):
status_endpoint = request.rel_url.query.get('status_endpoint', None) status_endpoint = request.rel_url.query.get('status_endpoint', None)
if auth_token is not None and get_workflow_endpoint_url is not None: if auth_token is not None and get_workflow_endpoint_url is not None:
async with aiohttp.ClientSession() as session:
headers = {'Authorization': f'Bearer {auth_token}'} headers = {'Authorization': f'Bearer {auth_token}'}
async with session.get(get_workflow_endpoint_url, headers=headers) as response: response = await async_request_with_retry('GET', get_workflow_endpoint_url, headers=headers)
if response.status == 200: if response.status == 200:
workflow = await response.json() workflow = await response.json()
print("Loaded workflow version ",workflow["version"]) logger.info(f"Loaded workflow version ${workflow['version']}")
streaming_prompt_metadata[sid] = StreamingPrompt( streaming_prompt_metadata[sid] = StreamingPrompt(
workflow_api=workflow["workflow_api"], workflow_api=workflow["workflow_api"],
@@ -493,7 +708,7 @@ async def websocket_handler(request):
# await send("workflow_api", workflow_api, sid) # await send("workflow_api", workflow_api, sid)
else: else:
error_message = await response.text() error_message = await response.text()
print(f"Failed to fetch workflow endpoint. Status: {response.status}, Error: {error_message}") logger.info(f"Failed to fetch workflow endpoint. Status: {response.status}, Error: {error_message}")
# await send("error", {"message": error_message}, sid) # await send("error", {"message": error_message}, sid)
try: try:
@@ -508,10 +723,10 @@ async def websocket_handler(request):
if msg.type == aiohttp.WSMsgType.TEXT: if msg.type == aiohttp.WSMsgType.TEXT:
try: try:
data = json.loads(msg.data) data = json.loads(msg.data)
print(data) logger.info(data)
event_type = data.get('event') event_type = data.get('event')
if event_type == 'input': if event_type == 'input':
print("Got input: ", data.get("inputs")) logger.info(f"Got input: ${data.get('inputs')}")
input = data.get('inputs') input = data.get('inputs')
streaming_prompt_metadata[sid].inputs.update(input) streaming_prompt_metadata[sid].inputs.update(input)
elif event_type == 'queue_prompt': elif event_type == 'queue_prompt':
@@ -521,7 +736,7 @@ async def websocket_handler(request):
# Handle other event types # Handle other event types
pass pass
except json.JSONDecodeError: except json.JSONDecodeError:
print('Failed to decode JSON from message') logger.info('Failed to decode JSON from message')
if msg.type == aiohttp.WSMsgType.BINARY: if msg.type == aiohttp.WSMsgType.BINARY:
data = msg.data data = msg.data
@@ -530,9 +745,9 @@ async def websocket_handler(request):
image_type_code, = struct.unpack("<I", data[4:8]) image_type_code, = struct.unpack("<I", data[4:8])
input_id_bytes = data[8:32] # Extract the next 24 bytes for the input ID input_id_bytes = data[8:32] # Extract the next 24 bytes for the input ID
input_id = input_id_bytes.decode('ascii').strip() # Decode the input ID from ASCII input_id = input_id_bytes.decode('ascii').strip() # Decode the input ID from ASCII
print(event_type) logger.info(event_type)
print(image_type_code) logger.info(image_type_code)
print(input_id) logger.info(input_id)
image_data = data[32:] # The rest is the image data image_data = data[32:] # The rest is the image data
if image_type_code == 1: if image_type_code == 1:
image_type = "JPEG" image_type = "JPEG"
@@ -541,7 +756,7 @@ async def websocket_handler(request):
elif image_type_code == 3: elif image_type_code == 3:
image_type = "WEBP" image_type = "WEBP"
else: else:
print("Unknown image type code:", image_type_code) logger.info(f"Unknown image type code: ${image_type_code}")
return return
image = Image.open(BytesIO(image_data)) image = Image.open(BytesIO(image_data))
# Check if the input ID already exists and replace the input with the new one # Check if the input ID already exists and replace the input with the new one
@@ -552,14 +767,14 @@ async def websocket_handler(request):
if hasattr(existing_image, 'close'): if hasattr(existing_image, 'close'):
existing_image.close() existing_image.close()
except Exception as e: except Exception as e:
print(f"Error closing previous image for input ID {input_id}: {e}") logger.info(f"Error closing previous image for input ID {input_id}: {e}")
streaming_prompt_metadata[sid].inputs[input_id] = image streaming_prompt_metadata[sid].inputs[input_id] = image
# clear_current_prompt(sid) # clear_current_prompt(sid)
# send_prompt(sid, streaming_prompt_metadata[sid]) # send_prompt(sid, streaming_prompt_metadata[sid])
print(f"Received {image_type} image of size {image.size} with input ID {input_id}") logger.info(f"Received {image_type} image of size {image.size} with input ID {input_id}")
if msg.type == aiohttp.WSMsgType.ERROR: if msg.type == aiohttp.WSMsgType.ERROR:
print('ws connection closed with exception %s' % ws.exception()) logger.info('ws connection closed with exception %s' % ws.exception())
finally: finally:
sockets.pop(sid, None) sockets.pop(sid, None)
@@ -604,16 +819,16 @@ async def send(event, data, sid=None):
if not ws.closed: # Check if the WebSocket connection is open and not closing if not ws.closed: # Check if the WebSocket connection is open and not closing
await ws.send_json({ 'event': event, 'data': data }) await ws.send_json({ 'event': event, 'data': data })
except Exception as e: except Exception as e:
print(f"Exception: {e}") logger.info(f"Exception: {e}")
traceback.print_exc() traceback.print_exc()
logging.basicConfig(level=logging.INFO)
prompt_server = server.PromptServer.instance prompt_server = server.PromptServer.instance
send_json = prompt_server.send_json send_json = prompt_server.send_json
async def send_json_override(self, event, data, sid=None): async def send_json_override(self, event, data, sid=None):
# print("INTERNAL:", event, data, sid) # logger.info("INTERNAL:", event, data, sid)
prompt_id = data.get('prompt_id') prompt_id = data.get('prompt_id')
target_sid = sid target_sid = sid
@@ -626,8 +841,19 @@ async def send_json_override(self, event, data, sid=None):
asyncio.create_task(self.send_json_original(event, data, sid)) asyncio.create_task(self.send_json_original(event, data, sid))
]) ])
if prompt_id in comfy_message_queues:
comfy_message_queues[prompt_id].put_nowait({
"event": event,
"data": data
})
# event_emitter.emit("send_json", {
# "event": event,
# "data": data
# })
if event == 'execution_start': if event == 'execution_start':
update_run(prompt_id, Status.RUNNING) await update_run(prompt_id, Status.RUNNING)
if prompt_id in prompt_metadata: if prompt_id in prompt_metadata:
prompt_metadata[prompt_id].start_time = time.perf_counter() prompt_metadata[prompt_id].start_time = time.perf_counter()
@@ -636,12 +862,12 @@ async def send_json_override(self, event, data, sid=None):
if event == 'executing' and data.get('node') is None: if event == 'executing' and data.get('node') is None:
mark_prompt_done(prompt_id=prompt_id) mark_prompt_done(prompt_id=prompt_id)
if not have_pending_upload(prompt_id): if not have_pending_upload(prompt_id):
update_run(prompt_id, Status.SUCCESS) await update_run(prompt_id, Status.SUCCESS)
if prompt_id in prompt_metadata: if prompt_id in prompt_metadata:
current_time = time.perf_counter() current_time = time.perf_counter()
if prompt_metadata[prompt_id].start_time is not None: if prompt_metadata[prompt_id].start_time is not None:
elapsed_time = current_time - prompt_metadata[prompt_id].start_time elapsed_time = current_time - prompt_metadata[prompt_id].start_time
print(f"Elapsed time: {elapsed_time} seconds") logger.info(f"Elapsed time: {elapsed_time} seconds")
await send("elapsed_time", { await send("elapsed_time", {
"prompt_id": prompt_id, "prompt_id": prompt_id,
"elapsed_time": elapsed_time "elapsed_time": elapsed_time
@@ -656,13 +882,14 @@ async def send_json_override(self, event, data, sid=None):
prompt_metadata[prompt_id].progress.add(node) prompt_metadata[prompt_id].progress.add(node)
calculated_progress = len(prompt_metadata[prompt_id].progress) / len(prompt_metadata[prompt_id].workflow_api) calculated_progress = len(prompt_metadata[prompt_id].progress) / len(prompt_metadata[prompt_id].workflow_api)
# print("calculated_progress", calculated_progress) calculated_progress = round(calculated_progress, 2)
# logger.info("calculated_progress", calculated_progress)
if prompt_metadata[prompt_id].last_updated_node is not None and prompt_metadata[prompt_id].last_updated_node == node: if prompt_metadata[prompt_id].last_updated_node is not None and prompt_metadata[prompt_id].last_updated_node == node:
return return
prompt_metadata[prompt_id].last_updated_node = node prompt_metadata[prompt_id].last_updated_node = node
class_type = prompt_metadata[prompt_id].workflow_api[node]['class_type'] class_type = prompt_metadata[prompt_id].workflow_api[node]['class_type']
print("updating run live status", class_type) logger.info(f"At: {calculated_progress * 100}% - {class_type}")
await send("live_status", { await send("live_status", {
"prompt_id": prompt_id, "prompt_id": prompt_id,
"current_node": class_type, "current_node": class_type,
@@ -683,18 +910,19 @@ async def send_json_override(self, event, data, sid=None):
if event == 'execution_error': if event == 'execution_error':
# Careful this might not be fully awaited. # Careful this might not be fully awaited.
await update_run_with_output(prompt_id, data) await update_run_with_output(prompt_id, data)
update_run(prompt_id, Status.FAILED) await update_run(prompt_id, Status.FAILED)
# await update_run_with_output(prompt_id, data) # await update_run_with_output(prompt_id, data)
if event == 'executed' and 'node' in data and 'output' in data: if event == 'executed' and 'node' in data and 'output' in data:
print("executed", data)
if prompt_id in prompt_metadata: if prompt_id in prompt_metadata:
node = data.get('node') node = data.get('node')
class_type = prompt_metadata[prompt_id].workflow_api[node]['class_type'] class_type = prompt_metadata[prompt_id].workflow_api[node]['class_type']
print("executed", class_type) logger.info(f"Executed {class_type} {data}")
if class_type == "PreviewImage": if class_type == "PreviewImage":
print("skipping preview image") logger.info("Skipping preview image")
return return
else:
logger.info(f"Executed {data}")
await update_run_with_output(prompt_id, data.get('output'), node_id=data.get('node')) await update_run_with_output(prompt_id, data.get('output'), node_id=data.get('node'))
# await update_run_with_output(prompt_id, data.get('output'), node_id=data.get('node')) # await update_run_with_output(prompt_id, data.get('output'), node_id=data.get('node'))
@@ -710,21 +938,34 @@ async def update_run_live_status(prompt_id, live_status, calculated_progress: fl
if prompt_metadata[prompt_id].is_realtime is True: if prompt_metadata[prompt_id].is_realtime is True:
return return
print("progress", calculated_progress)
status_endpoint = prompt_metadata[prompt_id].status_endpoint status_endpoint = prompt_metadata[prompt_id].status_endpoint
if (status_endpoint is None):
return
# logger.info(f"progress {calculated_progress}")
body = { body = {
"run_id": prompt_id, "run_id": prompt_id,
"live_status": live_status, "live_status": live_status,
"progress": calculated_progress "progress": calculated_progress
} }
if prompt_id in comfy_message_queues:
comfy_message_queues[prompt_id].put_nowait({
"event": "live_status",
"data": {
"prompt_id": prompt_id,
"live_status": live_status,
"progress": calculated_progress
}
})
# requests.post(status_endpoint, json=body) # requests.post(status_endpoint, json=body)
async with aiohttp.ClientSession() as session: await async_request_with_retry('POST', status_endpoint, json=body)
async with session.post(status_endpoint, json=body) as response:
pass
def update_run(prompt_id: str, status: Status): async def update_run(prompt_id: str, status: Status):
global last_read_line_number global last_read_line_number
if prompt_id not in prompt_metadata: if prompt_id not in prompt_metadata:
@@ -747,18 +988,20 @@ def update_run(prompt_id: str, status: Status):
"run_id": prompt_id, "run_id": prompt_id,
"status": status.value, "status": status.value,
} }
print(f"Status: {status.value}") logger.info(f"Status: {status.value}")
try: try:
requests.post(status_endpoint, json=body) # requests.post(status_endpoint, json=body)
if (status_endpoint is not None):
await async_request_with_retry('POST', status_endpoint, json=body)
if cd_enable_run_log and (status == Status.SUCCESS or status == Status.FAILED): if (status_endpoint is not None) and cd_enable_run_log and (status == Status.SUCCESS or status == Status.FAILED):
try: try:
with open(comfyui_file_path, 'r') as log_file: with open(comfyui_file_path, 'r') as log_file:
# log_data = log_file.read() # log_data = log_file.read()
# Move to the last read line # Move to the last read line
all_log_data = log_file.read() # Read all log data all_log_data = log_file.read() # Read all log data
print("All log data before skipping:", all_log_data) # Log all data before skipping # logger.info("All log data before skipping: ") # Log all data before skipping
log_file.seek(0) # Reset file pointer to the beginning log_file.seek(0) # Reset file pointer to the beginning
for _ in range(last_read_line_number): for _ in range(last_read_line_number):
@@ -766,9 +1009,9 @@ def update_run(prompt_id: str, status: Status):
log_data = log_file.read() log_data = log_file.read()
# Update the last read line number # Update the last read line number
last_read_line_number += log_data.count('\n') last_read_line_number += log_data.count('\n')
print("last_read_line_number", last_read_line_number) # logger.info("last_read_line_number", last_read_line_number)
print("log_data", log_data) # logger.info("log_data", log_data)
print("log_data.count(n)", log_data.count('\n')) # logger.info("log_data.count(n)", log_data.count('\n'))
body = { body = {
"run_id": prompt_id, "run_id": prompt_id,
@@ -779,16 +1022,26 @@ def update_run(prompt_id: str, status: Status):
} }
] ]
} }
requests.post(status_endpoint, json=body)
await async_request_with_retry('POST', status_endpoint, json=body)
# requests.post(status_endpoint, json=body)
except Exception as log_error: except Exception as log_error:
print(f"Error reading log file: {log_error}") logger.info(f"Error reading log file: {log_error}")
except Exception as e: except Exception as e:
error_type = type(e).__name__ error_type = type(e).__name__
stack_trace = traceback.format_exc().strip() stack_trace = traceback.format_exc().strip()
print(f"Error occurred while updating run: {e} {stack_trace}") logger.info(f"Error occurred while updating run: {e} {stack_trace}")
finally: finally:
prompt_metadata[prompt_id].status = status prompt_metadata[prompt_id].status = status
if prompt_id in comfy_message_queues:
comfy_message_queues[prompt_id].put_nowait({
"event": "status",
"data": {
"prompt_id": prompt_id,
"status": status.value,
}
})
async def upload_file(prompt_id, filename, subfolder=None, content_type="image/png", type="output"): async def upload_file(prompt_id, filename, subfolder=None, content_type="image/png", type="output"):
@@ -806,7 +1059,7 @@ async def upload_file(prompt_id, filename, subfolder=None, content_type="image/p
output_dir = folder_paths.get_directory_by_type(type) output_dir = folder_paths.get_directory_by_type(type)
if output_dir is None: if output_dir is None:
print(filename, "Upload failed: output_dir is None") logger.info(f"{filename} Upload failed: output_dir is None")
return return
if subfolder != None: if subfolder != None:
@@ -818,7 +1071,7 @@ async def upload_file(prompt_id, filename, subfolder=None, content_type="image/p
filename = os.path.basename(filename) filename = os.path.basename(filename)
file = os.path.join(output_dir, filename) file = os.path.join(output_dir, filename)
print("uploading file", file) logger.info(f"Uploading file {file}")
file_upload_endpoint = prompt_metadata[prompt_id].file_upload_endpoint file_upload_endpoint = prompt_metadata[prompt_id].file_upload_endpoint
@@ -831,7 +1084,7 @@ async def upload_file(prompt_id, filename, subfolder=None, content_type="image/p
start_time = time.time() # Start timing here start_time = time.time() # Start timing here
result = requests.get(target_url) result = requests.get(target_url)
end_time = time.time() # End timing after the request is complete end_time = time.time() # End timing after the request is complete
print("Time taken for getting file upload endpoint: {:.2f} seconds".format(end_time - start_time)) logger.info("Time taken for getting file upload endpoint: {:.2f} seconds".format(end_time - start_time))
ok = result.json() ok = result.json()
start_time = time.time() # Start timing here start_time = time.time() # Start timing here
@@ -844,18 +1097,17 @@ async def upload_file(prompt_id, filename, subfolder=None, content_type="image/p
"Content-Length": str(len(data)), "Content-Length": str(len(data)),
} }
# response = requests.put(ok.get("url"), headers=headers, data=data) # response = requests.put(ok.get("url"), headers=headers, data=data)
async with aiohttp.ClientSession() as session: response = await async_request_with_retry('PUT', ok.get("url"), headers=headers, data=data)
async with session.put(ok.get("url"), headers=headers, data=data) as response: logger.info(f"Upload file response status: {response.status}, status text: {response.reason}")
print("Upload file response", response.status)
end_time = time.time() # End timing after the request is complete end_time = time.time() # End timing after the request is complete
print("Upload time: {:.2f} seconds".format(end_time - start_time)) logger.info("Upload time: {:.2f} seconds".format(end_time - start_time))
def have_pending_upload(prompt_id): def have_pending_upload(prompt_id):
if prompt_id in prompt_metadata and len(prompt_metadata[prompt_id].uploading_nodes) > 0: if prompt_id in prompt_metadata and len(prompt_metadata[prompt_id].uploading_nodes) > 0:
print("have pending upload ", len(prompt_metadata[prompt_id].uploading_nodes)) logger.info(f"Have pending upload {len(prompt_metadata[prompt_id].uploading_nodes)}")
return True return True
print("no pending upload") logger.info("No pending upload")
return False return False
def mark_prompt_done(prompt_id): def mark_prompt_done(prompt_id):
@@ -867,7 +1119,7 @@ def mark_prompt_done(prompt_id):
""" """
if prompt_id in prompt_metadata: if prompt_id in prompt_metadata:
prompt_metadata[prompt_id].done = True prompt_metadata[prompt_id].done = True
print("Prompt done") logger.info("Prompt done")
def is_prompt_done(prompt_id: str): def is_prompt_done(prompt_id: str):
""" """
@@ -899,8 +1151,8 @@ async def handle_error(prompt_id, data, e: Exception):
} }
} }
await update_file_status(prompt_id, data, False, have_error=True) await update_file_status(prompt_id, data, False, have_error=True)
print(body) logger.info(body)
print(f"Error occurred while uploading file: {e}") logger.info(f"Error occurred while uploading file: {e}")
# Mark the current prompt requires upload, and block it from being marked as success # Mark the current prompt requires upload, and block it from being marked as success
async def update_file_status(prompt_id: str, data, uploading, have_error=False, node_id=None): async def update_file_status(prompt_id: str, data, uploading, have_error=False, node_id=None):
@@ -913,11 +1165,11 @@ async def update_file_status(prompt_id: str, data, uploading, have_error=False,
else: else:
prompt_metadata[prompt_id].uploading_nodes.discard(node_id) prompt_metadata[prompt_id].uploading_nodes.discard(node_id)
print(prompt_metadata[prompt_id].uploading_nodes) logger.info(f"Remaining uploads: {prompt_metadata[prompt_id].uploading_nodes}")
# Update the remote status # Update the remote status
if have_error: if have_error:
update_run(prompt_id, Status.FAILED) await update_run(prompt_id, Status.FAILED)
await send("failed", { await send("failed", {
"prompt_id": prompt_id, "prompt_id": prompt_id,
}) })
@@ -926,15 +1178,15 @@ async def update_file_status(prompt_id: str, data, uploading, have_error=False,
# if there are still nodes that are uploading, then we set the status to uploading # if there are still nodes that are uploading, then we set the status to uploading
if uploading: if uploading:
if prompt_metadata[prompt_id].status != Status.UPLOADING: if prompt_metadata[prompt_id].status != Status.UPLOADING:
update_run(prompt_id, Status.UPLOADING) await update_run(prompt_id, Status.UPLOADING)
await send("uploading", { await send("uploading", {
"prompt_id": prompt_id, "prompt_id": prompt_id,
}) })
# if there are no nodes that are uploading, then we set the status to success # if there are no nodes that are uploading, then we set the status to success
elif not uploading and not have_pending_upload(prompt_id) and is_prompt_done(prompt_id=prompt_id): elif not uploading and not have_pending_upload(prompt_id) and is_prompt_done(prompt_id=prompt_id):
update_run(prompt_id, Status.SUCCESS) await update_run(prompt_id, Status.SUCCESS)
# print("Status: SUCCUSS") # logger.info("Status: SUCCUSS")
await send("success", { await send("success", {
"prompt_id": prompt_id, "prompt_id": prompt_id,
}) })
@@ -997,7 +1249,7 @@ async def update_run_with_output(prompt_id, data, node_id=None):
if have_upload_media: if have_upload_media:
try: try:
print("\nhave_upload", have_upload_media, node_id) logger.info(f"\nHave_upload {have_upload_media} Node Id: {node_id}")
if have_upload_media: if have_upload_media:
await update_file_status(prompt_id, data, True, node_id=node_id) await update_file_status(prompt_id, data, True, node_id=node_id)
@@ -1008,7 +1260,9 @@ async def update_run_with_output(prompt_id, data, node_id=None):
except Exception as e: except Exception as e:
await handle_error(prompt_id, data, e) await handle_error(prompt_id, data, e)
requests.post(status_endpoint, json=body) # requests.post(status_endpoint, json=body)
if status_endpoint is not None:
await async_request_with_retry('POST', status_endpoint, json=body)
await send('outputs_uploaded', { await send('outputs_uploaded', {
"prompt_id": prompt_id "prompt_id": prompt_id
+5 -4
View File
@@ -22,12 +22,13 @@ class StreamingPrompt(BaseModel):
auth_token: str auth_token: str
inputs: dict[str, Union[str, bytes, Image.Image]] inputs: dict[str, Union[str, bytes, Image.Image]]
running_prompt_ids: set[str] = set() running_prompt_ids: set[str] = set()
status_endpoint: str status_endpoint: Optional[str]
file_upload_endpoint: str file_upload_endpoint: Optional[str]
class SimplePrompt(BaseModel): class SimplePrompt(BaseModel):
status_endpoint: str status_endpoint: Optional[str]
file_upload_endpoint: str file_upload_endpoint: Optional[str]
workflow_api: dict workflow_api: dict
status: Status = Status.NOT_STARTED status: Status = Status.NOT_STARTED
progress: set = set() progress: set = set()
+1
View File
@@ -2,3 +2,4 @@ aiofiles
pydantic pydantic
opencv-python opencv-python
imageio-ffmpeg imageio-ffmpeg
# logfire
+263 -47
View File
@@ -13,6 +13,84 @@ function sendEventToCD(event, data) {
window.parent.postMessage(JSON.stringify(message), "*"); window.parent.postMessage(JSON.stringify(message), "*");
} }
function dispatchAPIEventData(data) {
const msg = JSON.parse(data);
// Custom parse error
if (msg.error) {
let message = msg.error.message;
if (msg.error.details) message += ": " + msg.error.details;
for (const [nodeID, nodeError] of Object.entries(msg.node_errors)) {
message += "\n" + nodeError.class_type + ":";
for (const errorReason of nodeError.errors) {
message +=
"\n - " +
errorReason.message +
": " +
errorReason.details;
}
}
app.ui.dialog.show(message);
if (msg.node_errors) {
app.lastNodeErrors = msg.node_errors;
app.canvas.draw(true, true);
}
}
switch (msg.event) {
case "error":
break;
case "status":
if (msg.data.sid) {
// this.clientId = msg.data.sid;
// window.name = this.clientId; // use window name so it isnt reused when duplicating tabs
// sessionStorage.setItem("clientId", this.clientId); // store in session storage so duplicate tab can load correct workflow
}
api.dispatchEvent(
new CustomEvent("status", { detail: msg.data.status })
);
break;
case "progress":
api.dispatchEvent(
new CustomEvent("progress", { detail: msg.data })
);
break;
case "executing":
api.dispatchEvent(
new CustomEvent("executing", { detail: msg.data.node })
);
break;
case "executed":
api.dispatchEvent(
new CustomEvent("executed", { detail: msg.data })
);
break;
case "execution_start":
api.dispatchEvent(
new CustomEvent("execution_start", { detail: msg.data })
);
break;
case "execution_error":
api.dispatchEvent(
new CustomEvent("execution_error", { detail: msg.data })
);
break;
case "execution_cached":
api.dispatchEvent(
new CustomEvent("execution_cached", { detail: msg.data })
);
break;
default:
api.dispatchEvent(new CustomEvent(msg.type, { detail: msg.data }));
// default:
// if (this.#registered.has(msg.type)) {
// } else {
// throw new Error(`Unknown message type ${msg.type}`);
// }
}
}
/** @typedef {import('../../../web/types/comfy.js').ComfyExtension} ComfyExtension*/ /** @typedef {import('../../../web/types/comfy.js').ComfyExtension} ComfyExtension*/
/** @type {ComfyExtension} */ /** @type {ComfyExtension} */
const ext = { const ext = {
@@ -33,8 +111,7 @@ const ext = {
sendEventToCD("cd_plugin_onInit"); sendEventToCD("cd_plugin_onInit");
app.queuePrompt = ((originalFunction) => app.queuePrompt = ((originalFunction) => async () => {
async () => {
// const prompt = await app.graphToPrompt(); // const prompt = await app.graphToPrompt();
sendEventToCD("cd_plugin_onQueuePromptTrigger"); sendEventToCD("cd_plugin_onQueuePromptTrigger");
})(app.queuePrompt); })(app.queuePrompt);
@@ -75,11 +152,13 @@ const ext = {
} }
if (!workflow_version_id) { if (!workflow_version_id) {
console.error("No workflow_version_id provided in query parameters."); console.error(
"No workflow_version_id provided in query parameters."
);
} else { } else {
loadingDialog.showLoading( loadingDialog.showLoading(
"Loading workflow from " + org_display, "Loading workflow from " + org_display,
"Please wait...", "Please wait..."
); );
fetch(endpoint + "/api/workflow-version/" + workflow_version_id, { fetch(endpoint + "/api/workflow-version/" + workflow_version_id, {
method: "GET", method: "GET",
@@ -92,7 +171,10 @@ const ext = {
const data = await res.json(); const data = await res.json();
const { workflow, workflow_id, error } = data; const { workflow, workflow_id, error } = data;
if (error) { if (error) {
infoDialog.showMessage("Unable to load this workflow", error); infoDialog.showMessage(
"Unable to load this workflow",
error
);
return; return;
} }
@@ -115,7 +197,7 @@ const ext = {
window.history.replaceState( window.history.replaceState(
{}, {},
document.title, document.title,
window.location.pathname, window.location.pathname
); );
}); });
} }
@@ -138,22 +220,37 @@ const ext = {
ComfyWidgets.STRING( ComfyWidgets.STRING(
this, this,
"workflow_name", "workflow_name",
["", { default: this.properties.workflow_name, multiline: false }], [
app, "",
{
default: this.properties.workflow_name,
multiline: false,
},
],
app
); );
ComfyWidgets.STRING( ComfyWidgets.STRING(
this, this,
"workflow_id", "workflow_id",
["", { default: this.properties.workflow_id, multiline: false }], [
app, "",
{
default: this.properties.workflow_id,
multiline: false,
},
],
app
); );
ComfyWidgets.STRING( ComfyWidgets.STRING(
this, this,
"version", "version",
["", { default: this.properties.version, multiline: false }], [
app, "",
{ default: this.properties.version, multiline: false },
],
app
); );
// this.widgets.forEach((w) => { // this.widgets.forEach((w) => {
@@ -180,7 +277,7 @@ const ext = {
title_mode: LiteGraph.NORMAL_TITLE, title_mode: LiteGraph.NORMAL_TITLE,
title: "Comfy Deploy", title: "Comfy Deploy",
collapsable: true, collapsable: true,
}), })
); );
ComfyDeploy.category = "deploy"; ComfyDeploy.category = "deploy";
@@ -190,32 +287,128 @@ const ext = {
// const graphCanvas = document.getElementById("graph-canvas"); // const graphCanvas = document.getElementById("graph-canvas");
window.addEventListener("message", async (event) => { window.addEventListener("message", async (event) => {
// console.log("message", event);
try { try {
const message = JSON.parse(event.data); const message = JSON.parse(event.data);
if (message.type === "graph_load") { if (message.type === "graph_load") {
const comfyUIWorkflow = message.data; const comfyUIWorkflow = message.data;
console.log("recieved: ", comfyUIWorkflow); // console.log("recieved: ", comfyUIWorkflow);
// Assuming there's a method to load the workflow data into the ComfyUI // Assuming there's a method to load the workflow data into the ComfyUI
// This part of the code would depend on how the ComfyUI expects to receive and process the workflow data // This part of the code would depend on how the ComfyUI expects to receive and process the workflow data
// For demonstration, let's assume there's a loadWorkflow method in the ComfyUI API // For demonstration, let's assume there's a loadWorkflow method in the ComfyUI API
if (comfyUIWorkflow && app && app.loadGraphData) { if (comfyUIWorkflow && app && app.loadGraphData) {
console.log("loadGraphData");
app.loadGraphData(comfyUIWorkflow); app.loadGraphData(comfyUIWorkflow);
} }
} else if (message.type === "deploy") { } else if (message.type === "deploy") {
// deployWorkflow(); // deployWorkflow();
const prompt = await app.graphToPrompt(); const prompt = await app.graphToPrompt();
// api.handlePromptGenerated(prompt);
sendEventToCD("cd_plugin_onDeployChanges", prompt); sendEventToCD("cd_plugin_onDeployChanges", prompt);
} else if (message.type === "queue_prompt") { } else if (message.type === "queue_prompt") {
const prompt = await app.graphToPrompt(); const prompt = await app.graphToPrompt();
sendEventToCD("cd_plugin_onQueuePrompt", prompt); if (typeof api.handlePromptGenerated === "function") {
api.handlePromptGenerated(prompt);
} else {
console.warn(
"api.handlePromptGenerated is not a function"
);
} }
sendEventToCD("cd_plugin_onQueuePrompt", prompt);
} else if (message.type === "get_prompt") {
const prompt = await app.graphToPrompt();
sendEventToCD("cd_plugin_onGetPrompt", prompt);
} else if (message.type === "event") {
dispatchAPIEventData(message.data);
} else if (message.type === "add_node") {
console.log("add node", message.data);
app.graph.beforeChange();
var node = LiteGraph.createNode(message.data.type);
node.configure({
widgets_values: message.data.widgets_values,
});
console.log("node", node);
const graphMouse = app.canvas.graph_mouse;
node.pos = [graphMouse[0], graphMouse[1]];
app.graph.add(node);
app.graph.afterChange();
} else if (message.type === "zoom_to_node") {
const nodeId = message.data.nodeId;
const position = message.data.position;
const node = app.graph.getNodeById(nodeId);
if (!node) return;
const canvas = app.canvas;
const targetScale = 1;
const targetOffsetX =
canvas.canvas.width / 4 -
position[0] -
node.size[0] / 2;
const targetOffsetY =
canvas.canvas.height / 4 -
position[1] -
node.size[1] / 2;
const startScale = canvas.ds.scale;
const startOffsetX = canvas.ds.offset[0];
const startOffsetY = canvas.ds.offset[1];
const duration = 400; // Animation duration in milliseconds
const startTime = Date.now();
function easeOutCubic(t) {
return 1 - Math.pow(1 - t, 3);
}
function lerp(start, end, t) {
return start * (1 - t) + end * t;
}
function animate() {
const currentTime = Date.now();
const elapsedTime = currentTime - startTime;
const t = Math.min(elapsedTime / duration, 1);
const easedT = easeOutCubic(t);
const currentScale = lerp(
startScale,
targetScale,
easedT
);
const currentOffsetX = lerp(
startOffsetX,
targetOffsetX,
easedT
);
const currentOffsetY = lerp(
startOffsetY,
targetOffsetY,
easedT
);
canvas.setZoom(currentScale);
canvas.ds.offset = [currentOffsetX, currentOffsetY];
canvas.draw(true, true);
if (t < 1) {
requestAnimationFrame(animate);
}
}
animate();
}
// else if (message.type === "refresh") {
// sendEventToCD("cd_plugin_onRefresh");
// }
} catch (error) { } catch (error) {
// console.error("Error processing message:", error); // console.error("Error processing message:", error);
} }
// if (!event.data.flow || Object.entries(event.data.flow).length <= 0)
// return;
// updateBlendshapesPrompts(event.data.flow);
}); });
api.addEventListener("executed", (evt) => { api.addEventListener("executed", (evt) => {
@@ -249,7 +442,7 @@ const ext = {
function showError(title, message) { function showError(title, message) {
infoDialog.show( infoDialog.show(
`<h3 style="margin: 0px; color: red;">${title}</h3><br><span>${message}</span> `, `<h3 style="margin: 0px; color: red;">${title}</h3><br><span>${message}</span> `
); );
} }
@@ -357,7 +550,7 @@ async function deployWorkflow() {
if (deployMeta.length == 0) { if (deployMeta.length == 0) {
const text = await inputDialog.input( const text = await inputDialog.input(
"Create your deployment", "Create your deployment",
"Workflow name", "Workflow name"
); );
if (!text) return; if (!text) return;
console.log(text); console.log(text);
@@ -398,7 +591,7 @@ async function deployWorkflow() {
<input id="reuse-hash" type="checkbox" checked>Reuse hash from last version</input> <input id="reuse-hash" type="checkbox" checked>Reuse hash from last version</input>
</label> </label>
</div> </div>
`, `
); );
if (!ok) return; if (!ok) return;
@@ -417,7 +610,7 @@ async function deployWorkflow() {
if (!snapshot) { if (!snapshot) {
showError( showError(
"Error when deploying", "Error when deploying",
"Unable to generate snapshot, please install ComfyUI Manager", "Unable to generate snapshot, please install ComfyUI Manager"
); );
return; return;
} }
@@ -438,7 +631,7 @@ async function deployWorkflow() {
"Content-Type": "application/json", "Content-Type": "application/json",
Authorization: "Bearer " + apiKey, Authorization: "Bearer " + apiKey,
}, },
}, }
) )
.then((x) => x.json()) .then((x) => x.json())
.catch(() => { .catch(() => {
@@ -457,7 +650,7 @@ async function deployWorkflow() {
// Match previous hash for models // Match previous hash for models
if (reuseHash && existing_workflow?.dependencies?.models) { if (reuseHash && existing_workflow?.dependencies?.models) {
const previousModelHash = Object.entries( const previousModelHash = Object.entries(
existing_workflow?.dependencies?.models, existing_workflow?.dependencies?.models
).flatMap(([key, value]) => { ).flatMap(([key, value]) => {
return Object.values(value).map((x) => ({ return Object.values(value).map((x) => ({
...x, ...x,
@@ -479,7 +672,9 @@ async function deployWorkflow() {
console.log(file); console.log(file);
loadingDialog.showLoading("Generating hash", file); loadingDialog.showLoading("Generating hash", file);
const hash = await fetch( const hash = await fetch(
`/comfyui-deploy/get-file-hash?file_path=${encodeURIComponent(file)}`, `/comfyui-deploy/get-file-hash?file_path=${encodeURIComponent(
file
)}`
).then((x) => x.json()); ).then((x) => x.json());
loadingDialog.showLoading("Generating hash", file); loadingDialog.showLoading("Generating hash", file);
console.log(hash); console.log(hash);
@@ -489,18 +684,24 @@ async function deployWorkflow() {
console.log("Uploading ", file); console.log("Uploading ", file);
loadingDialog.showLoading("Uploading file", file); loadingDialog.showLoading("Uploading file", file);
try { try {
const { download_url } = await fetch(`/comfyui-deploy/upload-file`, { const { download_url } = await fetch(
`/comfyui-deploy/upload-file`,
{
method: "POST", method: "POST",
body: JSON.stringify({ body: JSON.stringify({
file_path: file, file_path: file,
token: apiKey, token: apiKey,
url: endpoint + "/api/upload-url", url: endpoint + "/api/upload-url",
}), }),
}) }
)
.then((x) => x.json()) .then((x) => x.json())
.catch(() => { .catch(() => {
loadingDialog.close(); loadingDialog.close();
confirmDialog.confirm("Error", "Unable to upload file " + file); confirmDialog.confirm(
"Error",
"Unable to upload file " + file
);
}); });
loadingDialog.showLoading("Uploaded file", file); loadingDialog.showLoading("Uploaded file", file);
console.log(download_url); console.log(download_url);
@@ -537,8 +738,8 @@ async function deployWorkflow() {
<iframe <iframe
style="z-index: 10; min-width: 600px; max-width: 1024px; min-height: 600px; border: none; background-color: transparent;" style="z-index: 10; min-width: 600px; max-width: 1024px; min-height: 600px; border: none; background-color: transparent;"
src="https://www.comfydeploy.com/dependency-graph?deps=${encodeURIComponent( src="https://www.comfydeploy.com/dependency-graph?deps=${encodeURIComponent(
JSON.stringify(deps), JSON.stringify(deps)
)}" />`, )}" />`
// createDynamicUIHtml(deps), // createDynamicUIHtml(deps),
); );
if (!depsOk) return; if (!depsOk) return;
@@ -599,7 +800,7 @@ async function deployWorkflow() {
graph.change(); graph.change();
infoDialog.show( infoDialog.show(
`<span style="color:green;">Deployed successfully!</span> <a style="color:white;" target="_blank" href=${endpoint}/workflows/${data.workflow_id}>-> View here</a> <br/> <br/> Workflow ID: ${data.workflow_id} <br/> Workflow Name: ${workflow_name} <br/> Workflow Version: ${data.version} <br/>`, `<span style="color:green;">Deployed successfully!</span> <a style="color:white;" target="_blank" href=${endpoint}/workflows/${data.workflow_id}>-> View here</a> <br/> <br/> Workflow ID: ${data.workflow_id} <br/> Workflow Name: ${workflow_name} <br/> Workflow Version: ${data.version} <br/>`
); );
setTimeout(() => { setTimeout(() => {
@@ -796,17 +997,22 @@ export class InputDialog extends InfoDialog {
type: "button", type: "button",
textContent: "Save", textContent: "Save",
onclick: () => { onclick: () => {
const input = this.textElement.querySelector("#input").value; const input =
this.textElement.querySelector("#input").value;
if (input.trim() === "") { if (input.trim() === "") {
showError("Input validation", "Input cannot be empty"); showError(
"Input validation",
"Input cannot be empty"
);
} else { } else {
this.callback?.(input); this.callback?.(input);
this.close(); this.close();
this.textElement.querySelector("#input").value = ""; this.textElement.querySelector("#input").value =
"";
} }
}, },
}), }),
], ]
), ),
]; ];
} }
@@ -867,7 +1073,7 @@ export class ConfirmDialog extends InfoDialog {
this.close(); this.close();
}, },
}), }),
], ]
), ),
]; ];
} }
@@ -924,7 +1130,7 @@ function getData(environment) {
function saveData(data) { function saveData(data) {
localStorage.setItem( localStorage.setItem(
"comfy_deploy_env_data_" + data.environment, "comfy_deploy_env_data_" + data.environment,
JSON.stringify(data), JSON.stringify(data)
); );
} }
@@ -939,7 +1145,9 @@ export class ConfigDialog extends ComfyDialog {
this.element.style.paddingBottom = "20px"; this.element.style.paddingBottom = "20px";
this.container = document.createElement("div"); this.container = document.createElement("div");
this.element.querySelector(".comfy-modal-content").prepend(this.container); this.element
.querySelector(".comfy-modal-content")
.prepend(this.container);
} }
createButtons() { createButtons() {
@@ -973,7 +1181,7 @@ export class ConfigDialog extends ComfyDialog {
this.close(); this.close();
}, },
}), }),
], ]
), ),
]; ];
} }
@@ -985,7 +1193,8 @@ export class ConfigDialog extends ComfyDialog {
} }
save(api_key, displayName) { save(api_key, displayName) {
const deployOption = this.container.querySelector("#deployOption").value; const deployOption =
this.container.querySelector("#deployOption").value;
localStorage.setItem("comfy_deploy_env", deployOption); localStorage.setItem("comfy_deploy_env", deployOption);
const endpoint = this.container.querySelector("#endpoint").value; const endpoint = this.container.querySelector("#endpoint").value;
@@ -1017,8 +1226,12 @@ export class ConfigDialog extends ComfyDialog {
<h3 style="margin: 0px;">Comfy Deploy Config</h3> <h3 style="margin: 0px;">Comfy Deploy Config</h3>
<label style="color: white; width: 100%;"> <label style="color: white; width: 100%;">
<select id="deployOption" style="margin: 8px 0px; width: 100%; height:30px; box-sizing: border-box;" > <select id="deployOption" style="margin: 8px 0px; width: 100%; height:30px; box-sizing: border-box;" >
<option value="cloud" ${data.environment === "cloud" ? "selected" : ""}>Cloud</option> <option value="cloud" ${
<option value="local" ${data.environment === "local" ? "selected" : ""}>Local</option> data.environment === "cloud" ? "selected" : ""
}>Cloud</option>
<option value="local" ${
data.environment === "local" ? "selected" : ""
}>Local</option>
</select> </select>
</label> </label>
<label style="color: white; width: 100%;"> <label style="color: white; width: 100%;">
@@ -1036,7 +1249,9 @@ export class ConfigDialog extends ComfyDialog {
}"> }">
<button id="loginButton" style="margin-top: 8px; width: 100%; height:40px; box-sizing: border-box; padding: 0px 6px;"> <button id="loginButton" style="margin-top: 8px; width: 100%; height:40px; box-sizing: border-box; padding: 0px 6px;">
${ ${
data.apiKey ? "Re-login with ComfyDeploy" : "Login with ComfyDeploy" data.apiKey
? "Re-login with ComfyDeploy"
: "Login with ComfyDeploy"
} }
</button> </button>
</div> </div>
@@ -1057,7 +1272,7 @@ export class ConfigDialog extends ComfyDialog {
clearInterval(poll); clearInterval(poll);
infoDialog.showMessage( infoDialog.showMessage(
"Timeout", "Timeout",
"Wait too long for the response, please try re-login", "Wait too long for the response, please try re-login"
); );
}, 30000); // Stop polling after 30 seconds }, 30000); // Stop polling after 30 seconds
@@ -1068,14 +1283,15 @@ export class ConfigDialog extends ComfyDialog {
if (json.api_key) { if (json.api_key) {
this.save(json.api_key, json.name); this.save(json.api_key, json.name);
this.close(); this.close();
this.container.querySelector("#apiKey").value = json.api_key; this.container.querySelector("#apiKey").value =
json.api_key;
// infoDialog.show(); // infoDialog.show();
clearInterval(this.poll); clearInterval(this.poll);
clearTimeout(this.timeout); clearTimeout(this.timeout);
// Refresh dialog // Refresh dialog
const a = await confirmDialog.confirm( const a = await confirmDialog.confirm(
"Authenticated", "Authenticated",
`<div>You will be able to upload workflow to <button style="font-size: 18px; width: fit;">${json.name}</button></div>`, `<div>You will be able to upload workflow to <button style="font-size: 18px; width: fit;">${json.name}</button></div>`
); );
configDialog.show(); configDialog.show();
} }
+1 -1
View File
@@ -74,7 +74,7 @@
"mitata": "^0.1.6", "mitata": "^0.1.6",
"ms": "^2.1.3", "ms": "^2.1.3",
"nanoid": "^5.0.4", "nanoid": "^5.0.4",
"next": "14.1", "next": "14.2",
"next-plausible": "^3.12.0", "next-plausible": "^3.12.0",
"next-themes": "^0.2.1", "next-themes": "^0.2.1",
"next-usequerystate": "^1.13.2", "next-usequerystate": "^1.13.2",
+3 -1
View File
@@ -51,7 +51,9 @@ const createRunRoute = createRoute({
export const registerCreateRunRoute = (app: App) => { export const registerCreateRunRoute = (app: App) => {
app.openapi(createRunRoute, async (c) => { app.openapi(createRunRoute, async (c) => {
const data = c.req.valid("json"); const data = c.req.valid("json");
const origin = new URL(c.req.url).origin; const proto = c.req.headers.get('x-forwarded-proto') || "http";
const host = c.req.headers.get('x-forwarded-host') || c.req.headers.get('host');
const origin = `${proto}://${host}` || new URL(c.req.url).origin;
const apiKeyTokenData = c.get("apiKeyTokenData")!; const apiKeyTokenData = c.get("apiKeyTokenData")!;
const { deployment_id, inputs } = data; const { deployment_id, inputs } = data;
+1 -1
View File
@@ -102,7 +102,7 @@ export const createRun = withServerPromise(
let prompt_id: string | undefined = undefined; let prompt_id: string | undefined = undefined;
const shareData = { const shareData = {
workflow_api: workflow_api, workflow_api_raw: workflow_api,
status_endpoint: `${origin}/api/update-run`, status_endpoint: `${origin}/api/update-run`,
file_upload_endpoint: `${origin}/api/file-upload`, file_upload_endpoint: `${origin}/api/file-upload`,
}; };