diff --git a/.prettierrc.json b/.prettierrc.json index de753c5..582f76b 100644 --- a/.prettierrc.json +++ b/.prettierrc.json @@ -1,3 +1,5 @@ { - "printWidth": 100 + "printWidth": 100, + "bracketSpacing": false, + "bracketSameLine": true } diff --git a/__build__.py b/__build__.py index f3cb685..c7832e6 100644 --- a/__build__.py +++ b/__build__.py @@ -15,6 +15,7 @@ start = time.time() parser = argparse.ArgumentParser() parser.add_argument("-t", "--with-tests", default=False, action="store_true") +parser.add_argument("-f", "--fix", default=False, action="store_true") args = parser.parse_args() THIS_DIR = os.path.dirname(os.path.abspath(__file__)) @@ -27,12 +28,16 @@ def log_step(msg=None, status=None): """ Logs a step keeping track of timing and initial msg. """ global step_msg # pylint: disable=W0601 global step_start # pylint: disable=W0601 + global step_warns # pylint: disable=W0601 if msg: tag = f'{COLORS["YELLOW"]}[ Notice ]' if status == 'Notice' else f'{COLORS["RESET"]}[Starting]' step_msg = f'โ–ป {tag}{COLORS["RESET"]} {msg}...' step_start = time.time() + step_warns = [] print(step_msg, end="\r") elif status: + if status != 'Error': + status = "Warn" if len(step_warns) > 0 else status step_time = round(time.time() - step_start, 3) if status == 'Error': status_msg = f'{COLORS["RED"]}โคซ {status}{COLORS["RESET"]}' @@ -41,8 +46,26 @@ def log_step(msg=None, status=None): else: status_msg = f'{COLORS["BRIGHT_GREEN"]}๐Ÿ—ธ {status}{COLORS["RESET"]}' print(f'{step_msg.ljust(64, ".")} {status_msg} ({step_time}s)') + for warning in step_warns: + print(warning) +if args.fix: + tss = glob(os.path.join(DIR_SRC_WEB, "**", "*.ts"), recursive=True) + log_step(msg=f'Fixing {len(tss)} ts files') + for ts in tss: + with open(ts, 'r', encoding="utf-8") as f: + content = f.read() + # (\s*from\s*['"](?!.*[.]js['"]).*?)(['"];) in vscode. + content, n = re.subn(r'(\s*from [\'"](?!.*[.]js[\'"]).*?)([\'"];)', '\\1.js\\2', content) + if n > 0: + filename = os.path.basename(ts) + step_warns.append( + f' - {filename} has {n} import{"s" if n > 1 else ""} that do not end in ".js"') + with open(ts, 'w', encoding="utf-8") as f: + f.write(content) + log_step(status="Done") + log_step(msg='Copying web directory') rmtree(DIR_WEB) copytree(DIR_SRC_WEB, DIR_WEB, ignore=ignore_patterns("typings*", "*.ts", "*.scss")) @@ -88,7 +111,6 @@ log_step(status="Done") # "../../rgthree/common" (which we map correctly in rgthree_server.py). log_step(msg='Cleaning Imports') js_files = glob(os.path.join(DIR_WEB, '**', '*.js'), recursive=True) -warns = [] for file in js_files: rel_path = file.replace(f'{DIR_WEB}/', "") with open(file, 'r', encoding="utf-8") as f: @@ -100,14 +122,13 @@ for file in js_files: else: filedata = re.sub(r'(from\s+["\'])rgthree/', f'\\1{"../" * num}', filedata) filedata = re.sub(r'(from\s+["\'])scripts/', f'\\1{"../" * (num + 1)}scripts/', filedata) - filedata, n = re.subn(r'(import .*from [\'"](?!.*[.]js[\'"]).*?)([\'"];)', '\\1.js\\2', filedata) + filedata, n = re.subn(r'(\s*from [\'"](?!.*[.]js[\'"]).*?)([\'"];)', '\\1.js\\2', filedata) if n > 0: filename = os.path.basename(file) - warns.append(f' - {filename} has {n} import{"s" if n > 1 else ""} that do not end in ".js"') + step_warns.append( + f' - {filename} has {n} import{"s" if n > 1 else ""} that do not end in ".js"') with open(file, 'w', encoding="utf-8") as f: f.write(filedata) -log_step(status="Warn" if len(warns) > 0 else "Done") -for warn in warns: - print(warn) +log_step(status="Done") print(f'Finished all in {round(time.time() - start, 3)}s') diff --git a/__init__.py b/__init__.py index fc6a48f..b5da9de 100644 --- a/__init__.py +++ b/__init__.py @@ -29,6 +29,8 @@ from .py.power_prompt import RgthreePowerPrompt from .py.power_prompt_simple import RgthreePowerPromptSimple from .py.image_inset_crop import RgthreeImageInsetCrop from .py.context_big import RgthreeBigContext +from .py.dynamic_context import RgthreeDynamicContext +from .py.dynamic_context_switch import RgthreeDynamicContextSwitch from .py.ksampler_config import RgthreeKSamplerConfig from .py.sdxl_power_prompt_postive import RgthreeSDXLPowerPromptPositive from .py.sdxl_power_prompt_simple import RgthreeSDXLPowerPromptSimple @@ -61,6 +63,10 @@ NODE_CLASS_MAPPINGS = { RgthreePowerLoraLoader.NAME: RgthreePowerLoraLoader, } +if get_config_value('unreleased.dynamic_context.enabled') is True: + NODE_CLASS_MAPPINGS[RgthreeDynamicContext.NAME] = RgthreeDynamicContext + NODE_CLASS_MAPPINGS[RgthreeDynamicContextSwitch.NAME] = RgthreeDynamicContextSwitch + # WEB_DIRECTORY is the comfyui nodes directory that ComfyUI will link and auto-load. WEB_DIRECTORY = "./web/comfyui" diff --git a/py/dynamic_context.py b/py/dynamic_context.py new file mode 100644 index 0000000..411210d --- /dev/null +++ b/py/dynamic_context.py @@ -0,0 +1,56 @@ +"""The Dynamic Context node.""" +from mimetypes import add_type +from .constants import get_category, get_name +from .utils import ByPassTypeTuple, FlexibleOptionalInputType + + +class RgthreeDynamicContext: + """The Dynamic Context node. + + Similar to the static Context and Context Big nodes, this allows users to add any number and + variety of inputs to a Dynamic Context node, and return the outputs by key name. + """ + + NAME = get_name("Dynamic Context") + CATEGORY = get_category() + + @classmethod + def INPUT_TYPES(cls): # pylint: disable = invalid-name,missing-function-docstring + return { + "required": {}, + "optional": FlexibleOptionalInputType(add_type), + "hidden": {}, + } + + RETURN_TYPES = ByPassTypeTuple(("RGTHREE_DYNAMIC_CONTEXT",)) + RETURN_NAMES = ByPassTypeTuple(("CONTEXT",)) + FUNCTION = "main" + + def main(self, **kwargs): + """Creates a new context from the provided data, with an optional base ctx to start. + + This node takes a list of named inputs that are the named keys (with an optional "+ " prefix) + which are to be stored within the ctx dict as well as a list of keys contained in `output_keys` + to determine the list of output data. + """ + + base_ctx = kwargs.get('base_ctx', None) + output_keys = kwargs.get('output_keys', None) + + new_ctx = base_ctx.copy() if base_ctx is not None else {} + + for key_raw, value in kwargs.items(): + if key_raw in ['base_ctx', 'output_keys']: + continue + key = key_raw.upper() + if key.startswith('+ '): + key = key[2:] + new_ctx[key] = value + + print(new_ctx) + + res = [new_ctx] + output_keys = output_keys.split(',') if output_keys is not None else [] + for key in output_keys: + res.append(new_ctx[key] if key in new_ctx else None) + return tuple(res) diff --git a/py/dynamic_context_switch.py b/py/dynamic_context_switch.py new file mode 100644 index 0000000..3d1b9e6 --- /dev/null +++ b/py/dynamic_context_switch.py @@ -0,0 +1,39 @@ +"""The original Context Switch.""" +from .constants import get_category, get_name +from .context_utils import is_context_empty +from .utils import ByPassTypeTuple, FlexibleOptionalInputType + + +class RgthreeDynamicContextSwitch: + """The initial Context Switch node.""" + + NAME = get_name("Dynamic Context Switch") + CATEGORY = get_category() + + @classmethod + def INPUT_TYPES(cls): # pylint: disable = invalid-name, missing-function-docstring + return { + "required": {}, + "optional": FlexibleOptionalInputType("RGTHREE_DYNAMIC_CONTEXT"), + } + + RETURN_TYPES = ByPassTypeTuple(("RGTHREE_DYNAMIC_CONTEXT",)) + RETURN_NAMES = ByPassTypeTuple(("CONTEXT",)) + FUNCTION = "switch" + + def switch(self, **kwargs): + """Chooses the first non-empty Context to output.""" + + output_keys = kwargs.get('output_keys', None) + + ctx = None + for key, value in kwargs.items(): + if key.startswith('ctx_') and not is_context_empty(value): + ctx = value + break + + res = [ctx] + output_keys = output_keys.split(',') if output_keys is not None else [] + for key in output_keys: + res.append(ctx[key] if ctx is not None and key in ctx else None) + return tuple(res) diff --git a/py/power_prompt.py b/py/power_prompt.py index d03aa4c..79426f0 100644 --- a/py/power_prompt.py +++ b/py/power_prompt.py @@ -81,8 +81,9 @@ class RgthreePowerPrompt: log_node_success(NODE_NAME, f'Loaded "{lora["lora"]}" from prompt') log_node_info(NODE_NAME, f'{len(loras)} Loras processed; stripping tags for TEXT output.') elif ' 1 and len(match[1]) else 1.0) - if strength == 0 and not silent: - log_node_info(log_node, f'Skipping "{tag_path}" with strength of zero') + if strength == 0: + if not silent: + log_node_info(log_node, f'Skipping "{tag_path}" with strength of zero') + skipped_loras.append({'lora': tag_path, 'strength': strength}) continue lora_path = get_lora_by_filename(tag_path, lora_paths, log_node=None if silent else log_node) diff --git a/py/utils.py b/py/utils.py index a2438c3..7fd0f95 100644 --- a/py/utils.py +++ b/py/utils.py @@ -100,3 +100,14 @@ def path_exists(path): if path is not None: return os.path.exists(path) return False + + +class ByPassTypeTuple(tuple): + """A special class that will return additional "AnyType" strings beyond defined values. + Credit to Trung0246 + """ + + def __getitem__(self, index): + if index > len(self) - 1: + return AnyType("*") + return super().__getitem__(index) diff --git a/src_web/comfyui/constants.ts b/src_web/comfyui/constants.ts index c642574..51bbace 100644 --- a/src_web/comfyui/constants.ts +++ b/src_web/comfyui/constants.ts @@ -1,4 +1,4 @@ -import { SERVICE as CONFIG_SERVICE } from "./services/config_service.js"; +import {SERVICE as CONFIG_SERVICE} from "./services/config_service.js"; export function addRgthree(str: string) { return str + " (rgthree)"; @@ -16,6 +16,8 @@ export const NodeTypesString = { CONTEXT_SWITCH_BIG: addRgthree("Context Switch Big"), CONTEXT_MERGE: addRgthree("Context Merge"), CONTEXT_MERGE_BIG: addRgthree("Context Merge Big"), + DYNAMIC_CONTEXT: addRgthree("Dynamic Context"), + DYNAMIC_CONTEXT_SWITCH: addRgthree("Dynamic Context Switch"), DISPLAY_ANY: addRgthree("Display Any"), NODE_MODE_RELAY: addRgthree("Mute / Bypass Relay"), NODE_MODE_REPEATER: addRgthree("Mute / Bypass Repeater"), @@ -47,5 +49,14 @@ export const NodeTypesString = { export function getNodeTypeStrings() { return Object.values(NodeTypesString) .map((i) => stripRgthree(i)) + .filter((i) => { + if ( + i.startsWith("Dynamic Context") && + !CONFIG_SERVICE.getConfigValue("unreleased.dynamic_context.enabled") + ) { + return false; + } + return true; + }) .sort(); } diff --git a/src_web/comfyui/context.ts b/src_web/comfyui/context.ts index 67ae688..b87583a 100644 --- a/src_web/comfyui/context.ts +++ b/src_web/comfyui/context.ts @@ -85,7 +85,7 @@ function findMatchingIndexByTypeOrName( /** * A Base Context node for other context based nodes to extend. */ -class BaseContextNode extends RgthreeBaseServerNode { +export class BaseContextNode extends RgthreeBaseServerNode { constructor(title: string) { super(title); } diff --git a/src_web/comfyui/dynamic_context.ts b/src_web/comfyui/dynamic_context.ts new file mode 100644 index 0000000..f637647 --- /dev/null +++ b/src_web/comfyui/dynamic_context.ts @@ -0,0 +1,297 @@ +import {app} from "scripts/app.js"; +import { + IoDirection, + followConnectionUntilType, + getConnectedInputInfosAndFilterPassThroughs, +} from "./utils.js"; +import {rgthree} from "./rgthree.js"; +import { + SERVICE as CONTEXT_SERVICE, + InputMutation, + InputMutationOperation, +} from "./services/context_service.js"; +import {NodeTypesString} from "./constants.js"; +import {removeUnusedInputsFromEnd} from "./utils_inputs_outputs.js"; +import {INodeInputSlot, INodeOutputSlot, INodeSlot, LGraphNode, LLink} from "typings/litegraph.js"; +import {ComfyNodeConstructor, ComfyObjectInfo} from "typings/comfy.js"; +import {DynamicContextNodeBase} from "./dynamic_context_base.js"; +import {SERVICE as CONFIG_SERVICE} from "./services/config_service.js"; + +const OWNED_PREFIX = "+"; +const REGEX_OWNED_PREFIX = /^\+\s*/; +const REGEX_EMPTY_INPUT = /^\+\s*$/; + +/** + * The Dynamic Context node. + */ +export class DynamicContextNode extends DynamicContextNodeBase { + static override title = NodeTypesString.DYNAMIC_CONTEXT; + static override type = NodeTypesString.DYNAMIC_CONTEXT; + static comfyClass = NodeTypesString.DYNAMIC_CONTEXT; + + constructor(title = DynamicContextNode.title) { + super(title); + } + + override onNodeCreated() { + this.addInput("base_ctx", "RGTHREE_DYNAMIC_CONTEXT"); + this.ensureOneRemainingNewInputSlot(); + super.onNodeCreated(); + } + + override onConnectionsChange( + type: number, + slotIndex: number, + isConnected: boolean, + link: LLink, + ioSlot: INodeSlot, + ): void { + super.onConnectionsChange?.call(this, type, slotIndex, isConnected, link, ioSlot); + if (this.configuring) { + return; + } + if (type === LiteGraph.INPUT) { + if (isConnected) { + this.handleInputConnected(slotIndex); + } else { + this.handleInputDisconnected(slotIndex); + } + } + } + + override onConnectInput( + inputIndex: number, + outputType: INodeOutputSlot["type"], + outputSlot: INodeOutputSlot, + outputNode: LGraphNode, + outputIndex: number, + ): boolean { + let canConnect = true; + if (super.onConnectInput) { + canConnect = super.onConnectInput.apply(this, [...arguments] as any); + } + if ( + canConnect && + outputNode instanceof DynamicContextNode && + outputIndex === 0 && + inputIndex !== 0 + ) { + const [n, v] = rgthree.logger.warnParts( + "Currently, you can only connect a context node in the first slot.", + ); + console[n]?.call(console, ...v); + canConnect = false; + } + return canConnect; + } + + handleInputConnected(slotIndex: number) { + const ioSlot = this.inputs[slotIndex]; + const connectedIndexes = []; + if (slotIndex === 0) { + let baseNodeInfos = getConnectedInputInfosAndFilterPassThroughs(this, this, 0); + const baseNodes = baseNodeInfos.map((n) => n.node)!; + const baseNodesDynamicCtx = baseNodes[0] as DynamicContextNodeBase; + if (baseNodesDynamicCtx?.provideInputsData) { + const inputsData = CONTEXT_SERVICE.getDynamicContextInputsData(baseNodesDynamicCtx); + console.log("inputsData", inputsData); + for (const input of baseNodesDynamicCtx.provideInputsData()) { + if (input.name === "base_ctx" || input.name === "+") { + continue; + } + this.addContextInput(input.name, input.type, input.index); + this.stabilizeNames(); + } + } + } else if (this.isInputSlotForNewInput(slotIndex)) { + this.handleNewInputConnected(slotIndex); + } + } + + isInputSlotForNewInput(slotIndex: number) { + const ioSlot = this.inputs[slotIndex]; + return ioSlot && ioSlot.name === "+" && ioSlot.type === "*"; + } + + handleNewInputConnected(slotIndex: number) { + if (!this.isInputSlotForNewInput(slotIndex)) { + throw new Error('Expected the incoming slot index to be the "new input" input.'); + } + const ioSlot = this.inputs[slotIndex]!; + let cxn = null; + if (ioSlot.link != null) { + cxn = followConnectionUntilType(this, IoDirection.INPUT, slotIndex, true); + } + if (cxn?.type && cxn?.name) { + let name = this.addOwnedPrefix(this.getNextUniqueNameForThisNode(cxn.name)); + if (name.match(/^\+\s*[A-Z_]+(\.\d+)?$/)) { + name = name.toLowerCase(); + } + ioSlot.name = name; + ioSlot.type = cxn.type as string; + ioSlot.removable = true; + while (!this.outputs[slotIndex]) { + this.addOutput("*", "*"); + } + this.outputs[slotIndex]!.type = cxn.type as string; + this.outputs[slotIndex]!.name = this.stripOwnedPrefix(name).toLocaleUpperCase(); + // This is a dumb override for ComfyUI's widgetinputs issues. + if (cxn.type === "COMBO" || cxn.type.includes(",") || Array.isArray(cxn.type)) { + (this.outputs[slotIndex] as any).widget = true; + } + this.inputsMutated({ + operation: InputMutationOperation.ADDED, + node: this, + slotIndex, + slot: ioSlot, + }); + this.stabilizeNames(); + this.ensureOneRemainingNewInputSlot(); + } + } + + handleInputDisconnected(slotIndex: number) { + const inputs = this.getContextInputsList(); + if (slotIndex === 0) { + for (let index = inputs.length - 1; index > 0; index--) { + if (index === 0 || index === inputs.length - 1) { + continue; + } + const input = inputs[index]!; + if (!this.isOwnedInput(input.name)) { + if (input.link || this.outputs[index]?.links?.length) { + this.renameContextInput(index, input.name, true); + } else { + this.removeContextInput(index); + } + } + } + this.setSize(this.computeSize()); + this.setDirtyCanvas(true, true); + } + } + + ensureOneRemainingNewInputSlot() { + removeUnusedInputsFromEnd(this, 1, REGEX_EMPTY_INPUT); + this.addInput(OWNED_PREFIX, "*"); + } + + getNextUniqueNameForThisNode(desiredName: string) { + const inputs = this.getContextInputsList(); + const allExistingKeys = inputs.map((i) => this.stripOwnedPrefix(i.name).toLocaleUpperCase()); + desiredName = this.stripOwnedPrefix(desiredName); + let newName = desiredName; + let n = 0; + while (allExistingKeys.includes(newName.toLocaleUpperCase())) { + newName = `${desiredName}.${++n}`; + } + return newName; + } + + override removeInput(slotIndex: number) { + const slot = this.inputs[slotIndex]!; + super.removeInput(slotIndex); + if (this.outputs[slotIndex]) { + this.removeOutput(slotIndex); + } + this.inputsMutated({operation: InputMutationOperation.REMOVED, node: this, slotIndex, slot}); + this.stabilizeNames(); + } + + stabilizeNames() { + const inputs = this.getContextInputsList(); + const names: string[] = []; + for (const [index, input] of inputs.entries()) { + if (index === 0 || index === inputs.length - 1) { + continue; + } + input.label = undefined; + this.outputs[index]!.label = undefined; + let origName = this.stripOwnedPrefix(input.name).replace(/\.\d+$/, ""); + let name = input.name; + if (!this.isOwnedInput(name)) { + names.push(name.toLocaleUpperCase()); + } else { + let n = 0; + name = this.addOwnedPrefix(origName); + while (names.includes(this.stripOwnedPrefix(name).toLocaleUpperCase())) { + name = `${this.addOwnedPrefix(origName)}.${++n}`; + } + names.push(this.stripOwnedPrefix(name).toLocaleUpperCase()); + if (input.name !== name) { + this.renameContextInput(index, name); + } + } + } + } + + override getSlotMenuOptions(slot: { + slot: number; + input?: INodeInputSlot | undefined; + output?: INodeOutputSlot | undefined; + }) { + const editable = this.isOwnedInput(slot.input!.name) && this.type !== "*"; + return [ + { + content: "โœ๏ธ Rename Input", + disabled: !editable, + callback: () => { + var dialog = app.canvas.createDialog( + "Name", + {}, + ); + var dialogInput = dialog.querySelector("input")!; + if (dialogInput) { + dialogInput.value = this.stripOwnedPrefix(slot.input!.name || ""); + } + var inner = () => { + this.handleContextMenuRenameInputDialog(slot.slot, dialogInput.value); + dialog.close(); + }; + dialog.querySelector("button")!.addEventListener("click", inner); + dialogInput.addEventListener("keydown", (e) => { + dialog.is_modified = true; + if (e.keyCode == 27) { + dialog.close(); + } else if (e.keyCode == 13) { + inner(); + } else if (e.keyCode != 13 && (e.target as HTMLElement)?.localName != "textarea") { + return; + } + e.preventDefault(); + e.stopPropagation(); + }); + dialogInput.focus(); + }, + }, + { + content: "๐Ÿ—‘๏ธ Delete Input", + disabled: !editable, + callback: () => { + this.removeInput(slot.slot); + }, + }, + ]; + } + + handleContextMenuRenameInputDialog(slotIndex: number, value: string) { + app.graph.beforeChange(); + this.renameContextInput(slotIndex, value); + this.stabilizeNames(); + this.setDirtyCanvas(true, true); + app.graph.afterChange(); + } +} + +const contextDynamicNodes = [DynamicContextNode]; +app.registerExtension({ + name: "rgthree.DynamicContext", + async beforeRegisterNodeDef(nodeType: ComfyNodeConstructor, nodeData: ComfyObjectInfo) { + if (!CONFIG_SERVICE.getConfigValue("unreleased.dynamic_context.enabled")) { + return; + } + if (nodeData.name === DynamicContextNode.type) { + DynamicContextNode.setUp(nodeType, nodeData); + } + }, +}); diff --git a/src_web/comfyui/dynamic_context_base.ts b/src_web/comfyui/dynamic_context_base.ts new file mode 100644 index 0000000..6fe703b --- /dev/null +++ b/src_web/comfyui/dynamic_context_base.ts @@ -0,0 +1,237 @@ +import type {INodeInputSlot} from "typings/litegraph.js"; + +import {BaseContextNode} from "./context.js"; +import {ComfyNodeConstructor, ComfyObjectInfo} from "typings/comfy.js"; +import {RgthreeBaseServerNode} from "./base_node.js"; +import {moveArrayItem, wait} from "rgthree/common/shared_utils.js"; +import {RgthreeInvisibleWidget} from "./utils_widgets.js"; +import { + getContextOutputName, + InputMutation, + InputMutationOperation, +} from "./services/context_service.js"; +import {app} from "scripts/app.js"; +import {SERVICE as CONTEXT_SERVICE} from "./services/context_service.js"; + +const OWNED_PREFIX = "+"; +const REGEX_OWNED_PREFIX = /^\+\s*/; +const REGEX_EMPTY_INPUT = /^\+\s*$/; + +export type InputLike = { + name: string; + type: string | -1; + label?: string; + link: number | null; + removable?: boolean; +}; + +/** + * The base context node that contains some shared between DynamicContext nodes. Not labels + * `abstract` so we can reference `this` in static methods. + */ +export class DynamicContextNodeBase extends BaseContextNode { + protected readonly hasShadowInputs: boolean = false; + + getContextInputsList(): InputLike[] { + return this.inputs; + } + + provideInputsData() { + const inputs = this.getContextInputsList(); + return inputs + .map((input, index) => ({ + name: this.stripOwnedPrefix(input.name), + type: String(input.type), + index, + })) + .filter((i) => i.type !== "*"); + } + + addOwnedPrefix(name: string) { + return `+ ${this.stripOwnedPrefix(name)}`; + } + + isOwnedInput(inputOrName: string | null | INodeInputSlot) { + const name = typeof inputOrName == "string" ? inputOrName : inputOrName?.name || ""; + return REGEX_OWNED_PREFIX.test(name); + } + + stripOwnedPrefix(name: string) { + return name.replace(REGEX_OWNED_PREFIX, ""); + } + + // handleUpstreamMutation(mutation: InputMutation) { + // throw new Error('handleUpstreamMutation not overridden!') + // } + + handleUpstreamMutation(mutation: InputMutation) { + console.log(`[node ${this.id}] handleUpstreamMutation`, mutation); + if (mutation.operation === InputMutationOperation.ADDED) { + const slot = mutation.slot; + if (!slot) { + throw new Error("Cannot have an ADDED mutation without a provided slot data."); + } + this.addContextInput( + this.stripOwnedPrefix(slot.name), + slot.type as string, + mutation.slotIndex, + ); + return; + } + if (mutation.operation === InputMutationOperation.REMOVED) { + const slot = mutation.slot; + if (!slot) { + throw new Error("Cannot have an REMOVED mutation without a provided slot data."); + } + this.removeContextInput(mutation.slotIndex); + return; + } + if (mutation.operation === InputMutationOperation.RENAMED) { + const slot = mutation.slot; + if (!slot) { + throw new Error("Cannot have an RENAMED mutation without a provided slot data."); + } + this.renameContextInput(mutation.slotIndex, slot.name); + return; + } + } + override clone() { + const cloned = super.clone(); + while (cloned.inputs.length > 1) { + cloned.removeInput(cloned.inputs.length - 1); + } + while (cloned.widgets.length > 1) { + cloned.removeWidget(cloned.widgets.length - 1); + } + while (cloned.outputs.length > 1) { + cloned.removeOutput(cloned.outputs.length - 1); + } + return cloned; + } + + /** + * Adds the basic output_keys widget. Should be called _after_ specific nodes setup their inputs + * or widgets. + */ + override onNodeCreated() { + const node = this; + this.addCustomWidget( + new RgthreeInvisibleWidget("output_keys", "RGTHREE_DYNAMIC_CONTEXT_OUTPUTS", "", () => { + return (node.outputs || []) + .map((o, i) => i > 0 && o.name) + .filter((n) => n !== false) + .join(","); + }), + ); + } + + addContextInput(name: string, type: string, slot = -1) { + const inputs = this.getContextInputsList(); + if (this.hasShadowInputs) { + inputs.push({name, type, link: null}); + } else { + this.addInput(name, type); + } + if (slot > -1) { + moveArrayItem(inputs, inputs.length - 1, slot); + } else { + slot = inputs.length - 1; + } + if (type !== "*") { + const output = this.addOutput(getContextOutputName(name), type); + if (type === "COMBO" || String(type).includes(",") || Array.isArray(type)) { + (output as any).widget = true; + } + if (slot > -1) { + moveArrayItem(this.outputs, this.outputs.length - 1, slot); + } + } + this.fixInputsOutputsLinkSlots(); + this.inputsMutated({ + operation: InputMutationOperation.ADDED, + node: this, + slotIndex: slot, + slot: inputs[slot]!, + }); + } + + removeContextInput(slotIndex: number) { + if (this.hasShadowInputs) { + const inputs = this.getContextInputsList(); + const input = inputs.splice(slotIndex, 1)[0]; + if (this.outputs[slotIndex]) { + this.removeOutput(slotIndex); + } + } else { + this.removeInput(slotIndex); + } + } + + renameContextInput(index: number, newName: string, forceOwnBool: boolean | null = null) { + const inputs = this.getContextInputsList(); + const input = inputs[index]!; + const oldName = input.name; + newName = this.stripOwnedPrefix(newName.trim() || this.getSlotDefaultInputLabel(index)); + if (forceOwnBool === true || (this.isOwnedInput(oldName) && forceOwnBool !== false)) { + newName = this.addOwnedPrefix(newName); + } + if (oldName !== newName) { + input.name = newName; + input.removable = this.isOwnedInput(newName); + this.outputs[index]!.name = getContextOutputName(inputs[index]!.name); + this.inputsMutated({ + node: this, + operation: InputMutationOperation.RENAMED, + slotIndex: index, + slot: input, + }); + } + } + + getSlotDefaultInputLabel(slotIndex: number) { + const inputs = this.getContextInputsList(); + const input = inputs[slotIndex]!; + let defaultLabel = this.stripOwnedPrefix(input.name).toLowerCase(); + return defaultLabel.toLocaleLowerCase(); + } + + inputsMutated(mutation: InputMutation) { + CONTEXT_SERVICE.onInputChanges(this, mutation); + } + + fixInputsOutputsLinkSlots() { + if (!this.hasShadowInputs) { + const inputs = this.getContextInputsList(); + for (let index = inputs.length - 1; index > 0; index--) { + const input = inputs[index]!; + if ((input === null || input === void 0 ? void 0 : input.link) != null) { + app.graph.links[input.link!]!.target_slot = index; + } + } + } + const outputs = this.outputs; + for (let index = outputs.length - 1; index > 0; index--) { + const output = outputs[index]; + if (output) { + output.nameLocked = true; + for (const link of output.links || []) { + app.graph.links[link!]!.origin_slot = index; + } + } + } + } + + static override setUp(comfyClass: ComfyNodeConstructor, nodeData: ComfyObjectInfo) { + RgthreeBaseServerNode.registerForOverride(comfyClass, nodeData, this); + // [๐Ÿคฎ] ComfyUI only adds "required" inputs to the outputs list when dragging an output to + // empty space, but since RGTHREE_CONTEXT is optional, it doesn't get added to the menu because + // ...of course. So, we'll manually add it. Of course, we also have to do this in a timeout + // because ComfyUI clears out `LiteGraph.slot_types_default_out` in its own 'Comfy.SlotDefaults' + // extension and we need to wait for that to happen. + wait(500).then(() => { + LiteGraph.slot_types_default_out["RGTHREE_DYNAMIC_CONTEXT"] = + LiteGraph.slot_types_default_out["RGTHREE_DYNAMIC_CONTEXT"] || []; + LiteGraph.slot_types_default_out["RGTHREE_DYNAMIC_CONTEXT"].push(comfyClass.comfyClass); + }); + } +} diff --git a/src_web/comfyui/dynamic_context_switch.ts b/src_web/comfyui/dynamic_context_switch.ts new file mode 100644 index 0000000..34c7db7 --- /dev/null +++ b/src_web/comfyui/dynamic_context_switch.ts @@ -0,0 +1,207 @@ +import type {ComfyNodeConstructor, ComfyObjectInfo} from "typings/comfy.js"; +import type {INodeSlot, LGraphNode, LLink, LGraphCanvas} from "typings/litegraph.js"; + +import {app} from "scripts/app.js"; +import {DynamicContextNodeBase, InputLike} from "./dynamic_context_base.js"; +import {NodeTypesString} from "./constants.js"; +import { + InputMutation, + SERVICE as CONTEXT_SERVICE, + stripContextInputPrefixes, + getContextOutputName, +} from "./services/context_service.js"; +import {getConnectedInputNodesAndFilterPassThroughs} from "./utils.js"; +import {debounce, moveArrayItem} from "rgthree/common/shared_utils.js"; +import {measureText} from "./utils_canvas.js"; +import {SERVICE as CONFIG_SERVICE} from "./services/config_service.js"; + +type ShadowInputData = { + node: LGraphNode; + slot: number; + shadowIndex: number; + shadowIndexIfShownSingularly: number; + shadowIndexFull: number; + nodeIndex: number; + type: string | -1; + name: string; + key: string; + // isDuplicatedBefore: boolean, + duplicatesBefore: number[]; + duplicatesAfter: number[]; +}; + +/** + * The Context Switch node. + */ +class DynamicContextSwitchNode extends DynamicContextNodeBase { + static override title = NodeTypesString.DYNAMIC_CONTEXT_SWITCH; + static override type = NodeTypesString.DYNAMIC_CONTEXT_SWITCH; + static comfyClass = NodeTypesString.DYNAMIC_CONTEXT_SWITCH; + + protected override readonly hasShadowInputs = true; + + // override hasShadowInputs = true; + + /** + * We should be able to assume that `lastInputsList` is the input list after the last, major + * synchronous change. Which should mean, if we're handling a change that is currently live, but + * not represented in our node (like, an upstream node has already removed an input), then we + * should be able to compar the current InputList to this `lastInputsList`. + */ + lastInputsList: ShadowInputData[] = []; + + private shadowInputs: (InputLike & {count: number})[] = [ + {name: "base_ctx", type: "RGTHREE_DYNAMIC_CONTEXT", link: null, count: 0}, + ]; + + constructor(title = DynamicContextSwitchNode.title) { + super(title); + } + + override getContextInputsList() { + return this.shadowInputs; + } + override handleUpstreamMutation(mutation: InputMutation) { + this.scheduleHardRefresh(); + } + + override onConnectionsChange( + type: number, + slotIndex: number, + isConnected: boolean, + link: LLink, + ioSlot: INodeSlot, + ): void { + super.onConnectionsChange?.call(this, type, slotIndex, isConnected, link, ioSlot); + if (this.configuring) { + return; + } + if (type === LiteGraph.INPUT) { + this.scheduleHardRefresh(); + } + } + + scheduleHardRefresh(ms = 64) { + return debounce(() => { + this.refreshInputsAndOutputs(); + }, ms); + } + + override onNodeCreated() { + this.addInput("ctx_1", "RGTHREE_DYNAMIC_CONTEXT"); + this.addInput("ctx_2", "RGTHREE_DYNAMIC_CONTEXT"); + this.addInput("ctx_3", "RGTHREE_DYNAMIC_CONTEXT"); + this.addInput("ctx_4", "RGTHREE_DYNAMIC_CONTEXT"); + this.addInput("ctx_5", "RGTHREE_DYNAMIC_CONTEXT"); + super.onNodeCreated(); + } + + override addContextInput(name: string, type: string, slot?: number): void {} + + /** + * This is a "hard" refresh of the list, but looping over the actual context inputs, and + * recompiling the shadowInputs and outputs. + */ + private refreshInputsAndOutputs() { + const inputs: (InputLike & {count: number})[] = [ + {name: "base_ctx", type: "RGTHREE_DYNAMIC_CONTEXT", link: null, count: 0}, + ]; + let numConnected = 0; + for (let i = 0; i < this.inputs.length; i++) { + const childCtxs = getConnectedInputNodesAndFilterPassThroughs( + this, + this, + i, + ) as DynamicContextNodeBase[]; + if (childCtxs.length > 1) { + throw new Error("How is there more than one input?"); + } + const ctx = childCtxs[0]; + if (!ctx) continue; + numConnected++; + const slotsData = CONTEXT_SERVICE.getDynamicContextInputsData(ctx); + console.log(slotsData); + for (const slotData of slotsData) { + const found = inputs.find( + (n) => getContextOutputName(slotData.name) === getContextOutputName(n.name), + ); + if (found) { + found.count += 1; + continue; + } + inputs.push({ + name: slotData.name, + type: slotData.type, + link: null, + count: 1, + }); + } + } + this.shadowInputs = inputs; + // First output is always CONTEXT, so "p" is the offset. + let i = 0; + for (i; i < this.shadowInputs.length; i++) { + const data = this.shadowInputs[i]!; + let existing = this.outputs.find( + (o) => getContextOutputName(o.name) === getContextOutputName(data.name), + ); + if (!existing) { + existing = this.addOutput(getContextOutputName(data.name), data.type); + } + moveArrayItem(this.outputs, existing, i); + delete existing.rgthree_status; + if (data.count !== numConnected) { + existing.rgthree_status = "WARN"; + } + } + while (this.outputs[i]) { + const output = this.outputs[i]; + if (output?.links?.length) { + output.rgthree_status = "ERROR"; + i++; + } else { + this.removeOutput(i); + } + } + this.fixInputsOutputsLinkSlots(); + } + + override onDrawForeground(ctx: CanvasRenderingContext2D, canvas: LGraphCanvas): void { + const low_quality = (canvas?.ds?.scale ?? 1) < 0.6; + if (low_quality || this.size[0] <= 10) { + return; + } + let y = LiteGraph.NODE_SLOT_HEIGHT - 1; + const w = this.size[0]; + ctx.save(); + ctx.font = "normal " + LiteGraph.NODE_SUBTEXT_SIZE + "px Arial"; + ctx.textAlign = "right"; + + for (const output of this.outputs) { + if (!output.rgthree_status) { + y += LiteGraph.NODE_SLOT_HEIGHT; + continue; + } + const x = w - 20 - measureText(ctx, output.name); + if (output.rgthree_status === "ERROR") { + ctx.fillText("๐Ÿ›‘", x, y); + } else if (output.rgthree_status === "WARN") { + ctx.fillText("โš ๏ธ", x, y); + } + y += LiteGraph.NODE_SLOT_HEIGHT; + } + ctx.restore(); + } +} + +app.registerExtension({ + name: "rgthree.DynamicContextSwitch", + async beforeRegisterNodeDef(nodeType: ComfyNodeConstructor, nodeData: ComfyObjectInfo) { + if (!CONFIG_SERVICE.getConfigValue("unreleased.dynamic_context.enabled")) { + return; + } + if (nodeData.name === DynamicContextSwitchNode.type) { + DynamicContextSwitchNode.setUp(nodeType, nodeData); + } + }, +}); diff --git a/src_web/comfyui/rgthree.ts b/src_web/comfyui/rgthree.ts index 54fce1d..165491b 100644 --- a/src_web/comfyui/rgthree.ts +++ b/src_web/comfyui/rgthree.ts @@ -1,996 +1,996 @@ -import type { - LGraphCanvas as TLGraphCanvas, - LGraphNode, - SerializedLGraphNode, - serializedLGraph, - ContextMenuItem, - LGraph as TLGraph, - AdjustedMouseEvent, - IContextMenuOptions, -} from "typings/litegraph.js"; -import type { ComfyApiFormat, ComfyApiPrompt, ComfyApp } from "typings/comfy.js"; -import { app } from "scripts/app.js"; -import { api } from "scripts/api.js"; -import { SERVICE as CONFIG_SERVICE } from "./services/config_service.js"; -import { fixBadLinks } from "rgthree/common/link_fixer.js"; -import { injectCss, wait } from "rgthree/common/shared_utils.js"; -import { replaceNode, waitForCanvas, waitForGraph } from "./utils.js"; -import { NodeTypesString, addRgthree, getNodeTypeStrings, stripRgthree } from "./constants.js"; -import { RgthreeProgressBar } from "rgthree/common/progress_bar.js"; -import { RgthreeConfigDialog } from "./config.js"; -import { - iconGear, - iconNode, - iconReplace, - iconStarFilled, - logoRgthree, -} from "rgthree/common/media/svgs.js"; -import type { Bookmark } from "./bookmark"; -import { createElement, query, queryOne } from "rgthree/common/utils_dom.js"; - -export enum LogLevel { - IMPORTANT = 1, - ERROR, - WARN, - INFO, - DEBUG, - DEV, -} - -const LogLevelKeyToLogLevel: { [key: string]: LogLevel } = { - IMPORTANT: LogLevel.IMPORTANT, - ERROR: LogLevel.ERROR, - WARN: LogLevel.WARN, - INFO: LogLevel.INFO, - DEBUG: LogLevel.DEBUG, - DEV: LogLevel.DEV, -}; - -type ConsoleLogFns = "log" | "error" | "warn" | "debug" | "info"; -const LogLevelToMethod: { [key in LogLevel]: ConsoleLogFns } = { - [LogLevel.IMPORTANT]: "log", - [LogLevel.ERROR]: "error", - [LogLevel.WARN]: "warn", - [LogLevel.INFO]: "info", - [LogLevel.DEBUG]: "log", - [LogLevel.DEV]: "log", -}; -const LogLevelToCSS: { [key in LogLevel]: string } = { - [LogLevel.IMPORTANT]: "font-weight: bold; color: blue;", - [LogLevel.ERROR]: "", - [LogLevel.WARN]: "", - [LogLevel.INFO]: "font-style: italic; color: blue;", - [LogLevel.DEBUG]: "font-style: italic; color: #444;", - [LogLevel.DEV]: "color: #004b68;", -}; - -let GLOBAL_LOG_LEVEL = LogLevel.ERROR; - -/** - * A blocklist of extensions to disallow hooking into rgthree's base classes when calling the - * `rgthree.invokeExtensionsAsync` method (which runs outside of ComfyNode's - * `app.invokeExtensionsAsync` which is private). - * - * In Apr 2024 the base rgthree node class added support for other extensions using `nodeCreated` - * and `beforeRegisterNodeDef` which allows other extensions to modify the class. However, since it - * had been months since divorcing the ComfyNode in rgthree-comfy due to instability and - * inflexibility, this was a bit risky as other extensions hadn't ever run with this ability. This - * list attempts to block extensions from being able to call into rgthree-comfy nodes via the - * `nodeCreated` and `beforeRegisterNodeDef` callbacks now that rgthree-comfy is utilizing them - * because they do not work. Oddly, it's ComfyUI's own extension that is broken. - */ -const INVOKE_EXTENSIONS_BLOCKLIST = [ - { - name: "Comfy.WidgetInputs", - reason: - "Major conflict with rgthree-comfy nodes' inputs causing instability and " + - "repeated link disconnections.", - }, - { - name: "efficiency.widgethider", - reason: - "Overrides value getter before widget getter is prepared. Can be lifted if/when " + - "https://github.com/jags111/efficiency-nodes-comfyui/pull/203 is pulled.", - }, -]; - -/** A basic wrapper around logger. */ -class Logger { - /** Logs a message to the console if it meets the current log level. */ - log(level: LogLevel, message: string, ...args: any[]) { - const [n, v] = this.logParts(level, message, ...args); - console[n]?.(...v); - } - - /** - * Returns a tuple of the console function and its arguments. Useful for callers to make the - * actual console. call to gain benefits of DevTools knowing the source line. - * - * If the input is invalid or the level doesn't meet the configuration level, then the return - * value is an unknown function and empty set of values. Callers can use optionla chaining - * successfully: - * - * const [fn, values] = logger.logPars(LogLevel.INFO, 'my message'); - * console[fn]?.(...values); // Will work even if INFO won't be logged. - * - */ - logParts(level: LogLevel, message: string, ...args: any[]): [ConsoleLogFns, any[]] { - if (level <= GLOBAL_LOG_LEVEL) { - const css = LogLevelToCSS[level] || ""; - if (level === LogLevel.DEV) { - message = `๐Ÿ”ง ${message}`; - } - return [LogLevelToMethod[level], [`%c${message}`, css, ...args]]; - } - return ["none" as "info", []]; - } -} - -/** - * A log session, with the name as the prefix. A new session will stack prefixes. - */ -class LogSession { - readonly logger = new Logger(); - readonly logsCache: { [key: string]: { lastShownTime: number } } = {}; - - constructor(readonly name?: string) {} - - /** - * Returns the console log method to use and the arguments to pass so the call site can log from - * there. This extra work at the call site allows for easier debugging in the dev console. - * - * const [logMethod, logArgs] = logger.logParts(LogLevel.DEBUG, message, ...args); - * console[logMethod]?.(...logArgs); - */ - logParts(level: LogLevel, message?: string, ...args: any[]): [ConsoleLogFns, any[]] { - message = `${this.name || ""}${message ? " " + message : ""}`; - return this.logger.logParts(level, message, ...args); - } - - logPartsOnceForTime( - level: LogLevel, - time: number, - message?: string, - ...args: any[] - ): [ConsoleLogFns, any[]] { - message = `${this.name || ""}${message ? " " + message : ""}`; - const cacheKey = `${level}:${message}`; - const cacheEntry = this.logsCache[cacheKey]; - const now = +new Date(); - if (cacheEntry && cacheEntry.lastShownTime + time > now) { - return ["none" as "info", []]; - } - const parts = this.logger.logParts(level, message, ...args); - if (console[parts[0]]) { - this.logsCache[cacheKey] = this.logsCache[cacheKey] || ({} as { lastShownTime: number }); - this.logsCache[cacheKey]!.lastShownTime = now; - } - return parts; - } - - debugParts(message?: string, ...args: any[]) { - return this.logParts(LogLevel.DEBUG, message, ...args); - } - - infoParts(message?: string, ...args: any[]) { - return this.logParts(LogLevel.INFO, message, ...args); - } - - warnParts(message?: string, ...args: any[]) { - return this.logParts(LogLevel.WARN, message, ...args); - } - - newSession(name?: string) { - return new LogSession(`${this.name}${name}`); - } -} - -export type RgthreeUiMessage = { - id: string; - message: string; - type?: "warn" | "info" | "success" | null; - timeout?: number; - // closeable?: boolean; // TODO - actions?: Array<{ - label: string; - href?: string; - callback?: (event: MouseEvent) => void; - }>; -}; - -/** - * A global class as 'rgthree'; exposed on wiindow. Lots can go in here. - */ -class Rgthree extends EventTarget { - /** Exposes the ComfyUI api instance on rgthree. */ - readonly api = api; - private settingsDialog: RgthreeConfigDialog | null = null; - private progressBarEl: RgthreeProgressBar | null = null; - private rgthreeCssPromise: Promise; - - /** Stores a node id that we will use to queu only that output node (with `queueOutputNode`). */ - private queueNodeIds: number[] | null = null; - - logger = new LogSession("[rgthree]"); - - monitorBadLinksAlerted = false; - monitorLinkTimeout: number | null = null; - - processingQueue = false; - loadingApiJson = false; - replacingReroute: number | null = null; - processingMouseDown = false; - processingMouseUp = false; - processingMouseMove = false; - lastAdjustedMouseEvent: AdjustedMouseEvent | null = null; - - // Comfy/LiteGraph states so nodes and tell what the hell is going on. - canvasCurrentlyCopyingToClipboard = false; - canvasCurrentlyCopyingToClipboardWithMultipleNodes = false; - initialGraphToPromptSerializedWorkflowBecauseComfyUIBrokeStuff: any = null; - - private elDebugKeydowns: HTMLDivElement | null = null; - - private readonly isMac: boolean = !!( - navigator.platform?.toLocaleUpperCase().startsWith("MAC") || - (navigator as any).userAgentData?.platform?.toLocaleUpperCase().startsWith("MAC") - ); - - constructor() { - super(); - - const logLevel = - LogLevelKeyToLogLevel[CONFIG_SERVICE.getConfigValue("log_level")] ?? GLOBAL_LOG_LEVEL; - this.setLogLevel(logLevel); - - this.initializeGraphAndCanvasHooks(); - this.initializeComfyUIHooks(); - this.initializeContextMenu(); - - this.rgthreeCssPromise = injectCss("extensions/rgthree-comfy/rgthree.css"); - - this.initializeProgressBar(); - - CONFIG_SERVICE.addEventListener("config-change", ((e: CustomEvent) => { - if (e.detail?.key?.includes("features.progress_bar")) { - this.initializeProgressBar(); - } - }) as EventListener); - } - - /** - * Initializes the top progress bar, if it's configured. - */ - async initializeProgressBar() { - if (CONFIG_SERVICE.getConfigValue("features.progress_bar.enabled")) { - await this.rgthreeCssPromise; - if (!this.progressBarEl) { - this.progressBarEl = RgthreeProgressBar.create(); - this.progressBarEl.setAttribute( - "title", - "Progress Bar by rgthree. right-click for rgthree menu.", - ); - - this.progressBarEl.addEventListener("contextmenu", async (e) => { - e.stopPropagation(); - e.preventDefault(); - }); - - this.progressBarEl.addEventListener("pointerdown", async (e) => { - LiteGraph.closeAllContextMenus(); - if (e.button == 2) { - const canvas = await waitForCanvas(); - new LiteGraph.ContextMenu( - this.getRgthreeContextMenuItems(), - { - title: `
${logoRgthree} rgthree-comfy
`, - left: e.clientX, - top: 5, - }, - canvas.getCanvasWindow(), - ); - return; - } - if (e.button == 0) { - const nodeId = this.progressBarEl?.currentNodeId; - if (nodeId) { - const [canvas, graph] = await Promise.all([waitForCanvas(), waitForGraph()]); - const node = graph.getNodeById(Number(nodeId)); - if (node) { - canvas.centerOnNode(node); - e.stopPropagation(); - e.preventDefault(); - } - } - return; - } - }); - } - // Handle both cases in case someone hasn't updated. Can probably just assume - // `isUpdatedComfyBodyClasses` is true in the near future. - const isUpdatedComfyBodyClasses = !!queryOne(".comfyui-body-top"); - const position = CONFIG_SERVICE.getConfigValue("features.progress_bar.position"); - this.progressBarEl.classList.toggle("rgthree-pos-bottom", position === "bottom"); - // If ComfyUI is updated with the body segments, then use that. - if (isUpdatedComfyBodyClasses) { - if (position === "bottom") { - queryOne(".comfyui-body-bottom")!.appendChild(this.progressBarEl); - } else { - queryOne(".comfyui-body-top")!.appendChild(this.progressBarEl); - } - } else { - document.body.appendChild(this.progressBarEl); - } - const height = CONFIG_SERVICE.getConfigValue("features.progress_bar.height") || 14; - this.progressBarEl.style.height = `${height}px`; - const fontSize = Math.max(10, Number(height) - 10); - this.progressBarEl.style.fontSize = `${fontSize}px`; - this.progressBarEl.style.fontWeight = fontSize <= 12 ? "bold" : "normal"; - } else { - this.progressBarEl?.remove(); - } - } - - /** - * Initialize a bunch of hooks into LiteGraph itself so we can either keep state or context on - * what's happening so nodes can respond appropriately. This is usually to fix broken assumptions - * in the unowned code [๐Ÿคฎ], but sometimes to add features or enhancements too [โญ]. - */ - private async initializeGraphAndCanvasHooks() { - const rgthree = this; - - // [๐Ÿคฎ] To mitigate changes from https://github.com/rgthree/rgthree-comfy/issues/69 - // and https://github.com/comfyanonymous/ComfyUI/issues/2193 we can try to store the workflow - // node so our nodes can find the seralized node. Works with method - // `getNodeFromInitialGraphToPromptSerializedWorkflowBecauseComfyUIBrokeStuff` to find a node - // while serializing. What a way to work around... - const graphSerialize = LGraph.prototype.serialize; - LGraph.prototype.serialize = function () { - const response = graphSerialize.apply(this, [...arguments] as any) as any; - rgthree.initialGraphToPromptSerializedWorkflowBecauseComfyUIBrokeStuff = response; - return response; - }; - - // Overrides LiteGraphs' processMouseDown to both keep state as well as dispatch a custom event. - const processMouseDown = LGraphCanvas.prototype.processMouseDown; - LGraphCanvas.prototype.processMouseDown = function (e: AdjustedMouseEvent) { - rgthree.processingMouseDown = true; - const returnVal = processMouseDown.apply(this, [...arguments] as any); - rgthree.dispatchCustomEvent("on-process-mouse-down", { originalEvent: e }); - rgthree.processingMouseDown = false; - return returnVal; - }; - - // Overrides LiteGraph's `adjustMouseEvent` to capture the last even coming in and out. Useful - // to capture the last `canvasX` and `canvasY` properties, which are not the same as LiteGraph's - // `canvas.last_mouse_position`, unfortunately. - const adjustMouseEvent = LGraphCanvas.prototype.adjustMouseEvent; - LGraphCanvas.prototype.adjustMouseEvent = function (e: PointerEvent) { - adjustMouseEvent.apply(this, [...arguments] as any); - rgthree.lastAdjustedMouseEvent = e as AdjustedMouseEvent; - }; - - // [๐Ÿคฎ] Copying to clipboard clones nodes and then manipulats the linking data manually which - // does not allow a node to handle connections. This harms nodes that manually handle inputs, - // like our any-input nodes that may start with one input, and manually add new ones when one is - // attached. - const copyToClipboard = LGraphCanvas.prototype.copyToClipboard; - LGraphCanvas.prototype.copyToClipboard = function (nodes: LGraphNode[]) { - rgthree.canvasCurrentlyCopyingToClipboard = true; - rgthree.canvasCurrentlyCopyingToClipboardWithMultipleNodes = - Object.values(nodes || this.selected_nodes || []).length > 1; - copyToClipboard.apply(this, [...arguments] as any); - rgthree.canvasCurrentlyCopyingToClipboard = false; - rgthree.canvasCurrentlyCopyingToClipboardWithMultipleNodes = false; - }; - - // [โญ] Make it so when we add a group, we get to name it immediately. - const onGroupAdd = LGraphCanvas.onGroupAdd; - LGraphCanvas.onGroupAdd = function (...args: any[]) { - const graph = app.graph as TLGraph; - onGroupAdd.apply(this, [...args] as any); - LGraphCanvas.onShowPropertyEditor( - {}, - null, - null, - null, - graph._groups[graph._groups.length - 1], - ); - }; - } - - /** - * [๐Ÿคฎ] Handles the same exact thing as ComfyApp's `invokeExtensionsAsync`, but done here since - * it is #private in ComfyApp because... of course it us. This is necessary since we purposefully - * avoid using the ComfyNode due to historical instability and inflexibility for all the advanced - * ui stuff rgthree-comfy nodes do, but we can still have other custom nodes know what's happening - * with rgthree-comfy; specifically, for `nodeCreated` as of now. - */ - async invokeExtensionsAsync(method: "nodeCreated", ...args: any[]) { - const comfyapp = app as ComfyApp; - if (CONFIG_SERVICE.getConfigValue("features.invoke_extensions_async.node_created") === false) { - const [m, a] = this.logParts( - LogLevel.INFO, - `Skipping invokeExtensionsAsync for applicable rgthree-comfy nodes`, - ); - console[m]?.(...a); - return Promise.resolve(); - } - return await Promise.all( - comfyapp.extensions.map(async (ext) => { - if (ext?.[method]) { - try { - const blocked = INVOKE_EXTENSIONS_BLOCKLIST.find((block) => - ext.name.toLowerCase().startsWith(block.name.toLowerCase()), - ); - if (blocked) { - const [n, v] = this.logger.logPartsOnceForTime( - LogLevel.WARN, - 5000, - `Blocked extension '${ext.name}' method '${method}' for rgthree-nodes because: ${blocked.reason}`, - ); - console[n]?.(...v); - return Promise.resolve(); - } - return await (ext[method] as Function)(...args, comfyapp); - } catch (error) { - const [n, v] = this.logParts( - LogLevel.ERROR, - `Error calling extension '${ext.name}' method '${method}' for rgthree-node.`, - { error }, - { extension: ext }, - { args }, - ); - console[n]?.(...v); - } - } - }), - ); - } - - /** - * Wraps `dispatchEvent` for easier CustomEvent dispatching. - */ - private dispatchCustomEvent(event: string, detail?: any) { - if (detail != null) { - return this.dispatchEvent(new CustomEvent(event, { detail })); - } - return this.dispatchEvent(new CustomEvent(event)); - } - - /** - * Initializes hooks specific to an rgthree-comfy context menu on the root menu. - */ - private async initializeContextMenu() { - const that = this; - setTimeout(async () => { - const getCanvasMenuOptions = LGraphCanvas.prototype.getCanvasMenuOptions; - LGraphCanvas.prototype.getCanvasMenuOptions = function (...args: any[]) { - let existingOptions = getCanvasMenuOptions.apply(this, [...args] as any); - - const options = []; - options.push(null); // Divider - options.push(null); // Divider - options.push(null); // Divider - options.push({ - content: logoRgthree + `rgthree-comfy`, - className: "rgthree-contextmenu-item rgthree-contextmenu-main-item-rgthree-comfy", - submenu: { - options: that.getRgthreeContextMenuItems(), - }, - }); - options.push(null); // Divider - options.push(null); // Divider - - let idx = null; - idx = idx || existingOptions.findIndex((o) => o?.content?.startsWith?.("Queue Group")) + 1; - idx = - idx || existingOptions.findIndex((o) => o?.content?.startsWith?.("Queue Selected")) + 1; - idx = idx || existingOptions.findIndex((o) => o?.content?.startsWith?.("Convert to Group")); - idx = idx || existingOptions.findIndex((o) => o?.content?.startsWith?.("Arrange (")); - idx = idx || existingOptions.findIndex((o) => !o) + 1; - idx = idx || 3; - existingOptions.splice(idx, 0, ...options); - for (let i = existingOptions.length; i > 0; i--) { - if (existingOptions[i] === null && existingOptions[i + 1] === null) { - existingOptions.splice(i, 1); - } - } - - return existingOptions; - }; - }, 1016); - } - - /** - * Returns the standard menu items for an rgthree-comfy context menu. - */ - private getRgthreeContextMenuItems(): ContextMenuItem[] { - const [canvas, graph] = [app.canvas as TLGraphCanvas, app.graph as TLGraph]; - const selectedNodes = Object.values(canvas.selected_nodes || {}); - let rerouteNodes: LGraphNode[] = []; - if (selectedNodes.length) { - rerouteNodes = selectedNodes.filter((n) => n.type === "Reroute"); - } else { - rerouteNodes = graph._nodes.filter((n) => n.type == "Reroute"); - } - const rerouteLabel = selectedNodes.length ? "selected" : "all"; - - const showBookmarks = CONFIG_SERVICE.getFeatureValue("menu_bookmarks.enabled"); - const bookmarkMenuItems = showBookmarks ? getBookmarks() : []; - - return [ - { - content: "Nodes", - disabled: true, - className: "rgthree-contextmenu-item rgthree-contextmenu-label", - }, - { - content: iconNode + "All", - className: "rgthree-contextmenu-item", - has_submenu: true, - submenu: { - options: getNodeTypeStrings() as unknown as ContextMenuItem[], - callback: ( - value: string | ContextMenuItem, - options: IContextMenuOptions, - event: MouseEvent, - ) => { - const node = LiteGraph.createNode(addRgthree(value as string)); - node.pos = [ - rgthree.lastAdjustedMouseEvent!.canvasX, - rgthree.lastAdjustedMouseEvent!.canvasY, - ]; - canvas.graph.add(node); - canvas.selectNode(node); - app.graph.setDirtyCanvas(true, true); - }, - extra: { rgthree_doNotNest: true }, - }, - }, - - { - content: "Actions", - disabled: true, - className: "rgthree-contextmenu-item rgthree-contextmenu-label", - }, - { - content: iconGear + "Settings (rgthree-comfy)", - disabled: !!this.settingsDialog, - className: "rgthree-contextmenu-item", - callback: (...args: any[]) => { - this.settingsDialog = new RgthreeConfigDialog().show(); - this.settingsDialog.addEventListener("close", (e) => { - this.settingsDialog = null; - }); - }, - }, - { - content: iconReplace + ` Convert ${rerouteLabel} Reroutes`, - disabled: !rerouteNodes.length, - className: "rgthree-contextmenu-item", - callback: (...args: any[]) => { - const msg = - `Convert ${rerouteLabel} ComfyUI Reroutes to Reroute (rgthree) nodes? \n` + - `(First save a copy of your workflow & check reroute connections afterwards)`; - if (!window.confirm(msg)) { - return; - } - (async () => { - for (const node of [...rerouteNodes]) { - if (node.type == "Reroute") { - this.replacingReroute = node.id; - await replaceNode(node, NodeTypesString.REROUTE); - this.replacingReroute = null; - } - } - })(); - }, - }, - ...bookmarkMenuItems, - { - content: "More...", - disabled: true, - className: "rgthree-contextmenu-item rgthree-contextmenu-label", - }, - { - content: iconStarFilled + "Star on Github", - className: "rgthree-contextmenu-item rgthree-contextmenu-github", - callback: (...args: any[]) => { - window.open("https://github.com/rgthree/rgthree-comfy", "_blank"); - }, - }, - ]; - } - - /** - * Wraps an `app.queuePrompt` call setting a specific node id that we will inspect and change the - * serialized graph right before being sent (below, in our `api.queuePrompt` override). - */ - async queueOutputNodes(nodeIds: number[]) { - try { - this.queueNodeIds = nodeIds; - await app.queuePrompt(); - } catch (e) { - const [n, v] = this.logParts( - LogLevel.ERROR, - `There was an error queuing nodes ${nodeIds}`, - e, - ); - console[n]?.(...v); - } finally { - this.queueNodeIds = null; - } - } - - /** - * Recusively walks backwards from a node adding its inputs to the `newOutput` from `oldOutput`. - */ - private recursiveAddNodes(nodeId: string, oldOutput: ComfyApiFormat, newOutput: ComfyApiFormat) { - let currentId = nodeId; - let currentNode = oldOutput[currentId]!; - if (newOutput[currentId] == null) { - newOutput[currentId] = currentNode; - for (const inputValue of Object.values(currentNode.inputs || [])) { - if (Array.isArray(inputValue)) { - this.recursiveAddNodes(inputValue[0], oldOutput, newOutput); - } - } - } - return newOutput; - } - - /** - * Initialize a bunch of hooks into ComfyUI and/or LiteGraph itself so we can either keep state or - * context on what's happening so nodes can respond appropriately. This is usually to fix broken - * assumptions in the unowned code [๐Ÿคฎ], but sometimes to add features or enhancements too [โญ]. - */ - private initializeComfyUIHooks() { - const rgthree = this; - - // Keep state for when the app is queuing the prompt. For instance, this is used for seed to - // understand if we're serializing because we're queueing (and return the random seed to use) or - // for saving the workflow (and keep -1, etc.). - const queuePrompt = app.queuePrompt as Function; - app.queuePrompt = async function () { - rgthree.processingQueue = true; - rgthree.dispatchCustomEvent("queue"); - try { - await queuePrompt.apply(app, [...arguments]); - } finally { - rgthree.processingQueue = false; - rgthree.dispatchCustomEvent("queue-end"); - } - }; - - // Keep state for when the app is in the middle of loading from an api JSON file. - const loadApiJson = app.loadApiJson; - app.loadApiJson = async function () { - rgthree.loadingApiJson = true; - try { - loadApiJson.apply(app, [...arguments] as any); - } finally { - rgthree.loadingApiJson = false; - } - }; - - // Keep state for when the app is serizalizing the graph to prompt. - const graphToPrompt = app.graphToPrompt; - app.graphToPrompt = async function () { - rgthree.dispatchCustomEvent("graph-to-prompt"); - let promise = graphToPrompt.apply(app, [...arguments] as any); - await promise; - rgthree.dispatchCustomEvent("graph-to-prompt-end"); - return promise; - }; - - // Override the queuePrompt for api to intercept the prompt output and, if queueNodeIds is set, - // then we only want to queue those nodes, by rewriting the api format (prompt 'output' field) - // so only those are evaluated. - const apiQueuePrompt = api.queuePrompt as Function; - api.queuePrompt = async function (index: number, prompt: ComfyApiPrompt) { - if (rgthree.queueNodeIds?.length && prompt.output) { - const oldOutput = prompt.output; - let newOutput = {}; - for (const queueNodeId of rgthree.queueNodeIds) { - rgthree.recursiveAddNodes(String(queueNodeId), oldOutput, newOutput); - } - prompt.output = newOutput; - } - rgthree.dispatchCustomEvent("comfy-api-queue-prompt-before", { - workflow: prompt.workflow, - output: prompt.output, - }); - const response = apiQueuePrompt.apply(app, [index, prompt]); - rgthree.dispatchCustomEvent("comfy-api-queue-prompt-end"); - return response; - }; - - // Hook into a clean call; allow us to clear and rgthree messages. - const clean = app.clean; - app.clean = function () { - rgthree.clearAllMessages(); - clean && clean.apply(app, [...arguments] as any); - }; - - // Hook into a data load, like from an image or JSON drop-in. This is (currently) used to - // monitor for bad linking data. - const loadGraphData = app.loadGraphData; - app.loadGraphData = function (graph: serializedLGraph) { - if (rgthree.monitorLinkTimeout) { - clearTimeout(rgthree.monitorLinkTimeout); - rgthree.monitorLinkTimeout = null; - } - rgthree.clearAllMessages(); - // Try to make a copy to use, because ComfyUI's loadGraphData will modify it. - let graphCopy: serializedLGraph | null; - try { - graphCopy = JSON.parse(JSON.stringify(graph)); - } catch (e) { - graphCopy = null; - } - setTimeout(() => { - const wasLoadingAborted = document - .querySelector(".comfy-modal-content") - ?.textContent?.includes("Loading aborted due"); - const graphToUse = wasLoadingAborted ? graphCopy || graph : app.graph; - const fixBadLinksResult = fixBadLinks(graphToUse as unknown as TLGraph); - if (fixBadLinksResult.hasBadLinks) { - const [n, v] = rgthree.logParts( - LogLevel.WARN, - `The workflow you've loaded has corrupt linking data. Open ${ - new URL(location.href).origin - }/rgthree/link_fixer to try to fix.`, - ); - console[n]?.(...v); - if (CONFIG_SERVICE.getConfigValue("features.show_alerts_for_corrupt_workflows")) { - rgthree.showMessage({ - id: "bad-links", - type: "warn", - message: - "The workflow you've loaded has corrupt linking data that may be able to be fixed.", - actions: [ - { - label: "Open fixer", - href: "/rgthree/link_fixer", - }, - { - label: "Fix in place", - href: "/rgthree/link_fixer", - callback: (event) => { - event.stopPropagation(); - event.preventDefault(); - if ( - confirm( - "This will attempt to fix in place. Please make sure to have a saved copy of your workflow.", - ) - ) { - try { - const fixBadLinksResult = fixBadLinks( - graphToUse as unknown as TLGraph, - true, - ); - if (!fixBadLinksResult.hasBadLinks) { - rgthree.hideMessage("bad-links"); - alert( - "Success! It's possible some valid links may have been affected. Please check and verify your workflow.", - ); - wasLoadingAborted && app.loadGraphData(fixBadLinksResult.graph); - if ( - CONFIG_SERVICE.getConfigValue("features.monitor_for_corrupt_links") || - CONFIG_SERVICE.getConfigValue("features.monitor_bad_links") - ) { - rgthree.monitorLinkTimeout = setTimeout(() => { - rgthree.monitorBadLinks(); - }, 5000); - } - } - } catch (e) { - console.error(e); - alert("Unsuccessful at fixing corrupt data. :("); - rgthree.hideMessage("bad-links"); - } - } - }, - }, - ], - }); - } - } else if ( - CONFIG_SERVICE.getConfigValue("features.monitor_for_corrupt_links") || - CONFIG_SERVICE.getConfigValue("features.monitor_bad_links") - ) { - rgthree.monitorLinkTimeout = setTimeout(() => { - rgthree.monitorBadLinks(); - }, 5000); - } - }, 100); - return loadGraphData && loadGraphData.apply(app, [...arguments] as any); - }; - } - - /** - * [๐Ÿคฎ] Finds a node in the currently serializing workflow from the hook setup above. This is to - * mitigate breakages from https://github.com/comfyanonymous/ComfyUI/issues/2193 we can try to - * store the workflow node so our nodes can find the seralized node. - */ - getNodeFromInitialGraphToPromptSerializedWorkflowBecauseComfyUIBrokeStuff( - node: LGraphNode, - ): SerializedLGraphNode | null { - return ( - this.initialGraphToPromptSerializedWorkflowBecauseComfyUIBrokeStuff?.nodes?.find( - (n: SerializedLGraphNode) => n.id === node.id, - ) ?? null - ); - } - - /** - * Shows a message in the UI. - */ - async showMessage(data: RgthreeUiMessage) { - let container = document.querySelector(".rgthree-top-messages-container"); - if (!container) { - container = document.createElement("div"); - container.classList.add("rgthree-top-messages-container"); - document.body.appendChild(container); - } - // If we have a dialog open then we want to append the message to the dialog so they show over - // the modal. - const dialogs = query("dialog[open]"); - if (dialogs.length) { - let dialog = dialogs[dialogs.length - 1]!; - dialog.appendChild(container); - dialog.addEventListener("close", (e) => { - document.body.appendChild(container!); - }); - } - // Hide if we exist. - await this.hideMessage(data.id); - - const messageContainer = document.createElement("div"); - messageContainer.setAttribute("type", data.type || "info"); - - const message = document.createElement("span"); - message.innerHTML = data.message; - messageContainer.appendChild(message); - - for (let a = 0; a < (data.actions || []).length; a++) { - const action = data.actions![a]!; - if (a > 0) { - const sep = document.createElement("span"); - sep.innerHTML = " | "; - messageContainer.appendChild(sep); - } - - const actionEl = document.createElement("a"); - actionEl.innerText = action.label; - if (action.href) { - actionEl.target = "_blank"; - actionEl.href = action.href; - } - if (action.callback) { - actionEl.onclick = (e) => { - return action.callback!(e); - }; - } - messageContainer.appendChild(actionEl); - } - - const messageAnimContainer = document.createElement("div"); - messageAnimContainer.setAttribute("msg-id", data.id); - messageAnimContainer.appendChild(messageContainer); - container.appendChild(messageAnimContainer); - - // Add. Wait. Measure. Wait. Anim. - await wait(64); - messageAnimContainer.style.marginTop = `-${messageAnimContainer.offsetHeight}px`; - await wait(64); - messageAnimContainer.classList.add("-show"); - - if (data.timeout) { - await wait(data.timeout); - this.hideMessage(data.id); - } - } - - /** - * Hides a message in the UI. - */ - async hideMessage(id: string) { - const msg = document.querySelector(`.rgthree-top-messages-container > [msg-id="${id}"]`); - if (msg?.classList.contains("-show")) { - msg.classList.remove("-show"); - await wait(750); - } - msg && msg.remove(); - } - - /** - * Clears all messages in the UI. - */ - async clearAllMessages() { - let container = document.querySelector(".rgthree-top-messages-container"); - container && (container.innerHTML = ""); - } - - setLogLevel(level?: LogLevel | string) { - if (typeof level === "string") { - level = LogLevelKeyToLogLevel[CONFIG_SERVICE.getConfigValue("log_level")]; - } - if (level != null) { - GLOBAL_LOG_LEVEL = level; - } - } - - logParts(level: LogLevel, message?: string, ...args: any[]) { - return this.logger.logParts(level, message, ...args); - } - - newLogSession(name?: string) { - return this.logger.newSession(name); - } - - isDevMode() { - if (window.location.href.includes("rgthree-dev=false")) { - return false; - } - return GLOBAL_LOG_LEVEL >= LogLevel.DEBUG || window.location.href.includes("rgthree-dev"); - } - - isDebugMode() { - if (!this.isDevMode() || window.location.href.includes("rgthree-debug=false")) { - return false; - } - return window.location.href.includes("rgthree-debug"); - } - - monitorBadLinks() { - const badLinksFound = fixBadLinks(app.graph); - if (badLinksFound.hasBadLinks && !this.monitorBadLinksAlerted) { - this.monitorBadLinksAlerted = true; - alert( - `Problematic links just found in live data. Can you save your workflow and file a bug with ` + - `the last few steps you took to trigger this at ` + - `https://github.com/rgthree/rgthree-comfy/issues. Thank you!`, - ); - } else if (!badLinksFound.hasBadLinks) { - // Clear the alert once fixed so we can alert again. - this.monitorBadLinksAlerted = false; - } - this.monitorLinkTimeout = setTimeout(() => { - this.monitorBadLinks(); - }, 5000); - } -} - -function getBookmarks(): ContextMenuItem[] { - const graph: TLGraph = app.graph; - - // Sorts by Title. - // I could see an option to sort by either Shortcut, Title, or Position. - const bookmarks = graph._nodes - .filter((n): n is Bookmark => n.type === NodeTypesString.BOOKMARK) - .sort((a, b) => a.title.localeCompare(b.title)) - .map((n) => ({ - content: `[${n.shortcutKey}] ${n.title}`, - className: "rgthree-contextmenu-item", - callback: () => { - n.canvasToBookmark(); - }, - })); - - return !bookmarks.length - ? [] - : [ - { - content: "๐Ÿ”– Bookmarks", - disabled: true, - className: "rgthree-contextmenu-item rgthree-contextmenu-label", - }, - ...bookmarks, - ]; -} - -export const rgthree = new Rgthree(); -// Expose it on window because, why not. -(window as any).rgthree = rgthree; +import type { + LGraphCanvas as TLGraphCanvas, + LGraphNode, + SerializedLGraphNode, + serializedLGraph, + ContextMenuItem, + LGraph as TLGraph, + AdjustedMouseEvent, + IContextMenuOptions, +} from "typings/litegraph.js"; +import type { ComfyApiFormat, ComfyApiPrompt, ComfyApp } from "typings/comfy.js"; +import { app } from "scripts/app.js"; +import { api } from "scripts/api.js"; +import { SERVICE as CONFIG_SERVICE } from "./services/config_service.js"; +import { fixBadLinks } from "rgthree/common/link_fixer.js"; +import { injectCss, wait } from "rgthree/common/shared_utils.js"; +import { replaceNode, waitForCanvas, waitForGraph } from "./utils.js"; +import { NodeTypesString, addRgthree, getNodeTypeStrings, stripRgthree } from "./constants.js"; +import { RgthreeProgressBar } from "rgthree/common/progress_bar.js"; +import { RgthreeConfigDialog } from "./config.js"; +import { + iconGear, + iconNode, + iconReplace, + iconStarFilled, + logoRgthree, +} from "rgthree/common/media/svgs.js"; +import type { Bookmark } from "./bookmark.js"; +import { createElement, query, queryOne } from "rgthree/common/utils_dom.js"; + +export enum LogLevel { + IMPORTANT = 1, + ERROR, + WARN, + INFO, + DEBUG, + DEV, +} + +const LogLevelKeyToLogLevel: { [key: string]: LogLevel } = { + IMPORTANT: LogLevel.IMPORTANT, + ERROR: LogLevel.ERROR, + WARN: LogLevel.WARN, + INFO: LogLevel.INFO, + DEBUG: LogLevel.DEBUG, + DEV: LogLevel.DEV, +}; + +type ConsoleLogFns = "log" | "error" | "warn" | "debug" | "info"; +const LogLevelToMethod: { [key in LogLevel]: ConsoleLogFns } = { + [LogLevel.IMPORTANT]: "log", + [LogLevel.ERROR]: "error", + [LogLevel.WARN]: "warn", + [LogLevel.INFO]: "info", + [LogLevel.DEBUG]: "log", + [LogLevel.DEV]: "log", +}; +const LogLevelToCSS: { [key in LogLevel]: string } = { + [LogLevel.IMPORTANT]: "font-weight: bold; color: blue;", + [LogLevel.ERROR]: "", + [LogLevel.WARN]: "", + [LogLevel.INFO]: "font-style: italic; color: blue;", + [LogLevel.DEBUG]: "font-style: italic; color: #444;", + [LogLevel.DEV]: "color: #004b68;", +}; + +let GLOBAL_LOG_LEVEL = LogLevel.ERROR; + +/** + * A blocklist of extensions to disallow hooking into rgthree's base classes when calling the + * `rgthree.invokeExtensionsAsync` method (which runs outside of ComfyNode's + * `app.invokeExtensionsAsync` which is private). + * + * In Apr 2024 the base rgthree node class added support for other extensions using `nodeCreated` + * and `beforeRegisterNodeDef` which allows other extensions to modify the class. However, since it + * had been months since divorcing the ComfyNode in rgthree-comfy due to instability and + * inflexibility, this was a bit risky as other extensions hadn't ever run with this ability. This + * list attempts to block extensions from being able to call into rgthree-comfy nodes via the + * `nodeCreated` and `beforeRegisterNodeDef` callbacks now that rgthree-comfy is utilizing them + * because they do not work. Oddly, it's ComfyUI's own extension that is broken. + */ +const INVOKE_EXTENSIONS_BLOCKLIST = [ + { + name: "Comfy.WidgetInputs", + reason: + "Major conflict with rgthree-comfy nodes' inputs causing instability and " + + "repeated link disconnections.", + }, + { + name: "efficiency.widgethider", + reason: + "Overrides value getter before widget getter is prepared. Can be lifted if/when " + + "https://github.com/jags111/efficiency-nodes-comfyui/pull/203 is pulled.", + }, +]; + +/** A basic wrapper around logger. */ +class Logger { + /** Logs a message to the console if it meets the current log level. */ + log(level: LogLevel, message: string, ...args: any[]) { + const [n, v] = this.logParts(level, message, ...args); + console[n]?.(...v); + } + + /** + * Returns a tuple of the console function and its arguments. Useful for callers to make the + * actual console. call to gain benefits of DevTools knowing the source line. + * + * If the input is invalid or the level doesn't meet the configuration level, then the return + * value is an unknown function and empty set of values. Callers can use optionla chaining + * successfully: + * + * const [fn, values] = logger.logPars(LogLevel.INFO, 'my message'); + * console[fn]?.(...values); // Will work even if INFO won't be logged. + * + */ + logParts(level: LogLevel, message: string, ...args: any[]): [ConsoleLogFns, any[]] { + if (level <= GLOBAL_LOG_LEVEL) { + const css = LogLevelToCSS[level] || ""; + if (level === LogLevel.DEV) { + message = `๐Ÿ”ง ${message}`; + } + return [LogLevelToMethod[level], [`%c${message}`, css, ...args]]; + } + return ["none" as "info", []]; + } +} + +/** + * A log session, with the name as the prefix. A new session will stack prefixes. + */ +class LogSession { + readonly logger = new Logger(); + readonly logsCache: { [key: string]: { lastShownTime: number } } = {}; + + constructor(readonly name?: string) {} + + /** + * Returns the console log method to use and the arguments to pass so the call site can log from + * there. This extra work at the call site allows for easier debugging in the dev console. + * + * const [logMethod, logArgs] = logger.logParts(LogLevel.DEBUG, message, ...args); + * console[logMethod]?.(...logArgs); + */ + logParts(level: LogLevel, message?: string, ...args: any[]): [ConsoleLogFns, any[]] { + message = `${this.name || ""}${message ? " " + message : ""}`; + return this.logger.logParts(level, message, ...args); + } + + logPartsOnceForTime( + level: LogLevel, + time: number, + message?: string, + ...args: any[] + ): [ConsoleLogFns, any[]] { + message = `${this.name || ""}${message ? " " + message : ""}`; + const cacheKey = `${level}:${message}`; + const cacheEntry = this.logsCache[cacheKey]; + const now = +new Date(); + if (cacheEntry && cacheEntry.lastShownTime + time > now) { + return ["none" as "info", []]; + } + const parts = this.logger.logParts(level, message, ...args); + if (console[parts[0]]) { + this.logsCache[cacheKey] = this.logsCache[cacheKey] || ({} as { lastShownTime: number }); + this.logsCache[cacheKey]!.lastShownTime = now; + } + return parts; + } + + debugParts(message?: string, ...args: any[]) { + return this.logParts(LogLevel.DEBUG, message, ...args); + } + + infoParts(message?: string, ...args: any[]) { + return this.logParts(LogLevel.INFO, message, ...args); + } + + warnParts(message?: string, ...args: any[]) { + return this.logParts(LogLevel.WARN, message, ...args); + } + + newSession(name?: string) { + return new LogSession(`${this.name}${name}`); + } +} + +export type RgthreeUiMessage = { + id: string; + message: string; + type?: "warn" | "info" | "success" | null; + timeout?: number; + // closeable?: boolean; // TODO + actions?: Array<{ + label: string; + href?: string; + callback?: (event: MouseEvent) => void; + }>; +}; + +/** + * A global class as 'rgthree'; exposed on wiindow. Lots can go in here. + */ +class Rgthree extends EventTarget { + /** Exposes the ComfyUI api instance on rgthree. */ + readonly api = api; + private settingsDialog: RgthreeConfigDialog | null = null; + private progressBarEl: RgthreeProgressBar | null = null; + private rgthreeCssPromise: Promise; + + /** Stores a node id that we will use to queu only that output node (with `queueOutputNode`). */ + private queueNodeIds: number[] | null = null; + + logger = new LogSession("[rgthree]"); + + monitorBadLinksAlerted = false; + monitorLinkTimeout: number | null = null; + + processingQueue = false; + loadingApiJson = false; + replacingReroute: number | null = null; + processingMouseDown = false; + processingMouseUp = false; + processingMouseMove = false; + lastAdjustedMouseEvent: AdjustedMouseEvent | null = null; + + // Comfy/LiteGraph states so nodes and tell what the hell is going on. + canvasCurrentlyCopyingToClipboard = false; + canvasCurrentlyCopyingToClipboardWithMultipleNodes = false; + initialGraphToPromptSerializedWorkflowBecauseComfyUIBrokeStuff: any = null; + + private elDebugKeydowns: HTMLDivElement | null = null; + + private readonly isMac: boolean = !!( + navigator.platform?.toLocaleUpperCase().startsWith("MAC") || + (navigator as any).userAgentData?.platform?.toLocaleUpperCase().startsWith("MAC") + ); + + constructor() { + super(); + + const logLevel = + LogLevelKeyToLogLevel[CONFIG_SERVICE.getConfigValue("log_level")] ?? GLOBAL_LOG_LEVEL; + this.setLogLevel(logLevel); + + this.initializeGraphAndCanvasHooks(); + this.initializeComfyUIHooks(); + this.initializeContextMenu(); + + this.rgthreeCssPromise = injectCss("extensions/rgthree-comfy/rgthree.css"); + + this.initializeProgressBar(); + + CONFIG_SERVICE.addEventListener("config-change", ((e: CustomEvent) => { + if (e.detail?.key?.includes("features.progress_bar")) { + this.initializeProgressBar(); + } + }) as EventListener); + } + + /** + * Initializes the top progress bar, if it's configured. + */ + async initializeProgressBar() { + if (CONFIG_SERVICE.getConfigValue("features.progress_bar.enabled")) { + await this.rgthreeCssPromise; + if (!this.progressBarEl) { + this.progressBarEl = RgthreeProgressBar.create(); + this.progressBarEl.setAttribute( + "title", + "Progress Bar by rgthree. right-click for rgthree menu.", + ); + + this.progressBarEl.addEventListener("contextmenu", async (e) => { + e.stopPropagation(); + e.preventDefault(); + }); + + this.progressBarEl.addEventListener("pointerdown", async (e) => { + LiteGraph.closeAllContextMenus(); + if (e.button == 2) { + const canvas = await waitForCanvas(); + new LiteGraph.ContextMenu( + this.getRgthreeContextMenuItems(), + { + title: `
${logoRgthree} rgthree-comfy
`, + left: e.clientX, + top: 5, + }, + canvas.getCanvasWindow(), + ); + return; + } + if (e.button == 0) { + const nodeId = this.progressBarEl?.currentNodeId; + if (nodeId) { + const [canvas, graph] = await Promise.all([waitForCanvas(), waitForGraph()]); + const node = graph.getNodeById(Number(nodeId)); + if (node) { + canvas.centerOnNode(node); + e.stopPropagation(); + e.preventDefault(); + } + } + return; + } + }); + } + // Handle both cases in case someone hasn't updated. Can probably just assume + // `isUpdatedComfyBodyClasses` is true in the near future. + const isUpdatedComfyBodyClasses = !!queryOne(".comfyui-body-top"); + const position = CONFIG_SERVICE.getConfigValue("features.progress_bar.position"); + this.progressBarEl.classList.toggle("rgthree-pos-bottom", position === "bottom"); + // If ComfyUI is updated with the body segments, then use that. + if (isUpdatedComfyBodyClasses) { + if (position === "bottom") { + queryOne(".comfyui-body-bottom")!.appendChild(this.progressBarEl); + } else { + queryOne(".comfyui-body-top")!.appendChild(this.progressBarEl); + } + } else { + document.body.appendChild(this.progressBarEl); + } + const height = CONFIG_SERVICE.getConfigValue("features.progress_bar.height") || 14; + this.progressBarEl.style.height = `${height}px`; + const fontSize = Math.max(10, Number(height) - 10); + this.progressBarEl.style.fontSize = `${fontSize}px`; + this.progressBarEl.style.fontWeight = fontSize <= 12 ? "bold" : "normal"; + } else { + this.progressBarEl?.remove(); + } + } + + /** + * Initialize a bunch of hooks into LiteGraph itself so we can either keep state or context on + * what's happening so nodes can respond appropriately. This is usually to fix broken assumptions + * in the unowned code [๐Ÿคฎ], but sometimes to add features or enhancements too [โญ]. + */ + private async initializeGraphAndCanvasHooks() { + const rgthree = this; + + // [๐Ÿคฎ] To mitigate changes from https://github.com/rgthree/rgthree-comfy/issues/69 + // and https://github.com/comfyanonymous/ComfyUI/issues/2193 we can try to store the workflow + // node so our nodes can find the seralized node. Works with method + // `getNodeFromInitialGraphToPromptSerializedWorkflowBecauseComfyUIBrokeStuff` to find a node + // while serializing. What a way to work around... + const graphSerialize = LGraph.prototype.serialize; + LGraph.prototype.serialize = function () { + const response = graphSerialize.apply(this, [...arguments] as any) as any; + rgthree.initialGraphToPromptSerializedWorkflowBecauseComfyUIBrokeStuff = response; + return response; + }; + + // Overrides LiteGraphs' processMouseDown to both keep state as well as dispatch a custom event. + const processMouseDown = LGraphCanvas.prototype.processMouseDown; + LGraphCanvas.prototype.processMouseDown = function (e: AdjustedMouseEvent) { + rgthree.processingMouseDown = true; + const returnVal = processMouseDown.apply(this, [...arguments] as any); + rgthree.dispatchCustomEvent("on-process-mouse-down", { originalEvent: e }); + rgthree.processingMouseDown = false; + return returnVal; + }; + + // Overrides LiteGraph's `adjustMouseEvent` to capture the last even coming in and out. Useful + // to capture the last `canvasX` and `canvasY` properties, which are not the same as LiteGraph's + // `canvas.last_mouse_position`, unfortunately. + const adjustMouseEvent = LGraphCanvas.prototype.adjustMouseEvent; + LGraphCanvas.prototype.adjustMouseEvent = function (e: PointerEvent) { + adjustMouseEvent.apply(this, [...arguments] as any); + rgthree.lastAdjustedMouseEvent = e as AdjustedMouseEvent; + }; + + // [๐Ÿคฎ] Copying to clipboard clones nodes and then manipulats the linking data manually which + // does not allow a node to handle connections. This harms nodes that manually handle inputs, + // like our any-input nodes that may start with one input, and manually add new ones when one is + // attached. + const copyToClipboard = LGraphCanvas.prototype.copyToClipboard; + LGraphCanvas.prototype.copyToClipboard = function (nodes: LGraphNode[]) { + rgthree.canvasCurrentlyCopyingToClipboard = true; + rgthree.canvasCurrentlyCopyingToClipboardWithMultipleNodes = + Object.values(nodes || this.selected_nodes || []).length > 1; + copyToClipboard.apply(this, [...arguments] as any); + rgthree.canvasCurrentlyCopyingToClipboard = false; + rgthree.canvasCurrentlyCopyingToClipboardWithMultipleNodes = false; + }; + + // [โญ] Make it so when we add a group, we get to name it immediately. + const onGroupAdd = LGraphCanvas.onGroupAdd; + LGraphCanvas.onGroupAdd = function (...args: any[]) { + const graph = app.graph as TLGraph; + onGroupAdd.apply(this, [...args] as any); + LGraphCanvas.onShowPropertyEditor( + {}, + null, + null, + null, + graph._groups[graph._groups.length - 1], + ); + }; + } + + /** + * [๐Ÿคฎ] Handles the same exact thing as ComfyApp's `invokeExtensionsAsync`, but done here since + * it is #private in ComfyApp because... of course it us. This is necessary since we purposefully + * avoid using the ComfyNode due to historical instability and inflexibility for all the advanced + * ui stuff rgthree-comfy nodes do, but we can still have other custom nodes know what's happening + * with rgthree-comfy; specifically, for `nodeCreated` as of now. + */ + async invokeExtensionsAsync(method: "nodeCreated", ...args: any[]) { + const comfyapp = app as ComfyApp; + if (CONFIG_SERVICE.getConfigValue("features.invoke_extensions_async.node_created") === false) { + const [m, a] = this.logParts( + LogLevel.INFO, + `Skipping invokeExtensionsAsync for applicable rgthree-comfy nodes`, + ); + console[m]?.(...a); + return Promise.resolve(); + } + return await Promise.all( + comfyapp.extensions.map(async (ext) => { + if (ext?.[method]) { + try { + const blocked = INVOKE_EXTENSIONS_BLOCKLIST.find((block) => + ext.name.toLowerCase().startsWith(block.name.toLowerCase()), + ); + if (blocked) { + const [n, v] = this.logger.logPartsOnceForTime( + LogLevel.WARN, + 5000, + `Blocked extension '${ext.name}' method '${method}' for rgthree-nodes because: ${blocked.reason}`, + ); + console[n]?.(...v); + return Promise.resolve(); + } + return await (ext[method] as Function)(...args, comfyapp); + } catch (error) { + const [n, v] = this.logParts( + LogLevel.ERROR, + `Error calling extension '${ext.name}' method '${method}' for rgthree-node.`, + { error }, + { extension: ext }, + { args }, + ); + console[n]?.(...v); + } + } + }), + ); + } + + /** + * Wraps `dispatchEvent` for easier CustomEvent dispatching. + */ + private dispatchCustomEvent(event: string, detail?: any) { + if (detail != null) { + return this.dispatchEvent(new CustomEvent(event, { detail })); + } + return this.dispatchEvent(new CustomEvent(event)); + } + + /** + * Initializes hooks specific to an rgthree-comfy context menu on the root menu. + */ + private async initializeContextMenu() { + const that = this; + setTimeout(async () => { + const getCanvasMenuOptions = LGraphCanvas.prototype.getCanvasMenuOptions; + LGraphCanvas.prototype.getCanvasMenuOptions = function (...args: any[]) { + let existingOptions = getCanvasMenuOptions.apply(this, [...args] as any); + + const options = []; + options.push(null); // Divider + options.push(null); // Divider + options.push(null); // Divider + options.push({ + content: logoRgthree + `rgthree-comfy`, + className: "rgthree-contextmenu-item rgthree-contextmenu-main-item-rgthree-comfy", + submenu: { + options: that.getRgthreeContextMenuItems(), + }, + }); + options.push(null); // Divider + options.push(null); // Divider + + let idx = null; + idx = idx || existingOptions.findIndex((o) => o?.content?.startsWith?.("Queue Group")) + 1; + idx = + idx || existingOptions.findIndex((o) => o?.content?.startsWith?.("Queue Selected")) + 1; + idx = idx || existingOptions.findIndex((o) => o?.content?.startsWith?.("Convert to Group")); + idx = idx || existingOptions.findIndex((o) => o?.content?.startsWith?.("Arrange (")); + idx = idx || existingOptions.findIndex((o) => !o) + 1; + idx = idx || 3; + existingOptions.splice(idx, 0, ...options); + for (let i = existingOptions.length; i > 0; i--) { + if (existingOptions[i] === null && existingOptions[i + 1] === null) { + existingOptions.splice(i, 1); + } + } + + return existingOptions; + }; + }, 1016); + } + + /** + * Returns the standard menu items for an rgthree-comfy context menu. + */ + private getRgthreeContextMenuItems(): ContextMenuItem[] { + const [canvas, graph] = [app.canvas as TLGraphCanvas, app.graph as TLGraph]; + const selectedNodes = Object.values(canvas.selected_nodes || {}); + let rerouteNodes: LGraphNode[] = []; + if (selectedNodes.length) { + rerouteNodes = selectedNodes.filter((n) => n.type === "Reroute"); + } else { + rerouteNodes = graph._nodes.filter((n) => n.type == "Reroute"); + } + const rerouteLabel = selectedNodes.length ? "selected" : "all"; + + const showBookmarks = CONFIG_SERVICE.getFeatureValue("menu_bookmarks.enabled"); + const bookmarkMenuItems = showBookmarks ? getBookmarks() : []; + + return [ + { + content: "Nodes", + disabled: true, + className: "rgthree-contextmenu-item rgthree-contextmenu-label", + }, + { + content: iconNode + "All", + className: "rgthree-contextmenu-item", + has_submenu: true, + submenu: { + options: getNodeTypeStrings() as unknown as ContextMenuItem[], + callback: ( + value: string | ContextMenuItem, + options: IContextMenuOptions, + event: MouseEvent, + ) => { + const node = LiteGraph.createNode(addRgthree(value as string)); + node.pos = [ + rgthree.lastAdjustedMouseEvent!.canvasX, + rgthree.lastAdjustedMouseEvent!.canvasY, + ]; + canvas.graph.add(node); + canvas.selectNode(node); + app.graph.setDirtyCanvas(true, true); + }, + extra: { rgthree_doNotNest: true }, + }, + }, + + { + content: "Actions", + disabled: true, + className: "rgthree-contextmenu-item rgthree-contextmenu-label", + }, + { + content: iconGear + "Settings (rgthree-comfy)", + disabled: !!this.settingsDialog, + className: "rgthree-contextmenu-item", + callback: (...args: any[]) => { + this.settingsDialog = new RgthreeConfigDialog().show(); + this.settingsDialog.addEventListener("close", (e) => { + this.settingsDialog = null; + }); + }, + }, + { + content: iconReplace + ` Convert ${rerouteLabel} Reroutes`, + disabled: !rerouteNodes.length, + className: "rgthree-contextmenu-item", + callback: (...args: any[]) => { + const msg = + `Convert ${rerouteLabel} ComfyUI Reroutes to Reroute (rgthree) nodes? \n` + + `(First save a copy of your workflow & check reroute connections afterwards)`; + if (!window.confirm(msg)) { + return; + } + (async () => { + for (const node of [...rerouteNodes]) { + if (node.type == "Reroute") { + this.replacingReroute = node.id; + await replaceNode(node, NodeTypesString.REROUTE); + this.replacingReroute = null; + } + } + })(); + }, + }, + ...bookmarkMenuItems, + { + content: "More...", + disabled: true, + className: "rgthree-contextmenu-item rgthree-contextmenu-label", + }, + { + content: iconStarFilled + "Star on Github", + className: "rgthree-contextmenu-item rgthree-contextmenu-github", + callback: (...args: any[]) => { + window.open("https://github.com/rgthree/rgthree-comfy", "_blank"); + }, + }, + ]; + } + + /** + * Wraps an `app.queuePrompt` call setting a specific node id that we will inspect and change the + * serialized graph right before being sent (below, in our `api.queuePrompt` override). + */ + async queueOutputNodes(nodeIds: number[]) { + try { + this.queueNodeIds = nodeIds; + await app.queuePrompt(); + } catch (e) { + const [n, v] = this.logParts( + LogLevel.ERROR, + `There was an error queuing nodes ${nodeIds}`, + e, + ); + console[n]?.(...v); + } finally { + this.queueNodeIds = null; + } + } + + /** + * Recusively walks backwards from a node adding its inputs to the `newOutput` from `oldOutput`. + */ + private recursiveAddNodes(nodeId: string, oldOutput: ComfyApiFormat, newOutput: ComfyApiFormat) { + let currentId = nodeId; + let currentNode = oldOutput[currentId]!; + if (newOutput[currentId] == null) { + newOutput[currentId] = currentNode; + for (const inputValue of Object.values(currentNode.inputs || [])) { + if (Array.isArray(inputValue)) { + this.recursiveAddNodes(inputValue[0], oldOutput, newOutput); + } + } + } + return newOutput; + } + + /** + * Initialize a bunch of hooks into ComfyUI and/or LiteGraph itself so we can either keep state or + * context on what's happening so nodes can respond appropriately. This is usually to fix broken + * assumptions in the unowned code [๐Ÿคฎ], but sometimes to add features or enhancements too [โญ]. + */ + private initializeComfyUIHooks() { + const rgthree = this; + + // Keep state for when the app is queuing the prompt. For instance, this is used for seed to + // understand if we're serializing because we're queueing (and return the random seed to use) or + // for saving the workflow (and keep -1, etc.). + const queuePrompt = app.queuePrompt as Function; + app.queuePrompt = async function () { + rgthree.processingQueue = true; + rgthree.dispatchCustomEvent("queue"); + try { + await queuePrompt.apply(app, [...arguments]); + } finally { + rgthree.processingQueue = false; + rgthree.dispatchCustomEvent("queue-end"); + } + }; + + // Keep state for when the app is in the middle of loading from an api JSON file. + const loadApiJson = app.loadApiJson; + app.loadApiJson = async function () { + rgthree.loadingApiJson = true; + try { + loadApiJson.apply(app, [...arguments] as any); + } finally { + rgthree.loadingApiJson = false; + } + }; + + // Keep state for when the app is serizalizing the graph to prompt. + const graphToPrompt = app.graphToPrompt; + app.graphToPrompt = async function () { + rgthree.dispatchCustomEvent("graph-to-prompt"); + let promise = graphToPrompt.apply(app, [...arguments] as any); + await promise; + rgthree.dispatchCustomEvent("graph-to-prompt-end"); + return promise; + }; + + // Override the queuePrompt for api to intercept the prompt output and, if queueNodeIds is set, + // then we only want to queue those nodes, by rewriting the api format (prompt 'output' field) + // so only those are evaluated. + const apiQueuePrompt = api.queuePrompt as Function; + api.queuePrompt = async function (index: number, prompt: ComfyApiPrompt) { + if (rgthree.queueNodeIds?.length && prompt.output) { + const oldOutput = prompt.output; + let newOutput = {}; + for (const queueNodeId of rgthree.queueNodeIds) { + rgthree.recursiveAddNodes(String(queueNodeId), oldOutput, newOutput); + } + prompt.output = newOutput; + } + rgthree.dispatchCustomEvent("comfy-api-queue-prompt-before", { + workflow: prompt.workflow, + output: prompt.output, + }); + const response = apiQueuePrompt.apply(app, [index, prompt]); + rgthree.dispatchCustomEvent("comfy-api-queue-prompt-end"); + return response; + }; + + // Hook into a clean call; allow us to clear and rgthree messages. + const clean = app.clean; + app.clean = function () { + rgthree.clearAllMessages(); + clean && clean.apply(app, [...arguments] as any); + }; + + // Hook into a data load, like from an image or JSON drop-in. This is (currently) used to + // monitor for bad linking data. + const loadGraphData = app.loadGraphData; + app.loadGraphData = function (graph: serializedLGraph) { + if (rgthree.monitorLinkTimeout) { + clearTimeout(rgthree.monitorLinkTimeout); + rgthree.monitorLinkTimeout = null; + } + rgthree.clearAllMessages(); + // Try to make a copy to use, because ComfyUI's loadGraphData will modify it. + let graphCopy: serializedLGraph | null; + try { + graphCopy = JSON.parse(JSON.stringify(graph)); + } catch (e) { + graphCopy = null; + } + setTimeout(() => { + const wasLoadingAborted = document + .querySelector(".comfy-modal-content") + ?.textContent?.includes("Loading aborted due"); + const graphToUse = wasLoadingAborted ? graphCopy || graph : app.graph; + const fixBadLinksResult = fixBadLinks(graphToUse as unknown as TLGraph); + if (fixBadLinksResult.hasBadLinks) { + const [n, v] = rgthree.logParts( + LogLevel.WARN, + `The workflow you've loaded has corrupt linking data. Open ${ + new URL(location.href).origin + }/rgthree/link_fixer to try to fix.`, + ); + console[n]?.(...v); + if (CONFIG_SERVICE.getConfigValue("features.show_alerts_for_corrupt_workflows")) { + rgthree.showMessage({ + id: "bad-links", + type: "warn", + message: + "The workflow you've loaded has corrupt linking data that may be able to be fixed.", + actions: [ + { + label: "Open fixer", + href: "/rgthree/link_fixer", + }, + { + label: "Fix in place", + href: "/rgthree/link_fixer", + callback: (event) => { + event.stopPropagation(); + event.preventDefault(); + if ( + confirm( + "This will attempt to fix in place. Please make sure to have a saved copy of your workflow.", + ) + ) { + try { + const fixBadLinksResult = fixBadLinks( + graphToUse as unknown as TLGraph, + true, + ); + if (!fixBadLinksResult.hasBadLinks) { + rgthree.hideMessage("bad-links"); + alert( + "Success! It's possible some valid links may have been affected. Please check and verify your workflow.", + ); + wasLoadingAborted && app.loadGraphData(fixBadLinksResult.graph); + if ( + CONFIG_SERVICE.getConfigValue("features.monitor_for_corrupt_links") || + CONFIG_SERVICE.getConfigValue("features.monitor_bad_links") + ) { + rgthree.monitorLinkTimeout = setTimeout(() => { + rgthree.monitorBadLinks(); + }, 5000); + } + } + } catch (e) { + console.error(e); + alert("Unsuccessful at fixing corrupt data. :("); + rgthree.hideMessage("bad-links"); + } + } + }, + }, + ], + }); + } + } else if ( + CONFIG_SERVICE.getConfigValue("features.monitor_for_corrupt_links") || + CONFIG_SERVICE.getConfigValue("features.monitor_bad_links") + ) { + rgthree.monitorLinkTimeout = setTimeout(() => { + rgthree.monitorBadLinks(); + }, 5000); + } + }, 100); + return loadGraphData && loadGraphData.apply(app, [...arguments] as any); + }; + } + + /** + * [๐Ÿคฎ] Finds a node in the currently serializing workflow from the hook setup above. This is to + * mitigate breakages from https://github.com/comfyanonymous/ComfyUI/issues/2193 we can try to + * store the workflow node so our nodes can find the seralized node. + */ + getNodeFromInitialGraphToPromptSerializedWorkflowBecauseComfyUIBrokeStuff( + node: LGraphNode, + ): SerializedLGraphNode | null { + return ( + this.initialGraphToPromptSerializedWorkflowBecauseComfyUIBrokeStuff?.nodes?.find( + (n: SerializedLGraphNode) => n.id === node.id, + ) ?? null + ); + } + + /** + * Shows a message in the UI. + */ + async showMessage(data: RgthreeUiMessage) { + let container = document.querySelector(".rgthree-top-messages-container"); + if (!container) { + container = document.createElement("div"); + container.classList.add("rgthree-top-messages-container"); + document.body.appendChild(container); + } + // If we have a dialog open then we want to append the message to the dialog so they show over + // the modal. + const dialogs = query("dialog[open]"); + if (dialogs.length) { + let dialog = dialogs[dialogs.length - 1]!; + dialog.appendChild(container); + dialog.addEventListener("close", (e) => { + document.body.appendChild(container!); + }); + } + // Hide if we exist. + await this.hideMessage(data.id); + + const messageContainer = document.createElement("div"); + messageContainer.setAttribute("type", data.type || "info"); + + const message = document.createElement("span"); + message.innerHTML = data.message; + messageContainer.appendChild(message); + + for (let a = 0; a < (data.actions || []).length; a++) { + const action = data.actions![a]!; + if (a > 0) { + const sep = document.createElement("span"); + sep.innerHTML = " | "; + messageContainer.appendChild(sep); + } + + const actionEl = document.createElement("a"); + actionEl.innerText = action.label; + if (action.href) { + actionEl.target = "_blank"; + actionEl.href = action.href; + } + if (action.callback) { + actionEl.onclick = (e) => { + return action.callback!(e); + }; + } + messageContainer.appendChild(actionEl); + } + + const messageAnimContainer = document.createElement("div"); + messageAnimContainer.setAttribute("msg-id", data.id); + messageAnimContainer.appendChild(messageContainer); + container.appendChild(messageAnimContainer); + + // Add. Wait. Measure. Wait. Anim. + await wait(64); + messageAnimContainer.style.marginTop = `-${messageAnimContainer.offsetHeight}px`; + await wait(64); + messageAnimContainer.classList.add("-show"); + + if (data.timeout) { + await wait(data.timeout); + this.hideMessage(data.id); + } + } + + /** + * Hides a message in the UI. + */ + async hideMessage(id: string) { + const msg = document.querySelector(`.rgthree-top-messages-container > [msg-id="${id}"]`); + if (msg?.classList.contains("-show")) { + msg.classList.remove("-show"); + await wait(750); + } + msg && msg.remove(); + } + + /** + * Clears all messages in the UI. + */ + async clearAllMessages() { + let container = document.querySelector(".rgthree-top-messages-container"); + container && (container.innerHTML = ""); + } + + setLogLevel(level?: LogLevel | string) { + if (typeof level === "string") { + level = LogLevelKeyToLogLevel[CONFIG_SERVICE.getConfigValue("log_level")]; + } + if (level != null) { + GLOBAL_LOG_LEVEL = level; + } + } + + logParts(level: LogLevel, message?: string, ...args: any[]) { + return this.logger.logParts(level, message, ...args); + } + + newLogSession(name?: string) { + return this.logger.newSession(name); + } + + isDevMode() { + if (window.location.href.includes("rgthree-dev=false")) { + return false; + } + return GLOBAL_LOG_LEVEL >= LogLevel.DEBUG || window.location.href.includes("rgthree-dev"); + } + + isDebugMode() { + if (!this.isDevMode() || window.location.href.includes("rgthree-debug=false")) { + return false; + } + return window.location.href.includes("rgthree-debug"); + } + + monitorBadLinks() { + const badLinksFound = fixBadLinks(app.graph); + if (badLinksFound.hasBadLinks && !this.monitorBadLinksAlerted) { + this.monitorBadLinksAlerted = true; + alert( + `Problematic links just found in live data. Can you save your workflow and file a bug with ` + + `the last few steps you took to trigger this at ` + + `https://github.com/rgthree/rgthree-comfy/issues. Thank you!`, + ); + } else if (!badLinksFound.hasBadLinks) { + // Clear the alert once fixed so we can alert again. + this.monitorBadLinksAlerted = false; + } + this.monitorLinkTimeout = setTimeout(() => { + this.monitorBadLinks(); + }, 5000); + } +} + +function getBookmarks(): ContextMenuItem[] { + const graph: TLGraph = app.graph; + + // Sorts by Title. + // I could see an option to sort by either Shortcut, Title, or Position. + const bookmarks = graph._nodes + .filter((n): n is Bookmark => n.type === NodeTypesString.BOOKMARK) + .sort((a, b) => a.title.localeCompare(b.title)) + .map((n) => ({ + content: `[${n.shortcutKey}] ${n.title}`, + className: "rgthree-contextmenu-item", + callback: () => { + n.canvasToBookmark(); + }, + })); + + return !bookmarks.length + ? [] + : [ + { + content: "๐Ÿ”– Bookmarks", + disabled: true, + className: "rgthree-contextmenu-item rgthree-contextmenu-label", + }, + ...bookmarks, + ]; +} + +export const rgthree = new Rgthree(); +// Expose it on window because, why not. +(window as any).rgthree = rgthree; diff --git a/src_web/comfyui/services/context_service.ts b/src_web/comfyui/services/context_service.ts new file mode 100644 index 0000000..0a0eeff --- /dev/null +++ b/src_web/comfyui/services/context_service.ts @@ -0,0 +1,76 @@ +import type {DynamicContextNodeBase} from "../dynamic_context_base.js"; + +import {app} from "scripts/app.js"; +import {NodeTypesString} from "../constants.js"; +import {getConnectedOutputNodesAndFilterPassThroughs} from "../utils.js"; +import {INodeInputSlot, INodeOutputSlot, INodeSlot, LGraphNode} from "typings/litegraph.js"; + +export let SERVICE: ContextService; + +const OWNED_PREFIX = "+"; +const REGEX_PREFIX = /^[\+โš ๏ธ]\s*/; +const REGEX_EMPTY_INPUT = /^\+\s*$/; + +export function stripContextInputPrefixes(name: string) { + return name.replace(REGEX_PREFIX, ""); +} + +export function getContextOutputName(inputName: string) { + if (inputName === "base_ctx") return "CONTEXT"; + return stripContextInputPrefixes(inputName).toUpperCase(); +} + +export enum InputMutationOperation { + "UNKNOWN", + "ADDED", + "REMOVED", + "RENAMED", +} + +export type InputMutation = { + operation: InputMutationOperation; + node: DynamicContextNodeBase; + slotIndex: number; + slot: INodeSlot; +}; + +export class ContextService { + + constructor() { + if (SERVICE) { + throw new Error("ContextService was already instantiated."); + } + } + + onInputChanges(node: any, mutation: InputMutation) { + const childCtxs = getConnectedOutputNodesAndFilterPassThroughs( + node, + node, + 0, + ) as DynamicContextNodeBase[]; + for (const childCtx of childCtxs) { + childCtx.handleUpstreamMutation(mutation); + } + } + + getDynamicContextInputsData(node: DynamicContextNodeBase) { + return node + .getContextInputsList() + .map((input: INodeInputSlot, index: number) => ({ + name: stripContextInputPrefixes(input.name), + type: String(input.type), + index, + })) + .filter((i) => i.type !== "*"); + } + + getDynamicContextOutputsData(node: LGraphNode) { + return node.outputs.map((output: INodeOutputSlot, index: number) => ({ + name: stripContextInputPrefixes(output.name), + type: String(output.type), + index, + })); + } +} + +SERVICE = new ContextService(); diff --git a/src_web/comfyui/tests/context_dynamic_tests.ts b/src_web/comfyui/tests/context_dynamic_tests.ts new file mode 100644 index 0000000..90d882e --- /dev/null +++ b/src_web/comfyui/tests/context_dynamic_tests.ts @@ -0,0 +1,191 @@ +import type { + LiteGraph as TLiteGraph, + LGraphCanvas as TLGraphCanvas, + LGraph as TLGraph, + LGraphNode as TLGraphNode, + Vector2, + LGraphNode, +} from "typings/litegraph.js"; +import {rgthree} from "../rgthree.js"; +import {NodeTypesString} from "../constants.js"; +import {wait} from "rgthree/common/shared_utils.js"; +import {describe, should, beforeEach, expect, describeRun} from "../testing/runner.js"; +import {ComfyUITestEnvironment} from "../testing/comfyui_env.js"; + +declare const LiteGraph: typeof TLiteGraph; + +const env = new ComfyUITestEnvironment(); + +function verifyInputAndOutputName( + node: LGraphNode, + index: number, + inputName: string | null, + isLinked?: boolean, +) { + if (inputName != null) { + expect(node.inputs[index]!.name).toBe(`input ${index} name`, inputName); + } + if (isLinked) { + expect(node.inputs[index]!.link).toBeANumber(`input ${index} connection`); + } else if (isLinked === false) { + expect(node.inputs[index]!.link).toBeNullOrUndefined(`input ${index} connection`); + } + if (inputName != null) { + if (inputName === "+") { + expect(node.outputs[index]).toBeUndefined(`output ${index}`); + } else { + let outputName = + inputName === "base_ctx" ? "CONTEXT" : inputName.replace(/^\+\s/, "").toUpperCase(); + expect(node.outputs[index]!.name).toBe(`output ${index} name`, outputName); + } + } +} + +function vertifyInputsStructure(node: LGraphNode, expectedLength: number) { + expect(node.inputs.length).toBe("inputs length", expectedLength); + expect(node.outputs.length).toBe("outputs length", expectedLength - 1); + verifyInputAndOutputName(node, expectedLength - 1, "+", false); +} + +(window as any).rgthree_tests = (window as any).rgthree_tests || {}; +(window as any).rgthree_tests.test_dynamic_context = describe("ContextDynamicTest", async () => { + let nodeConfig!: TLGraphNode; + let nodeCtx!: TLGraphNode; + + let lastNode: LGraphNode | null = null; + + await beforeEach(async () => { + await env.clear(); + lastNode = nodeConfig = await env.addNode(NodeTypesString.KSAMPLER_CONFIG); + lastNode = nodeCtx = await env.addNode(NodeTypesString.DYNAMIC_CONTEXT); + nodeConfig.connect(0, nodeCtx, 1); // steps + nodeConfig.connect(2, nodeCtx, 2); // cfg + nodeConfig.connect(4, nodeCtx, 3); // scheduler + nodeConfig.connect(0, nodeCtx, 4); // This is the step.1 + nodeConfig.connect(0, nodeCtx, 5); // This is the step.2 + nodeCtx.disconnectInput(2); + nodeCtx.disconnectInput(5); + nodeConfig.connect(0, nodeCtx, 6); // This is the step.3 + nodeCtx.disconnectInput(6); + await wait(); + }); + + await should("add correct inputs", async () => { + vertifyInputsStructure(nodeCtx, 8); + let i = 0; + verifyInputAndOutputName(nodeCtx, i++, "base_ctx", false); + verifyInputAndOutputName(nodeCtx, i++, "+ steps", true); + verifyInputAndOutputName(nodeCtx, i++, "+ cfg", false); + verifyInputAndOutputName(nodeCtx, i++, "+ scheduler", true); + verifyInputAndOutputName(nodeCtx, i++, "+ steps.1", true); + verifyInputAndOutputName(nodeCtx, i++, "+ steps.2", false); + verifyInputAndOutputName(nodeCtx, i++, "+ steps.3", false); + }); + + await should("add evaluate correct outputs", async () => { + const displayAny1 = await env.addNode(NodeTypesString.DISPLAY_ANY, {placement: "right"}); + const displayAny2 = await env.addNode(NodeTypesString.DISPLAY_ANY, {placement: "under"}); + const displayAny3 = await env.addNode(NodeTypesString.DISPLAY_ANY, {placement: "under"}); + const displayAny4 = await env.addNode(NodeTypesString.DISPLAY_ANY, {placement: "under"}); + + nodeCtx.connect(1, displayAny1, 0); // steps + nodeCtx.connect(3, displayAny2, 0); // scheduler + nodeCtx.connect(4, displayAny3, 0); // steps.1 + nodeCtx.connect(6, displayAny4, 0); // steps.3 (unlinked) + + await env.queuePrompt(); + + expect(displayAny1.widgets![0]!.value).toBe("output 1", 30); + expect(displayAny2.widgets![0]!.value).toBe("output 3", '"normal"'); + expect(displayAny3.widgets![0]!.value).toBe("output 4", 30); + expect(displayAny4.widgets![0]!.value).toBe("output 6", "None"); + }); + + await describeRun("Nested", async () => { + let nodeConfig2!: TLGraphNode; + let nodeCtx2!: TLGraphNode; + + await beforeEach(async () => { + nodeConfig2 = await env.addNode(NodeTypesString.KSAMPLER_CONFIG, {placement: "start"}); + nodeConfig2.widgets[0]!.value = 111; + nodeConfig2.widgets[2]!.value = 11.1; + nodeCtx2 = await env.addNode(NodeTypesString.DYNAMIC_CONTEXT, {placement: "right"}); + nodeConfig2.connect(0, nodeCtx2, 1); // steps + nodeConfig2.connect(2, nodeCtx2, 2); // cfg + nodeConfig2.connect(3, nodeCtx2, 3); // sampler + nodeConfig2.connect(2, nodeCtx2, 4); // This is the cfg.1 + nodeConfig2.connect(0, nodeCtx2, 5); // This is the steps.1 + nodeCtx2.disconnectInput(2); + nodeCtx2.disconnectInput(5); + nodeConfig2.connect(2, nodeCtx2, 6); // This is the cfg.2 + nodeCtx2.disconnectInput(6); + + await wait(); + }); + + await should("disallow context node to be connected to non-first spot.", async () => { + // Connect to first node. + let expectedInputs = 8; + + nodeCtx2.connect(0, nodeCtx, expectedInputs - 1); + console.log(nodeCtx.inputs); + + vertifyInputsStructure(nodeCtx, expectedInputs); + verifyInputAndOutputName(nodeCtx, 0, "base_ctx", false); + verifyInputAndOutputName(nodeCtx, nodeCtx.inputs.length - 1, null, false); + + nodeCtx2.connect(0, nodeCtx, 0); + expectedInputs = 14; + vertifyInputsStructure(nodeCtx, expectedInputs); + verifyInputAndOutputName(nodeCtx, 0, "base_ctx", true); + verifyInputAndOutputName(nodeCtx, expectedInputs - 1, null, false); + }); + + await should("add inputs from connected above owned.", async () => { + // Connect to first node. + nodeCtx2.connect(0, nodeCtx, 0); + + let expectedInputs = 14; + vertifyInputsStructure(nodeCtx, expectedInputs); + let i = 0; + verifyInputAndOutputName(nodeCtx, i++, "base_ctx", true); + verifyInputAndOutputName(nodeCtx, i++, "steps", false); + verifyInputAndOutputName(nodeCtx, i++, "cfg", false); + verifyInputAndOutputName(nodeCtx, i++, "sampler", false); + verifyInputAndOutputName(nodeCtx, i++, "cfg.1", false); + verifyInputAndOutputName(nodeCtx, i++, "steps.1", false); + verifyInputAndOutputName(nodeCtx, i++, "cfg.2", false); + verifyInputAndOutputName(nodeCtx, i++, "+ steps.2", true); + verifyInputAndOutputName(nodeCtx, i++, "+ cfg.3", false); + verifyInputAndOutputName(nodeCtx, i++, "+ scheduler", true); + verifyInputAndOutputName(nodeCtx, i++, "+ steps.3", true); + verifyInputAndOutputName(nodeCtx, i++, "+ steps.4", false); + verifyInputAndOutputName(nodeCtx, i++, "+ steps.5", false); + verifyInputAndOutputName(nodeCtx, i++, "+", false); + }); + + await should("add then remove inputs when disconnected.", async () => { + // Connect to first node. + nodeCtx2.connect(0, nodeCtx, 0); + + let expectedInputs = 14; + expect(nodeCtx.inputs.length).toBe("inputs length", expectedInputs); + expect(nodeCtx.outputs.length).toBe("outputs length", expectedInputs - 1); + + nodeCtx.disconnectInput(0); + + expectedInputs = 8; + expect(nodeCtx.inputs.length).toBe("inputs length", expectedInputs); + expect(nodeCtx.outputs.length).toBe("outputs length", expectedInputs - 1); + let i = 0; + verifyInputAndOutputName(nodeCtx, i++, "base_ctx", false); + verifyInputAndOutputName(nodeCtx, i++, "+ steps", true); + verifyInputAndOutputName(nodeCtx, i++, "+ cfg", false); + verifyInputAndOutputName(nodeCtx, i++, "+ scheduler", true); + verifyInputAndOutputName(nodeCtx, i++, "+ steps.1", true); + verifyInputAndOutputName(nodeCtx, i++, "+ steps.2", false); + verifyInputAndOutputName(nodeCtx, i++, "+ steps.3", false); + verifyInputAndOutputName(nodeCtx, i++, "+", false); + }); + }); +}); diff --git a/src_web/comfyui/utils.ts b/src_web/comfyui/utils.ts index 774b6bf..1cb1a9a 100644 --- a/src_web/comfyui/utils.ts +++ b/src_web/comfyui/utils.ts @@ -1,927 +1,927 @@ -import type { ComfyApp, ComfyNodeConstructor, ComfyObjectInfo } from "typings/comfy"; -import type { - Vector2, - LGraphCanvas, - ContextMenuItem, - LLink, - LGraph, - IContextMenuOptions, - ContextMenu, - LGraphNode, - INodeSlot, - INodeInputSlot, - INodeOutputSlot, -} from "typings/litegraph.js"; -import type { Constructor } from "typings/index.js"; -import { app } from "scripts/app.js"; -import { api } from "scripts/api.js"; -import { Resolver, getResolver, wait } from "rgthree/common/shared_utils.js"; -import { RgthreeHelpDialog } from "rgthree/common/dialog.js"; - -/** - * Override the api.getNodeDefs call to add a hook for refreshing node defs. - * This is necessary for power prompt's custom combos. Since API implements - * add/removeEventListener already, this is rather trivial. - */ -const oldApiGetNodeDefs = api.getNodeDefs; -api.getNodeDefs = async function () { - const defs = await oldApiGetNodeDefs.call(api); - this.dispatchEvent(new CustomEvent("fresh-node-defs", { detail: defs })); - return defs; -}; - -export enum IoDirection { - INPUT, - OUTPUT, -} - -const PADDING = 0; - -type LiteGraphDir = - | typeof LiteGraph.LEFT - | typeof LiteGraph.RIGHT - | typeof LiteGraph.UP - | typeof LiteGraph.DOWN; -export const LAYOUT_LABEL_TO_DATA: { [label: string]: [LiteGraphDir, Vector2, Vector2] } = { - Left: [LiteGraph.LEFT, [0, 0.5], [PADDING, 0]], - Right: [LiteGraph.RIGHT, [1, 0.5], [-PADDING, 0]], - Top: [LiteGraph.UP, [0.5, 0], [0, PADDING]], - Bottom: [LiteGraph.DOWN, [0.5, 1], [0, -PADDING]], -}; -export const LAYOUT_LABEL_OPPOSITES: { [label: string]: string } = { - Left: "Right", - Right: "Left", - Top: "Bottom", - Bottom: "Top", -}; -export const LAYOUT_CLOCKWISE = ["Top", "Right", "Bottom", "Left"]; - -interface MenuConfig { - name: string | ((node: LGraphNode) => string); - property?: string; - prepareValue?: (value: string, node: LGraphNode) => any; - callback?: (node: LGraphNode, value?: string) => void; - subMenuOptions?: (string | null)[] | ((node: LGraphNode) => (string | null)[]); -} - -export function addMenuItem( - node: Constructor, - _app: ComfyApp, - config: MenuConfig, - after = "Shape", -) { - const oldGetExtraMenuOptions = node.prototype.getExtraMenuOptions; - node.prototype.getExtraMenuOptions = function ( - canvas: LGraphCanvas, - menuOptions: ContextMenuItem[], - ) { - oldGetExtraMenuOptions && oldGetExtraMenuOptions.apply(this, [canvas, menuOptions]); - addMenuItemOnExtraMenuOptions(this, config, menuOptions, after); - }; -} - -/** - * Waits for the canvas to be available on app using a single promise. - */ -let canvasResolver: Resolver | null = null; -export function waitForCanvas() { - if (canvasResolver === null) { - canvasResolver = getResolver(); - function _waitForCanvas() { - if (!canvasResolver!.completed) { - if (app?.canvas) { - canvasResolver!.resolve(app.canvas); - } else { - requestAnimationFrame(_waitForCanvas); - } - } - } - _waitForCanvas(); - } - return canvasResolver.promise; -} - -/** - * Waits for the graph to be available on app using a single promise. - */ -let graphResolver: Resolver | null = null; -export function waitForGraph() { - if (graphResolver === null) { - graphResolver = getResolver(); - function _wait() { - if (!graphResolver!.completed) { - if (app?.graph) { - graphResolver!.resolve(app.graph); - } else { - requestAnimationFrame(_wait); - } - } - } - _wait(); - } - return graphResolver.promise; -} - -export function addMenuItemOnExtraMenuOptions( - node: LGraphNode, - config: MenuConfig, - menuOptions: ContextMenuItem[], - after = "Shape", -) { - let idx = menuOptions - .slice() - .reverse() - .findIndex((option) => (option as any)?.isRgthree); - if (idx == -1) { - idx = menuOptions.findIndex((option) => option?.content.includes(after)) + 1; - if (!idx) { - idx = menuOptions.length - 1; - } - // Add a separator, and move to the next one. - menuOptions.splice(idx, 0, null); - idx++; - } else { - idx = menuOptions.length - idx; - } - - const subMenuOptions = - typeof config.subMenuOptions === "function" - ? config.subMenuOptions(node) - : config.subMenuOptions; - - menuOptions.splice(idx, 0, { - content: typeof config.name == "function" ? config.name(node) : config.name, - has_submenu: !!subMenuOptions?.length, - isRgthree: true, // Mark it, so we can find it. - callback: ( - value: ContextMenuItem, - _options: IContextMenuOptions, - event: MouseEvent, - parentMenu: ContextMenu | undefined, - _node: LGraphNode, - ) => { - if (!!subMenuOptions?.length) { - new LiteGraph.ContextMenu( - subMenuOptions.map((option) => (option ? { content: option } : null)), - { - event, - parentMenu, - callback: ( - subValue: ContextMenuItem, - _options: IContextMenuOptions, - _event: MouseEvent, - _parentMenu: ContextMenu | undefined, - _node: LGraphNode, - ) => { - if (config.property) { - node.properties = node.properties || {}; - node.properties[config.property] = config.prepareValue - ? config.prepareValue(subValue!.content, node) - : subValue!.content; - } - config.callback && config.callback(node, subValue?.content); - }, - }, - ); - return; - } - if (config.property) { - node.properties = node.properties || {}; - node.properties[config.property] = config.prepareValue - ? config.prepareValue(node.properties[config.property], node) - : !node.properties[config.property]; - } - config.callback && config.callback(node, value?.content); - }, - } as ContextMenuItem); -} - -export function addConnectionLayoutSupport( - node: Constructor, - app: ComfyApp, - options = [ - ["Left", "Right"], - ["Right", "Left"], - ], - callback?: (node: LGraphNode) => void, -) { - addMenuItem(node, app, { - name: "Connections Layout", - property: "connections_layout", - subMenuOptions: options.map((option) => option[0] + (option[1] ? " -> " + option[1] : "")), - prepareValue: (value, node) => { - const values = value.split(" -> "); - if (!values[1] && !node.outputs?.length) { - values[1] = LAYOUT_LABEL_OPPOSITES[values[0]!]!; - } - if (!LAYOUT_LABEL_TO_DATA[values[0]!] || !LAYOUT_LABEL_TO_DATA[values[1]!]) { - throw new Error(`New Layout invalid: [${values[0]}, ${values[1]}]`); - } - return values; - }, - callback: (node) => { - callback && callback(node); - app.graph.setDirtyCanvas(true, true); - }, - }); - - // const oldGetConnectionPos = node.prototype.getConnectionPos; - node.prototype.getConnectionPos = function (isInput: boolean, slotNumber: number, out: Vector2) { - // Purposefully do not need to call the old one. - // oldGetConnectionPos && oldGetConnectionPos.apply(this, [isInput, slotNumber, out]); - return getConnectionPosForLayout(this, isInput, slotNumber, out); - }; -} - -export function setConnectionsLayout(node: LGraphNode, newLayout: [string, string]) { - newLayout = newLayout || (node as any).defaultConnectionsLayout || ["Left", "Right"]; - // If we didn't supply an output layout, and there's no outputs, then just choose the opposite of the - // input as a safety. - if (!newLayout[1] && !node.outputs?.length) { - newLayout[1] = LAYOUT_LABEL_OPPOSITES[newLayout[0]!]!; - } - if (!LAYOUT_LABEL_TO_DATA[newLayout[0]] || !LAYOUT_LABEL_TO_DATA[newLayout[1]]) { - throw new Error(`New Layout invalid: [${newLayout[0]}, ${newLayout[1]}]`); - } - node.properties = node.properties || {}; - node.properties["connections_layout"] = newLayout; -} - -/** Allows collapsing of connections into one. Pretty unusable, unless you're the muter. */ -export function setConnectionsCollapse( - node: LGraphNode, - collapseConnections: boolean | null = null, -) { - node.properties = node.properties || {}; - collapseConnections = - collapseConnections !== null ? collapseConnections : !node.properties["collapse_connections"]; - node.properties["collapse_connections"] = collapseConnections; -} - -export function getConnectionPosForLayout( - node: LGraphNode, - isInput: boolean, - slotNumber: number, - out: Vector2, -) { - out = out || new Float32Array(2); - node.properties = node.properties || {}; - const layout = node.properties["connections_layout"] || - (node as any).defaultConnectionsLayout || ["Left", "Right"]; - const collapseConnections = node.properties["collapse_connections"] || false; - const offset = (node.constructor as any).layout_slot_offset ?? LiteGraph.NODE_SLOT_HEIGHT * 0.5; - let side = isInput ? layout[0] : layout[1]; - const otherSide = isInput ? layout[1] : layout[0]; - let data = LAYOUT_LABEL_TO_DATA[side]!; // || LAYOUT_LABEL_TO_DATA[isInput ? 'Left' : 'Right']; - const slotList = node[isInput ? "inputs" : "outputs"]; - const cxn = slotList[slotNumber]; - if (!cxn) { - console.log("No connection found.. weird", isInput, slotNumber); - return out; - } - // Experimental; doesn't work without node.clip_area set (so it won't draw outside), - // but litegraph.core inexplicably clips the title off which we want... so, no go. - // if (cxn.hidden) { - // out[0] = node.pos[0] - 100000 - // out[1] = node.pos[1] - 100000 - // return out - // } - if (cxn.disabled) { - // Let's store the original colors if have them and haven't yet overridden - if (cxn.color_on !== "#666665") { - (cxn as any)._color_on_org = (cxn as any)._color_on_org || cxn.color_on; - (cxn as any)._color_off_org = (cxn as any)._color_off_org || cxn.color_off; - } - cxn.color_on = "#666665"; - cxn.color_off = "#666665"; - } else if (cxn.color_on === "#666665") { - cxn.color_on = (cxn as any)._color_on_org || undefined; - cxn.color_off = (cxn as any)._color_off_org || undefined; - } - const displaySlot = collapseConnections - ? 0 - : slotNumber - - slotList.reduce((count, ioput, index) => { - count += index < slotNumber && ioput.hidden ? 1 : 0; - return count; - }, 0); - // Set the direction first. This is how the connection line will be drawn. - cxn.dir = data[0]; - - // If we are only 10px tall or wide, then look at connections_dir for the direction. - if ((node.size[0] == 10 || node.size[1] == 10) && node.properties["connections_dir"]) { - cxn.dir = node.properties["connections_dir"][isInput ? 0 : 1]!; - } - - if (side === "Left") { - if (node.flags.collapsed) { - var w = (node as any)._collapsed_width || LiteGraph.NODE_COLLAPSED_WIDTH; - out[0] = node.pos[0]; - out[1] = node.pos[1] - LiteGraph.NODE_TITLE_HEIGHT * 0.5; - } else { - // If we're an output, then the litegraph.core hates us; we need to blank out the name - // because it's not flexible enough to put the text on the inside. - toggleConnectionLabel(cxn, !isInput || collapseConnections || !!(node as any).hideSlotLabels); - out[0] = node.pos[0] + offset; - if ((node.constructor as any)?.type.includes("Reroute")) { - out[1] = node.pos[1] + node.size[1] * 0.5; - } else { - out[1] = - node.pos[1] + - (displaySlot + 0.7) * LiteGraph.NODE_SLOT_HEIGHT + - ((node.constructor as any).slot_start_y || 0); - } - } - } else if (side === "Right") { - if (node.flags.collapsed) { - var w = (node as any)._collapsed_width || LiteGraph.NODE_COLLAPSED_WIDTH; - out[0] = node.pos[0] + w; - out[1] = node.pos[1] - LiteGraph.NODE_TITLE_HEIGHT * 0.5; - } else { - // If we're an input, then the litegraph.core hates us; we need to blank out the name - // because it's not flexible enough to put the text on the inside. - toggleConnectionLabel(cxn, isInput || collapseConnections || !!(node as any).hideSlotLabels); - out[0] = node.pos[0] + node.size[0] + 1 - offset; - if ((node.constructor as any)?.type.includes("Reroute")) { - out[1] = node.pos[1] + node.size[1] * 0.5; - } else { - out[1] = - node.pos[1] + - (displaySlot + 0.7) * LiteGraph.NODE_SLOT_HEIGHT + - ((node.constructor as any).slot_start_y || 0); - } - } - - // Right now, only reroute uses top/bottom, so this may not work for other nodes - // (like, applying to nodes with titles, collapsed, multiple inputs/outputs, etc). - } else if (side === "Top") { - if (!(cxn as any).has_old_label) { - (cxn as any).has_old_label = true; - (cxn as any).old_label = cxn.label; - cxn.label = " "; - } - out[0] = node.pos[0] + node.size[0] * 0.5; - out[1] = node.pos[1] + offset; - } else if (side === "Bottom") { - if (!(cxn as any).has_old_label) { - (cxn as any).has_old_label = true; - (cxn as any).old_label = cxn.label; - cxn.label = " "; - } - out[0] = node.pos[0] + node.size[0] * 0.5; - out[1] = node.pos[1] + node.size[1] - offset; - } - return out; -} - -function toggleConnectionLabel(cxn: any, hide = true) { - if (hide) { - if (!(cxn as any).has_old_label) { - (cxn as any).has_old_label = true; - (cxn as any).old_label = cxn.label; - } - cxn.label = " "; - } else if (!hide && (cxn as any).has_old_label) { - (cxn as any).has_old_label = false; - cxn.label = (cxn as any).old_label; - (cxn as any).old_label = undefined; - } - return cxn; -} - -export function addHelpMenuItem(node: LGraphNode, content: string, menuOptions: ContextMenuItem[]) { - addMenuItemOnExtraMenuOptions( - node, - { - name: "๐Ÿ›Ÿ Node Help", - callback: (node) => { - if ((node as any).showHelp) { - (node as any).showHelp(); - } else { - new RgthreeHelpDialog(node, content).show(); - } - }, - }, - menuOptions, - "Properties Panel", - ); -} - -export enum PassThroughFollowing { - ALL, - NONE, - REROUTE_ONLY, -} - -/** - * Determines if, when doing a chain lookup for connected nodes, we want to pass through this node, - * like reroutes, etc. - */ -export function shouldPassThrough( - node?: LGraphNode | null, - passThroughFollowing = PassThroughFollowing.ALL, -) { - const type = (node?.constructor as typeof LGraphNode)?.type; - if (!type || passThroughFollowing === PassThroughFollowing.NONE) { - return false; - } - if (passThroughFollowing === PassThroughFollowing.REROUTE_ONLY) { - return type.includes("Reroute"); - } - return ( - type.includes("Reroute") || type.includes("Node Combiner") || type.includes("Node Collector") - ); -} - - -function filterOutPassthroughNodes( - infos: ConnectedNodeInfo[], - passThroughFollowing = PassThroughFollowing.ALL, -) { - return infos.filter((i) => !shouldPassThrough(i.node, passThroughFollowing)); -} - -/** - * Looks through the immediate chain of a node to collect all connected nodes, passing through nodes - * like reroute, etc. Will also disconnect duplicate nodes from a provided node - */ -export function getConnectedInputNodes( - startNode: LGraphNode, - currentNode?: LGraphNode, - slot?: number, - passThroughFollowing = PassThroughFollowing.ALL, -): LGraphNode[] { - return getConnectedNodesInfo( - startNode, - IoDirection.INPUT, - currentNode, - slot, - passThroughFollowing, - ).map((n) => n.node); -} -export function getConnectedInputInfosAndFilterPassThroughs( - startNode: LGraphNode, - currentNode?: LGraphNode, - slot?: number, - passThroughFollowing = PassThroughFollowing.ALL) { - return filterOutPassthroughNodes( - getConnectedNodesInfo(startNode, IoDirection.INPUT, currentNode, slot, passThroughFollowing), - passThroughFollowing); -} -export function getConnectedInputNodesAndFilterPassThroughs( - startNode: LGraphNode, - currentNode?: LGraphNode, - slot?: number, - passThroughFollowing = PassThroughFollowing.ALL, -): LGraphNode[] { - return getConnectedInputInfosAndFilterPassThroughs(startNode, currentNode, slot, passThroughFollowing).map(n => n.node); -} - -export function getConnectedOutputNodes( - startNode: LGraphNode, - currentNode?: LGraphNode, - slot?: number, - passThroughFollowing = PassThroughFollowing.ALL, -): LGraphNode[] { - return getConnectedNodesInfo( - startNode, - IoDirection.OUTPUT, - currentNode, - slot, - passThroughFollowing, - ).map((n) => n.node); -} - -export function getConnectedOutputNodesAndFilterPassThroughs( - startNode: LGraphNode, - currentNode?: LGraphNode, - slot?: number, - passThroughFollowing = PassThroughFollowing.ALL, -): LGraphNode[] { - return filterOutPassthroughNodes( - getConnectedNodesInfo(startNode, IoDirection.OUTPUT, currentNode, slot, passThroughFollowing), - passThroughFollowing, - ).map(n => n.node); -} - -export type ConnectedNodeInfo = { - node: LGraphNode; - travelFromSlot: number; - travelToSlot: number; - originTravelFromSlot: number; -}; - -export function getConnectedNodesInfo( - startNode: LGraphNode, - dir = IoDirection.INPUT, - currentNode?: LGraphNode, - slot?: number, - passThroughFollowing = PassThroughFollowing.ALL, - originTravelFromSlot?: number, -): ConnectedNodeInfo[] { - currentNode = currentNode || startNode; - let rootNodes: ConnectedNodeInfo[] = []; - if (startNode === currentNode || shouldPassThrough(currentNode, passThroughFollowing)) { - let linkIds: Array; - - slot = slot != null && slot > -1 ? slot : undefined; - if (dir == IoDirection.OUTPUT) { - if (slot != null) { - linkIds = [...(currentNode.outputs?.[slot]?.links || [])]; - } else { - linkIds = currentNode.outputs?.flatMap((i) => i.links) || []; - } - } else { - if (slot != null) { - linkIds = [currentNode.inputs?.[slot]?.link]; - } else { - linkIds = currentNode.inputs?.map((i) => i.link) || []; - } - } - let graph = app.graph as LGraph; - for (const linkId of linkIds) { - let link: LLink | null = null; - if (typeof linkId == "number") { - link = graph.links[linkId] as LLink; - } - if (!link) { - continue; - } - const travelFromSlot = dir == IoDirection.OUTPUT ? link.origin_slot : link.target_slot; - const connectedId = dir == IoDirection.OUTPUT ? link.target_id : link.origin_id; - const travelToSlot = dir == IoDirection.OUTPUT ? link.target_slot : link.origin_slot; - originTravelFromSlot = originTravelFromSlot != null ? originTravelFromSlot : travelFromSlot; - const originNode: LGraphNode = graph.getNodeById(connectedId)!; - if (!link) { - console.error("No connected node found... weird"); - continue; - } - if (rootNodes.some((n) => n.node == originNode)) { - console.log( - `${startNode.title} (${startNode.id}) seems to have two links to ${originNode.title} (${ - originNode.id - }). One may be stale: ${linkIds.join(", ")}`, - ); - } else { - // Add the node and, if it's a pass through, let's collect all its nodes as well. - rootNodes.push({ node: originNode, travelFromSlot, travelToSlot, originTravelFromSlot }); - if (shouldPassThrough(originNode, passThroughFollowing)) { - for (const foundNode of getConnectedNodesInfo( - startNode, - dir, - originNode, - undefined, - undefined, - originTravelFromSlot, - )) { - if (!rootNodes.map((n) => n.node).includes(foundNode.node)) { - rootNodes.push(foundNode); - } - } - } - } - } - } - return rootNodes; -} - -export type ConnectionType = { - type: string | string[]; - name: string | undefined; - label: string | undefined; -}; - -/** - * Follows a connection until we find a type associated with a slot. - * `skipSelf` skips the current slot, useful when we may have a dynamic slot that we want to start - * from, but find a type _after_ it (in case it needs to change). - */ -export function followConnectionUntilType( - node: LGraphNode, - dir: IoDirection, - slotNum?: number, - skipSelf = false, -): ConnectionType | null { - const slots = dir === IoDirection.OUTPUT ? node.outputs : node.inputs; - if (!slots || !slots.length) { - return null; - } - let type: ConnectionType | null = null; - if (slotNum) { - if (!slots[slotNum]) { - return null; - } - type = getTypeFromSlot(slots[slotNum], dir, skipSelf); - } else { - for (const slot of slots) { - type = getTypeFromSlot(slot, dir, skipSelf); - if (type) { - break; - } - } - } - return type; -} - -/** - * Gets the type from a slot. If the type is '*' then it will follow the node to find the next slot. - */ -function getTypeFromSlot( - slot: INodeInputSlot | INodeOutputSlot | undefined, - dir: IoDirection, - skipSelf = false, -): ConnectionType | null { - let graph = app.graph as LGraph; - let type = slot?.type; - if (!skipSelf && type != null && type != "*") { - return { type: type as string, label: slot?.label, name: slot?.name }; - } - const links = getSlotLinks(slot); - for (const link of links) { - const connectedId = dir == IoDirection.OUTPUT ? link.link.target_id : link.link.origin_id; - const connectedSlotNum = - dir == IoDirection.OUTPUT ? link.link.target_slot : link.link.origin_slot; - const connectedNode: LGraphNode = graph.getNodeById(connectedId)!; - // Reversed since if we're traveling down the output we want the connected node's input, etc. - const connectedSlots = - dir === IoDirection.OUTPUT ? connectedNode.inputs : connectedNode.outputs; - let connectedSlot = connectedSlots[connectedSlotNum]; - if (connectedSlot?.type != null && connectedSlot?.type != "*") { - return { - type: connectedSlot.type as string, - label: connectedSlot?.label, - name: connectedSlot?.name, - }; - } else if (connectedSlot?.type == "*") { - return followConnectionUntilType(connectedNode, dir); - } - } - return null; -} - -export async function replaceNode( - existingNode: LGraphNode, - typeOrNewNode: string | LGraphNode, - inputNameMap?: Map, -) { - const existingCtor = existingNode.constructor as typeof LGraphNode; - - const newNode = - typeof typeOrNewNode === "string" ? LiteGraph.createNode(typeOrNewNode) : typeOrNewNode; - // Port title (maybe) the position, size, and properties from the old node. - if (existingNode.title != existingCtor.title) { - newNode.title = existingNode.title; - } - newNode.pos = [...existingNode.pos]; - newNode.properties = { ...existingNode.properties }; - const oldComputeSize = [...existingNode.computeSize()]; - // oldSize to use. If we match the smallest size (computeSize) then don't record and we'll use - // the smalles side after conversion. - const oldSize = [ - existingNode.size[0] === oldComputeSize[0] ? null : existingNode.size[0], - existingNode.size[1] === oldComputeSize[1] ? null : existingNode.size[1], - ]; - - let setSizeIters = 0; - const setSizeFn = () => { - // Size gets messed up when ComfyUI adds the text widget, so reset after a delay. - // Since we could be adding many more slots, let's take the larger of the two. - const newComputesize = newNode.computeSize(); - newNode.size[0] = Math.max(oldSize[0] || 0, newComputesize[0]); - newNode.size[1] = Math.max(oldSize[1] || 0, newComputesize[1]); - setSizeIters++; - if (setSizeIters > 10) { - requestAnimationFrame(setSizeFn); - } - }; - setSizeFn(); - - // We now collect the links data, inputs and outputs, of the old node since these will be - // lost when we remove it. - const links: { - node: LGraphNode; - slot: number | string; - targetNode: LGraphNode; - targetSlot: number | string; - }[] = []; - for (const [index, output] of existingNode.outputs.entries()) { - for (const linkId of output.links || []) { - const link: LLink = (app.graph as LGraph).links[linkId]!; - if (!link) continue; - const targetNode = app.graph.getNodeById(link.target_id)!; - links.push({ node: newNode, slot: output.name, targetNode, targetSlot: link.target_slot }); - } - } - for (const [index, input] of existingNode.inputs.entries()) { - const linkId = input.link; - if (linkId) { - const link: LLink = (app.graph as LGraph).links[linkId]!; - const originNode = app.graph.getNodeById(link.origin_id)!; - links.push({ - node: originNode, - slot: link.origin_slot, - targetNode: newNode, - targetSlot: inputNameMap?.has(input.name) - ? inputNameMap.get(input.name)! - : input.name || index, - }); - } - } - // Add the new node, remove the old node. - app.graph.add(newNode); - await wait(); - // Now go through and connect the other nodes up as they were. - for (const link of links) { - link.node.connect(link.slot, link.targetNode, link.targetSlot); - } - await wait(); - app.graph.remove(existingNode); - newNode.size = newNode.computeSize(); - newNode.setDirtyCanvas(true, true); - return newNode; -} - -export function getOriginNodeByLink(linkId?: number | null) { - let node: LGraphNode | null = null; - if (linkId != null) { - const link: LLink = app.graph.links[linkId]!; - node = (link != null && app.graph.getNodeById(link.origin_id)) || null; - } - return node; -} - -export function applyMixins(original: Constructor, constructors: any[]) { - constructors.forEach((baseCtor) => { - Object.getOwnPropertyNames(baseCtor.prototype).forEach((name) => { - Object.defineProperty( - original.prototype, - name, - Object.getOwnPropertyDescriptor(baseCtor.prototype, name) || Object.create(null), - ); - }); - }); -} - -/** - * Retruns a list of `{id: number, link: LLlink}` for a given input or output. - * - * Obviously, for an input, this will be a max of one. - */ -export function getSlotLinks(inputOrOutput?: INodeInputSlot | INodeOutputSlot | null) { - const links: { id: number; link: LLink }[] = []; - if (!inputOrOutput) { - return links; - } - if ((inputOrOutput as INodeOutputSlot).links?.length) { - const output = inputOrOutput as INodeOutputSlot; - for (const linkId of output.links || []) { - const link: LLink = (app.graph as LGraph).links[linkId]!; - if (link) { - links.push({ id: linkId, link: link }); - } - } - } - if ((inputOrOutput as INodeInputSlot).link) { - const input = inputOrOutput as INodeInputSlot; - const link: LLink = (app.graph as LGraph).links[input.link!]!; - if (link) { - links.push({ id: input.link!, link: link }); - } - } - return links; -} - -/** - * Given a node, whether we're dealing with INPUTS or OUTPUTS, and the server data, re-arrange then - * slots to match the order. - */ -export async function matchLocalSlotsToServer( - node: LGraphNode, - direction: IoDirection, - serverNodeData: ComfyObjectInfo, -) { - const serverSlotNames = - direction == IoDirection.INPUT - ? Object.keys(serverNodeData.input?.optional || {}) - : serverNodeData.output_name; - const serverSlotTypes = - direction == IoDirection.INPUT - ? (Object.values(serverNodeData.input?.optional || {}).map((i) => i[0]) as string[]) - : serverNodeData.output; - const slots = direction == IoDirection.INPUT ? node.inputs : node.outputs; - - // Let's go through the node data names and make sure our current ones match, and update if not. - let firstIndex = slots.findIndex((o, i) => i !== serverSlotNames.indexOf(o.name)); - if (firstIndex > -1) { - // Have mismatches. First, let's go through and save all our links by name. - const links: { [key: string]: { id: number; link: LLink }[] } = {}; - slots.map((slot) => { - // There's a chance we have duplicate names on an upgrade, so we'll collect all links to one - // name so we don't ovewrite our list per name. - links[slot.name] = links[slot.name] || []; - links[slot.name]?.push(...getSlotLinks(slot)); - }); - - // Now, go through and rearrange outputs by splicing - for (const [index, serverSlotName] of serverSlotNames.entries()) { - const currentNodeSlot = slots.map((s) => s.name).indexOf(serverSlotName); - if (currentNodeSlot > -1) { - if (currentNodeSlot != index) { - const splicedItem = slots.splice(currentNodeSlot, 1)[0]!; - slots.splice(index, 0, splicedItem as any); - } - } else if (currentNodeSlot === -1) { - const splicedItem = { - name: serverSlotName, - type: serverSlotTypes![index], - links: [], - }; - slots.splice(index, 0, splicedItem as any); - } - } - - if (slots.length > serverSlotNames.length) { - for (let i = slots.length - 1; i > serverSlotNames.length - 1; i--) { - if (direction == IoDirection.INPUT) { - node.disconnectInput(i); - node.removeInput(i); - } else { - node.disconnectOutput(i); - node.removeOutput(i); - } - } - } - - // Now, go through the link data again and make sure the origin_slot is the correct slot. - for (const [name, slotLinks] of Object.entries(links)) { - let currentNodeSlot = slots.map((s) => s.name).indexOf(name); - if (currentNodeSlot > -1) { - for (const linkData of slotLinks) { - if (direction == IoDirection.INPUT) { - linkData.link.target_slot = currentNodeSlot; - } else { - linkData.link.origin_slot = currentNodeSlot; - // If our next node is a Reroute, then let's get it to update the type. - const nextNode = app.graph.getNodeById(linkData.link.target_id); - // (Check nextNode, as sometimes graphs seem to have very stale data and that node id - // doesn't exist). - if ( - nextNode && - (nextNode.constructor as ComfyNodeConstructor)?.type!.includes("Reroute") - ) { - (nextNode as any).stabilize && (nextNode as any).stabilize(); - } - } - } - } - } - } -} - -export function isValidConnection(ioA?: INodeSlot | null, ioB?: INodeSlot | null) { - if (!ioA || !ioB) { - return false; - } - const typeA = String(ioA.type); - const typeB = String(ioB.type); - // What does litegraph think, which includes looking at array values. - let isValid = LiteGraph.isValidConnection(typeA, typeB); - - // This is here to fix the churn happening in list types in comfyui itself.. - // https://github.com/comfyanonymous/ComfyUI/issues/1674 - if (!isValid) { - let areCombos = - (typeA.includes(",") && typeB === "COMBO") || (typeA === "COMBO" && typeB.includes(",")); - // We don't want to let any old combo connect to any old combo, so we'll look at the names too. - if (areCombos) { - // Some nodes use "_name" and some use "model" and "ckpt", so normalize - const nameA = ioA.name.toUpperCase().replace("_NAME", "").replace("CKPT", "MODEL"); - const nameB = ioB.name.toUpperCase().replace("_NAME", "").replace("CKPT", "MODEL"); - isValid = nameA.includes(nameB) || nameB.includes(nameA); - } - } - return isValid; -} - -/** - * Patches the LiteGraph.isValidConnection so old nodes can connect to this new COMBO type for all - * lists (without users needing to go through and re-create all their nodes one by one). - */ -const oldIsValidConnection = LiteGraph.isValidConnection; -LiteGraph.isValidConnection = function (typeA: string | string[], typeB: string | string[]) { - let isValid = oldIsValidConnection.call(LiteGraph, typeA, typeB); - if (!isValid) { - typeA = String(typeA); - typeB = String(typeB); - // This is waaaay too liberal and now any combos can connect to any combos. But we only have the - // types (not names like my util above), and connecting too liberally is better than old nodes - // with lists not being able to connect to this new COMBO type. And, anyway, it matches the - // current behavior today with new nodes anyway, where all lists are COMBO types. - // Refs: https://github.com/comfyanonymous/ComfyUI/issues/1674 - // https://github.com/comfyanonymous/ComfyUI/pull/1675 - let areCombos = - (typeA.includes(",") && typeB === "COMBO") || (typeA === "COMBO" && typeB.includes(",")); - isValid = areCombos; - } - return isValid; -}; +import type { ComfyApp, ComfyNodeConstructor, ComfyObjectInfo } from "typings/comfy.js"; +import type { + Vector2, + LGraphCanvas, + ContextMenuItem, + LLink, + LGraph, + IContextMenuOptions, + ContextMenu, + LGraphNode, + INodeSlot, + INodeInputSlot, + INodeOutputSlot, +} from "typings/litegraph.js"; +import type { Constructor } from "typings/index.js"; +import { app } from "scripts/app.js"; +import { api } from "scripts/api.js"; +import { Resolver, getResolver, wait } from "rgthree/common/shared_utils.js"; +import { RgthreeHelpDialog } from "rgthree/common/dialog.js"; + +/** + * Override the api.getNodeDefs call to add a hook for refreshing node defs. + * This is necessary for power prompt's custom combos. Since API implements + * add/removeEventListener already, this is rather trivial. + */ +const oldApiGetNodeDefs = api.getNodeDefs; +api.getNodeDefs = async function () { + const defs = await oldApiGetNodeDefs.call(api); + this.dispatchEvent(new CustomEvent("fresh-node-defs", { detail: defs })); + return defs; +}; + +export enum IoDirection { + INPUT, + OUTPUT, +} + +const PADDING = 0; + +type LiteGraphDir = + | typeof LiteGraph.LEFT + | typeof LiteGraph.RIGHT + | typeof LiteGraph.UP + | typeof LiteGraph.DOWN; +export const LAYOUT_LABEL_TO_DATA: { [label: string]: [LiteGraphDir, Vector2, Vector2] } = { + Left: [LiteGraph.LEFT, [0, 0.5], [PADDING, 0]], + Right: [LiteGraph.RIGHT, [1, 0.5], [-PADDING, 0]], + Top: [LiteGraph.UP, [0.5, 0], [0, PADDING]], + Bottom: [LiteGraph.DOWN, [0.5, 1], [0, -PADDING]], +}; +export const LAYOUT_LABEL_OPPOSITES: { [label: string]: string } = { + Left: "Right", + Right: "Left", + Top: "Bottom", + Bottom: "Top", +}; +export const LAYOUT_CLOCKWISE = ["Top", "Right", "Bottom", "Left"]; + +interface MenuConfig { + name: string | ((node: LGraphNode) => string); + property?: string; + prepareValue?: (value: string, node: LGraphNode) => any; + callback?: (node: LGraphNode, value?: string) => void; + subMenuOptions?: (string | null)[] | ((node: LGraphNode) => (string | null)[]); +} + +export function addMenuItem( + node: Constructor, + _app: ComfyApp, + config: MenuConfig, + after = "Shape", +) { + const oldGetExtraMenuOptions = node.prototype.getExtraMenuOptions; + node.prototype.getExtraMenuOptions = function ( + canvas: LGraphCanvas, + menuOptions: ContextMenuItem[], + ) { + oldGetExtraMenuOptions && oldGetExtraMenuOptions.apply(this, [canvas, menuOptions]); + addMenuItemOnExtraMenuOptions(this, config, menuOptions, after); + }; +} + +/** + * Waits for the canvas to be available on app using a single promise. + */ +let canvasResolver: Resolver | null = null; +export function waitForCanvas() { + if (canvasResolver === null) { + canvasResolver = getResolver(); + function _waitForCanvas() { + if (!canvasResolver!.completed) { + if (app?.canvas) { + canvasResolver!.resolve(app.canvas); + } else { + requestAnimationFrame(_waitForCanvas); + } + } + } + _waitForCanvas(); + } + return canvasResolver.promise; +} + +/** + * Waits for the graph to be available on app using a single promise. + */ +let graphResolver: Resolver | null = null; +export function waitForGraph() { + if (graphResolver === null) { + graphResolver = getResolver(); + function _wait() { + if (!graphResolver!.completed) { + if (app?.graph) { + graphResolver!.resolve(app.graph); + } else { + requestAnimationFrame(_wait); + } + } + } + _wait(); + } + return graphResolver.promise; +} + +export function addMenuItemOnExtraMenuOptions( + node: LGraphNode, + config: MenuConfig, + menuOptions: ContextMenuItem[], + after = "Shape", +) { + let idx = menuOptions + .slice() + .reverse() + .findIndex((option) => (option as any)?.isRgthree); + if (idx == -1) { + idx = menuOptions.findIndex((option) => option?.content.includes(after)) + 1; + if (!idx) { + idx = menuOptions.length - 1; + } + // Add a separator, and move to the next one. + menuOptions.splice(idx, 0, null); + idx++; + } else { + idx = menuOptions.length - idx; + } + + const subMenuOptions = + typeof config.subMenuOptions === "function" + ? config.subMenuOptions(node) + : config.subMenuOptions; + + menuOptions.splice(idx, 0, { + content: typeof config.name == "function" ? config.name(node) : config.name, + has_submenu: !!subMenuOptions?.length, + isRgthree: true, // Mark it, so we can find it. + callback: ( + value: ContextMenuItem, + _options: IContextMenuOptions, + event: MouseEvent, + parentMenu: ContextMenu | undefined, + _node: LGraphNode, + ) => { + if (!!subMenuOptions?.length) { + new LiteGraph.ContextMenu( + subMenuOptions.map((option) => (option ? { content: option } : null)), + { + event, + parentMenu, + callback: ( + subValue: ContextMenuItem, + _options: IContextMenuOptions, + _event: MouseEvent, + _parentMenu: ContextMenu | undefined, + _node: LGraphNode, + ) => { + if (config.property) { + node.properties = node.properties || {}; + node.properties[config.property] = config.prepareValue + ? config.prepareValue(subValue!.content, node) + : subValue!.content; + } + config.callback && config.callback(node, subValue?.content); + }, + }, + ); + return; + } + if (config.property) { + node.properties = node.properties || {}; + node.properties[config.property] = config.prepareValue + ? config.prepareValue(node.properties[config.property], node) + : !node.properties[config.property]; + } + config.callback && config.callback(node, value?.content); + }, + } as ContextMenuItem); +} + +export function addConnectionLayoutSupport( + node: Constructor, + app: ComfyApp, + options = [ + ["Left", "Right"], + ["Right", "Left"], + ], + callback?: (node: LGraphNode) => void, +) { + addMenuItem(node, app, { + name: "Connections Layout", + property: "connections_layout", + subMenuOptions: options.map((option) => option[0] + (option[1] ? " -> " + option[1] : "")), + prepareValue: (value, node) => { + const values = value.split(" -> "); + if (!values[1] && !node.outputs?.length) { + values[1] = LAYOUT_LABEL_OPPOSITES[values[0]!]!; + } + if (!LAYOUT_LABEL_TO_DATA[values[0]!] || !LAYOUT_LABEL_TO_DATA[values[1]!]) { + throw new Error(`New Layout invalid: [${values[0]}, ${values[1]}]`); + } + return values; + }, + callback: (node) => { + callback && callback(node); + app.graph.setDirtyCanvas(true, true); + }, + }); + + // const oldGetConnectionPos = node.prototype.getConnectionPos; + node.prototype.getConnectionPos = function (isInput: boolean, slotNumber: number, out: Vector2) { + // Purposefully do not need to call the old one. + // oldGetConnectionPos && oldGetConnectionPos.apply(this, [isInput, slotNumber, out]); + return getConnectionPosForLayout(this, isInput, slotNumber, out); + }; +} + +export function setConnectionsLayout(node: LGraphNode, newLayout: [string, string]) { + newLayout = newLayout || (node as any).defaultConnectionsLayout || ["Left", "Right"]; + // If we didn't supply an output layout, and there's no outputs, then just choose the opposite of the + // input as a safety. + if (!newLayout[1] && !node.outputs?.length) { + newLayout[1] = LAYOUT_LABEL_OPPOSITES[newLayout[0]!]!; + } + if (!LAYOUT_LABEL_TO_DATA[newLayout[0]] || !LAYOUT_LABEL_TO_DATA[newLayout[1]]) { + throw new Error(`New Layout invalid: [${newLayout[0]}, ${newLayout[1]}]`); + } + node.properties = node.properties || {}; + node.properties["connections_layout"] = newLayout; +} + +/** Allows collapsing of connections into one. Pretty unusable, unless you're the muter. */ +export function setConnectionsCollapse( + node: LGraphNode, + collapseConnections: boolean | null = null, +) { + node.properties = node.properties || {}; + collapseConnections = + collapseConnections !== null ? collapseConnections : !node.properties["collapse_connections"]; + node.properties["collapse_connections"] = collapseConnections; +} + +export function getConnectionPosForLayout( + node: LGraphNode, + isInput: boolean, + slotNumber: number, + out: Vector2, +) { + out = out || new Float32Array(2); + node.properties = node.properties || {}; + const layout = node.properties["connections_layout"] || + (node as any).defaultConnectionsLayout || ["Left", "Right"]; + const collapseConnections = node.properties["collapse_connections"] || false; + const offset = (node.constructor as any).layout_slot_offset ?? LiteGraph.NODE_SLOT_HEIGHT * 0.5; + let side = isInput ? layout[0] : layout[1]; + const otherSide = isInput ? layout[1] : layout[0]; + let data = LAYOUT_LABEL_TO_DATA[side]!; // || LAYOUT_LABEL_TO_DATA[isInput ? 'Left' : 'Right']; + const slotList = node[isInput ? "inputs" : "outputs"]; + const cxn = slotList[slotNumber]; + if (!cxn) { + console.log("No connection found.. weird", isInput, slotNumber); + return out; + } + // Experimental; doesn't work without node.clip_area set (so it won't draw outside), + // but litegraph.core inexplicably clips the title off which we want... so, no go. + // if (cxn.hidden) { + // out[0] = node.pos[0] - 100000 + // out[1] = node.pos[1] - 100000 + // return out + // } + if (cxn.disabled) { + // Let's store the original colors if have them and haven't yet overridden + if (cxn.color_on !== "#666665") { + (cxn as any)._color_on_org = (cxn as any)._color_on_org || cxn.color_on; + (cxn as any)._color_off_org = (cxn as any)._color_off_org || cxn.color_off; + } + cxn.color_on = "#666665"; + cxn.color_off = "#666665"; + } else if (cxn.color_on === "#666665") { + cxn.color_on = (cxn as any)._color_on_org || undefined; + cxn.color_off = (cxn as any)._color_off_org || undefined; + } + const displaySlot = collapseConnections + ? 0 + : slotNumber - + slotList.reduce((count, ioput, index) => { + count += index < slotNumber && ioput.hidden ? 1 : 0; + return count; + }, 0); + // Set the direction first. This is how the connection line will be drawn. + cxn.dir = data[0]; + + // If we are only 10px tall or wide, then look at connections_dir for the direction. + if ((node.size[0] == 10 || node.size[1] == 10) && node.properties["connections_dir"]) { + cxn.dir = node.properties["connections_dir"][isInput ? 0 : 1]!; + } + + if (side === "Left") { + if (node.flags.collapsed) { + var w = (node as any)._collapsed_width || LiteGraph.NODE_COLLAPSED_WIDTH; + out[0] = node.pos[0]; + out[1] = node.pos[1] - LiteGraph.NODE_TITLE_HEIGHT * 0.5; + } else { + // If we're an output, then the litegraph.core hates us; we need to blank out the name + // because it's not flexible enough to put the text on the inside. + toggleConnectionLabel(cxn, !isInput || collapseConnections || !!(node as any).hideSlotLabels); + out[0] = node.pos[0] + offset; + if ((node.constructor as any)?.type.includes("Reroute")) { + out[1] = node.pos[1] + node.size[1] * 0.5; + } else { + out[1] = + node.pos[1] + + (displaySlot + 0.7) * LiteGraph.NODE_SLOT_HEIGHT + + ((node.constructor as any).slot_start_y || 0); + } + } + } else if (side === "Right") { + if (node.flags.collapsed) { + var w = (node as any)._collapsed_width || LiteGraph.NODE_COLLAPSED_WIDTH; + out[0] = node.pos[0] + w; + out[1] = node.pos[1] - LiteGraph.NODE_TITLE_HEIGHT * 0.5; + } else { + // If we're an input, then the litegraph.core hates us; we need to blank out the name + // because it's not flexible enough to put the text on the inside. + toggleConnectionLabel(cxn, isInput || collapseConnections || !!(node as any).hideSlotLabels); + out[0] = node.pos[0] + node.size[0] + 1 - offset; + if ((node.constructor as any)?.type.includes("Reroute")) { + out[1] = node.pos[1] + node.size[1] * 0.5; + } else { + out[1] = + node.pos[1] + + (displaySlot + 0.7) * LiteGraph.NODE_SLOT_HEIGHT + + ((node.constructor as any).slot_start_y || 0); + } + } + + // Right now, only reroute uses top/bottom, so this may not work for other nodes + // (like, applying to nodes with titles, collapsed, multiple inputs/outputs, etc). + } else if (side === "Top") { + if (!(cxn as any).has_old_label) { + (cxn as any).has_old_label = true; + (cxn as any).old_label = cxn.label; + cxn.label = " "; + } + out[0] = node.pos[0] + node.size[0] * 0.5; + out[1] = node.pos[1] + offset; + } else if (side === "Bottom") { + if (!(cxn as any).has_old_label) { + (cxn as any).has_old_label = true; + (cxn as any).old_label = cxn.label; + cxn.label = " "; + } + out[0] = node.pos[0] + node.size[0] * 0.5; + out[1] = node.pos[1] + node.size[1] - offset; + } + return out; +} + +function toggleConnectionLabel(cxn: any, hide = true) { + if (hide) { + if (!(cxn as any).has_old_label) { + (cxn as any).has_old_label = true; + (cxn as any).old_label = cxn.label; + } + cxn.label = " "; + } else if (!hide && (cxn as any).has_old_label) { + (cxn as any).has_old_label = false; + cxn.label = (cxn as any).old_label; + (cxn as any).old_label = undefined; + } + return cxn; +} + +export function addHelpMenuItem(node: LGraphNode, content: string, menuOptions: ContextMenuItem[]) { + addMenuItemOnExtraMenuOptions( + node, + { + name: "๐Ÿ›Ÿ Node Help", + callback: (node) => { + if ((node as any).showHelp) { + (node as any).showHelp(); + } else { + new RgthreeHelpDialog(node, content).show(); + } + }, + }, + menuOptions, + "Properties Panel", + ); +} + +export enum PassThroughFollowing { + ALL, + NONE, + REROUTE_ONLY, +} + +/** + * Determines if, when doing a chain lookup for connected nodes, we want to pass through this node, + * like reroutes, etc. + */ +export function shouldPassThrough( + node?: LGraphNode | null, + passThroughFollowing = PassThroughFollowing.ALL, +) { + const type = (node?.constructor as typeof LGraphNode)?.type; + if (!type || passThroughFollowing === PassThroughFollowing.NONE) { + return false; + } + if (passThroughFollowing === PassThroughFollowing.REROUTE_ONLY) { + return type.includes("Reroute"); + } + return ( + type.includes("Reroute") || type.includes("Node Combiner") || type.includes("Node Collector") + ); +} + + +function filterOutPassthroughNodes( + infos: ConnectedNodeInfo[], + passThroughFollowing = PassThroughFollowing.ALL, +) { + return infos.filter((i) => !shouldPassThrough(i.node, passThroughFollowing)); +} + +/** + * Looks through the immediate chain of a node to collect all connected nodes, passing through nodes + * like reroute, etc. Will also disconnect duplicate nodes from a provided node + */ +export function getConnectedInputNodes( + startNode: LGraphNode, + currentNode?: LGraphNode, + slot?: number, + passThroughFollowing = PassThroughFollowing.ALL, +): LGraphNode[] { + return getConnectedNodesInfo( + startNode, + IoDirection.INPUT, + currentNode, + slot, + passThroughFollowing, + ).map((n) => n.node); +} +export function getConnectedInputInfosAndFilterPassThroughs( + startNode: LGraphNode, + currentNode?: LGraphNode, + slot?: number, + passThroughFollowing = PassThroughFollowing.ALL) { + return filterOutPassthroughNodes( + getConnectedNodesInfo(startNode, IoDirection.INPUT, currentNode, slot, passThroughFollowing), + passThroughFollowing); +} +export function getConnectedInputNodesAndFilterPassThroughs( + startNode: LGraphNode, + currentNode?: LGraphNode, + slot?: number, + passThroughFollowing = PassThroughFollowing.ALL, +): LGraphNode[] { + return getConnectedInputInfosAndFilterPassThroughs(startNode, currentNode, slot, passThroughFollowing).map(n => n.node); +} + +export function getConnectedOutputNodes( + startNode: LGraphNode, + currentNode?: LGraphNode, + slot?: number, + passThroughFollowing = PassThroughFollowing.ALL, +): LGraphNode[] { + return getConnectedNodesInfo( + startNode, + IoDirection.OUTPUT, + currentNode, + slot, + passThroughFollowing, + ).map((n) => n.node); +} + +export function getConnectedOutputNodesAndFilterPassThroughs( + startNode: LGraphNode, + currentNode?: LGraphNode, + slot?: number, + passThroughFollowing = PassThroughFollowing.ALL, +): LGraphNode[] { + return filterOutPassthroughNodes( + getConnectedNodesInfo(startNode, IoDirection.OUTPUT, currentNode, slot, passThroughFollowing), + passThroughFollowing, + ).map(n => n.node); +} + +export type ConnectedNodeInfo = { + node: LGraphNode; + travelFromSlot: number; + travelToSlot: number; + originTravelFromSlot: number; +}; + +export function getConnectedNodesInfo( + startNode: LGraphNode, + dir = IoDirection.INPUT, + currentNode?: LGraphNode, + slot?: number, + passThroughFollowing = PassThroughFollowing.ALL, + originTravelFromSlot?: number, +): ConnectedNodeInfo[] { + currentNode = currentNode || startNode; + let rootNodes: ConnectedNodeInfo[] = []; + if (startNode === currentNode || shouldPassThrough(currentNode, passThroughFollowing)) { + let linkIds: Array; + + slot = slot != null && slot > -1 ? slot : undefined; + if (dir == IoDirection.OUTPUT) { + if (slot != null) { + linkIds = [...(currentNode.outputs?.[slot]?.links || [])]; + } else { + linkIds = currentNode.outputs?.flatMap((i) => i.links) || []; + } + } else { + if (slot != null) { + linkIds = [currentNode.inputs?.[slot]?.link]; + } else { + linkIds = currentNode.inputs?.map((i) => i.link) || []; + } + } + let graph = app.graph as LGraph; + for (const linkId of linkIds) { + let link: LLink | null = null; + if (typeof linkId == "number") { + link = graph.links[linkId] as LLink; + } + if (!link) { + continue; + } + const travelFromSlot = dir == IoDirection.OUTPUT ? link.origin_slot : link.target_slot; + const connectedId = dir == IoDirection.OUTPUT ? link.target_id : link.origin_id; + const travelToSlot = dir == IoDirection.OUTPUT ? link.target_slot : link.origin_slot; + originTravelFromSlot = originTravelFromSlot != null ? originTravelFromSlot : travelFromSlot; + const originNode: LGraphNode = graph.getNodeById(connectedId)!; + if (!link) { + console.error("No connected node found... weird"); + continue; + } + if (rootNodes.some((n) => n.node == originNode)) { + console.log( + `${startNode.title} (${startNode.id}) seems to have two links to ${originNode.title} (${ + originNode.id + }). One may be stale: ${linkIds.join(", ")}`, + ); + } else { + // Add the node and, if it's a pass through, let's collect all its nodes as well. + rootNodes.push({ node: originNode, travelFromSlot, travelToSlot, originTravelFromSlot }); + if (shouldPassThrough(originNode, passThroughFollowing)) { + for (const foundNode of getConnectedNodesInfo( + startNode, + dir, + originNode, + undefined, + undefined, + originTravelFromSlot, + )) { + if (!rootNodes.map((n) => n.node).includes(foundNode.node)) { + rootNodes.push(foundNode); + } + } + } + } + } + } + return rootNodes; +} + +export type ConnectionType = { + type: string | string[]; + name: string | undefined; + label: string | undefined; +}; + +/** + * Follows a connection until we find a type associated with a slot. + * `skipSelf` skips the current slot, useful when we may have a dynamic slot that we want to start + * from, but find a type _after_ it (in case it needs to change). + */ +export function followConnectionUntilType( + node: LGraphNode, + dir: IoDirection, + slotNum?: number, + skipSelf = false, +): ConnectionType | null { + const slots = dir === IoDirection.OUTPUT ? node.outputs : node.inputs; + if (!slots || !slots.length) { + return null; + } + let type: ConnectionType | null = null; + if (slotNum) { + if (!slots[slotNum]) { + return null; + } + type = getTypeFromSlot(slots[slotNum], dir, skipSelf); + } else { + for (const slot of slots) { + type = getTypeFromSlot(slot, dir, skipSelf); + if (type) { + break; + } + } + } + return type; +} + +/** + * Gets the type from a slot. If the type is '*' then it will follow the node to find the next slot. + */ +function getTypeFromSlot( + slot: INodeInputSlot | INodeOutputSlot | undefined, + dir: IoDirection, + skipSelf = false, +): ConnectionType | null { + let graph = app.graph as LGraph; + let type = slot?.type; + if (!skipSelf && type != null && type != "*") { + return { type: type as string, label: slot?.label, name: slot?.name }; + } + const links = getSlotLinks(slot); + for (const link of links) { + const connectedId = dir == IoDirection.OUTPUT ? link.link.target_id : link.link.origin_id; + const connectedSlotNum = + dir == IoDirection.OUTPUT ? link.link.target_slot : link.link.origin_slot; + const connectedNode: LGraphNode = graph.getNodeById(connectedId)!; + // Reversed since if we're traveling down the output we want the connected node's input, etc. + const connectedSlots = + dir === IoDirection.OUTPUT ? connectedNode.inputs : connectedNode.outputs; + let connectedSlot = connectedSlots[connectedSlotNum]; + if (connectedSlot?.type != null && connectedSlot?.type != "*") { + return { + type: connectedSlot.type as string, + label: connectedSlot?.label, + name: connectedSlot?.name, + }; + } else if (connectedSlot?.type == "*") { + return followConnectionUntilType(connectedNode, dir); + } + } + return null; +} + +export async function replaceNode( + existingNode: LGraphNode, + typeOrNewNode: string | LGraphNode, + inputNameMap?: Map, +) { + const existingCtor = existingNode.constructor as typeof LGraphNode; + + const newNode = + typeof typeOrNewNode === "string" ? LiteGraph.createNode(typeOrNewNode) : typeOrNewNode; + // Port title (maybe) the position, size, and properties from the old node. + if (existingNode.title != existingCtor.title) { + newNode.title = existingNode.title; + } + newNode.pos = [...existingNode.pos]; + newNode.properties = { ...existingNode.properties }; + const oldComputeSize = [...existingNode.computeSize()]; + // oldSize to use. If we match the smallest size (computeSize) then don't record and we'll use + // the smalles side after conversion. + const oldSize = [ + existingNode.size[0] === oldComputeSize[0] ? null : existingNode.size[0], + existingNode.size[1] === oldComputeSize[1] ? null : existingNode.size[1], + ]; + + let setSizeIters = 0; + const setSizeFn = () => { + // Size gets messed up when ComfyUI adds the text widget, so reset after a delay. + // Since we could be adding many more slots, let's take the larger of the two. + const newComputesize = newNode.computeSize(); + newNode.size[0] = Math.max(oldSize[0] || 0, newComputesize[0]); + newNode.size[1] = Math.max(oldSize[1] || 0, newComputesize[1]); + setSizeIters++; + if (setSizeIters > 10) { + requestAnimationFrame(setSizeFn); + } + }; + setSizeFn(); + + // We now collect the links data, inputs and outputs, of the old node since these will be + // lost when we remove it. + const links: { + node: LGraphNode; + slot: number | string; + targetNode: LGraphNode; + targetSlot: number | string; + }[] = []; + for (const [index, output] of existingNode.outputs.entries()) { + for (const linkId of output.links || []) { + const link: LLink = (app.graph as LGraph).links[linkId]!; + if (!link) continue; + const targetNode = app.graph.getNodeById(link.target_id)!; + links.push({ node: newNode, slot: output.name, targetNode, targetSlot: link.target_slot }); + } + } + for (const [index, input] of existingNode.inputs.entries()) { + const linkId = input.link; + if (linkId) { + const link: LLink = (app.graph as LGraph).links[linkId]!; + const originNode = app.graph.getNodeById(link.origin_id)!; + links.push({ + node: originNode, + slot: link.origin_slot, + targetNode: newNode, + targetSlot: inputNameMap?.has(input.name) + ? inputNameMap.get(input.name)! + : input.name || index, + }); + } + } + // Add the new node, remove the old node. + app.graph.add(newNode); + await wait(); + // Now go through and connect the other nodes up as they were. + for (const link of links) { + link.node.connect(link.slot, link.targetNode, link.targetSlot); + } + await wait(); + app.graph.remove(existingNode); + newNode.size = newNode.computeSize(); + newNode.setDirtyCanvas(true, true); + return newNode; +} + +export function getOriginNodeByLink(linkId?: number | null) { + let node: LGraphNode | null = null; + if (linkId != null) { + const link: LLink = app.graph.links[linkId]!; + node = (link != null && app.graph.getNodeById(link.origin_id)) || null; + } + return node; +} + +export function applyMixins(original: Constructor, constructors: any[]) { + constructors.forEach((baseCtor) => { + Object.getOwnPropertyNames(baseCtor.prototype).forEach((name) => { + Object.defineProperty( + original.prototype, + name, + Object.getOwnPropertyDescriptor(baseCtor.prototype, name) || Object.create(null), + ); + }); + }); +} + +/** + * Retruns a list of `{id: number, link: LLlink}` for a given input or output. + * + * Obviously, for an input, this will be a max of one. + */ +export function getSlotLinks(inputOrOutput?: INodeInputSlot | INodeOutputSlot | null) { + const links: { id: number; link: LLink }[] = []; + if (!inputOrOutput) { + return links; + } + if ((inputOrOutput as INodeOutputSlot).links?.length) { + const output = inputOrOutput as INodeOutputSlot; + for (const linkId of output.links || []) { + const link: LLink = (app.graph as LGraph).links[linkId]!; + if (link) { + links.push({ id: linkId, link: link }); + } + } + } + if ((inputOrOutput as INodeInputSlot).link) { + const input = inputOrOutput as INodeInputSlot; + const link: LLink = (app.graph as LGraph).links[input.link!]!; + if (link) { + links.push({ id: input.link!, link: link }); + } + } + return links; +} + +/** + * Given a node, whether we're dealing with INPUTS or OUTPUTS, and the server data, re-arrange then + * slots to match the order. + */ +export async function matchLocalSlotsToServer( + node: LGraphNode, + direction: IoDirection, + serverNodeData: ComfyObjectInfo, +) { + const serverSlotNames = + direction == IoDirection.INPUT + ? Object.keys(serverNodeData.input?.optional || {}) + : serverNodeData.output_name; + const serverSlotTypes = + direction == IoDirection.INPUT + ? (Object.values(serverNodeData.input?.optional || {}).map((i) => i[0]) as string[]) + : serverNodeData.output; + const slots = direction == IoDirection.INPUT ? node.inputs : node.outputs; + + // Let's go through the node data names and make sure our current ones match, and update if not. + let firstIndex = slots.findIndex((o, i) => i !== serverSlotNames.indexOf(o.name)); + if (firstIndex > -1) { + // Have mismatches. First, let's go through and save all our links by name. + const links: { [key: string]: { id: number; link: LLink }[] } = {}; + slots.map((slot) => { + // There's a chance we have duplicate names on an upgrade, so we'll collect all links to one + // name so we don't ovewrite our list per name. + links[slot.name] = links[slot.name] || []; + links[slot.name]?.push(...getSlotLinks(slot)); + }); + + // Now, go through and rearrange outputs by splicing + for (const [index, serverSlotName] of serverSlotNames.entries()) { + const currentNodeSlot = slots.map((s) => s.name).indexOf(serverSlotName); + if (currentNodeSlot > -1) { + if (currentNodeSlot != index) { + const splicedItem = slots.splice(currentNodeSlot, 1)[0]!; + slots.splice(index, 0, splicedItem as any); + } + } else if (currentNodeSlot === -1) { + const splicedItem = { + name: serverSlotName, + type: serverSlotTypes![index], + links: [], + }; + slots.splice(index, 0, splicedItem as any); + } + } + + if (slots.length > serverSlotNames.length) { + for (let i = slots.length - 1; i > serverSlotNames.length - 1; i--) { + if (direction == IoDirection.INPUT) { + node.disconnectInput(i); + node.removeInput(i); + } else { + node.disconnectOutput(i); + node.removeOutput(i); + } + } + } + + // Now, go through the link data again and make sure the origin_slot is the correct slot. + for (const [name, slotLinks] of Object.entries(links)) { + let currentNodeSlot = slots.map((s) => s.name).indexOf(name); + if (currentNodeSlot > -1) { + for (const linkData of slotLinks) { + if (direction == IoDirection.INPUT) { + linkData.link.target_slot = currentNodeSlot; + } else { + linkData.link.origin_slot = currentNodeSlot; + // If our next node is a Reroute, then let's get it to update the type. + const nextNode = app.graph.getNodeById(linkData.link.target_id); + // (Check nextNode, as sometimes graphs seem to have very stale data and that node id + // doesn't exist). + if ( + nextNode && + (nextNode.constructor as ComfyNodeConstructor)?.type!.includes("Reroute") + ) { + (nextNode as any).stabilize && (nextNode as any).stabilize(); + } + } + } + } + } + } +} + +export function isValidConnection(ioA?: INodeSlot | null, ioB?: INodeSlot | null) { + if (!ioA || !ioB) { + return false; + } + const typeA = String(ioA.type); + const typeB = String(ioB.type); + // What does litegraph think, which includes looking at array values. + let isValid = LiteGraph.isValidConnection(typeA, typeB); + + // This is here to fix the churn happening in list types in comfyui itself.. + // https://github.com/comfyanonymous/ComfyUI/issues/1674 + if (!isValid) { + let areCombos = + (typeA.includes(",") && typeB === "COMBO") || (typeA === "COMBO" && typeB.includes(",")); + // We don't want to let any old combo connect to any old combo, so we'll look at the names too. + if (areCombos) { + // Some nodes use "_name" and some use "model" and "ckpt", so normalize + const nameA = ioA.name.toUpperCase().replace("_NAME", "").replace("CKPT", "MODEL"); + const nameB = ioB.name.toUpperCase().replace("_NAME", "").replace("CKPT", "MODEL"); + isValid = nameA.includes(nameB) || nameB.includes(nameA); + } + } + return isValid; +} + +/** + * Patches the LiteGraph.isValidConnection so old nodes can connect to this new COMBO type for all + * lists (without users needing to go through and re-create all their nodes one by one). + */ +const oldIsValidConnection = LiteGraph.isValidConnection; +LiteGraph.isValidConnection = function (typeA: string | string[], typeB: string | string[]) { + let isValid = oldIsValidConnection.call(LiteGraph, typeA, typeB); + if (!isValid) { + typeA = String(typeA); + typeB = String(typeB); + // This is waaaay too liberal and now any combos can connect to any combos. But we only have the + // types (not names like my util above), and connecting too liberally is better than old nodes + // with lists not being able to connect to this new COMBO type. And, anyway, it matches the + // current behavior today with new nodes anyway, where all lists are COMBO types. + // Refs: https://github.com/comfyanonymous/ComfyUI/issues/1674 + // https://github.com/comfyanonymous/ComfyUI/pull/1675 + let areCombos = + (typeA.includes(",") && typeB === "COMBO") || (typeA === "COMBO" && typeB.includes(",")); + isValid = areCombos; + } + return isValid; +}; diff --git a/src_web/comfyui/utils_inputs_outputs.ts b/src_web/comfyui/utils_inputs_outputs.ts index c965154..ae007d5 100644 --- a/src_web/comfyui/utils_inputs_outputs.ts +++ b/src_web/comfyui/utils_inputs_outputs.ts @@ -1,14 +1,14 @@ -import type { LGraphNode } from "typings/litegraph"; - -/** Removes all inputs from the end. */ -export function removeUnusedInputsFromEnd(node: LGraphNode, minNumber = 1, nameMatch?: RegExp) { - for (let i = node.inputs.length - 1; i >= minNumber; i--) { - if (!node.inputs[i]?.link) { - if (!nameMatch || nameMatch.test(node.inputs[i]!.name)) { - node.removeInput(i); - } - continue; - } - break; - } +import type { LGraphNode } from "typings/litegraph.js"; + +/** Removes all inputs from the end. */ +export function removeUnusedInputsFromEnd(node: LGraphNode, minNumber = 1, nameMatch?: RegExp) { + for (let i = node.inputs.length - 1; i >= minNumber; i--) { + if (!node.inputs[i]?.link) { + if (!nameMatch || nameMatch.test(node.inputs[i]!.name)) { + node.removeInput(i); + } + continue; + } + break; + } } \ No newline at end of file diff --git a/src_web/common/link_fixer.ts b/src_web/common/link_fixer.ts index 651ff23..54491da 100644 --- a/src_web/common/link_fixer.ts +++ b/src_web/common/link_fixer.ts @@ -1,392 +1,392 @@ -import type { BadLinksData, SerializedGraph, SerializedLink, SerializedNode } from "typings/index"; -import type { LGraph, LGraphNode, LLink, serializedLGraph } from "typings/litegraph"; - -enum IoDirection { - INPUT, - OUTPUT, -} - -function getNodeById(graph: SerializedGraph | LGraph | serializedLGraph, id: number) { - if ((graph as LGraph).getNodeById) { - return (graph as LGraph).getNodeById(id); - } - graph = graph as SerializedGraph; - return graph.nodes.find((n) => n.id === id)!; -} - -function extendLink(link: SerializedLink) { - return { - link: link, - id: link[0], - origin_id: link[1], - origin_slot: link[2], - target_id: link[3], - target_slot: link[4], - type: link[5], - }; -} - -/** - * Takes a SerializedGraph or live LGraph and inspects the links and nodes to ensure the linking - * makes logical sense. Can apply fixes when passed the `fix` argument as true. - * - * Note that fixes are a best-effort attempt. Seems to get it correct in most cases, but there is a - * chance it correct an anomoly that results in placing an incorrect link (say, if there were two - * links in the data). Users should take care to not overwrite work until manually checking the - * result. - */ -export function fixBadLinks( - graph: SerializedGraph | LGraph, - fix = false, - silent = false, - logger: { log: (...args: any[]) => void } = console, -): BadLinksData { - const patchedNodeSlots: { - [nodeId: string]: { - inputs?: { [slot: number]: number | null }; - outputs?: { - [slots: number]: { - links: number[]; - changes: { [linkId: number]: "ADD" | "REMOVE" }; - }; - }; - }; - } = {}; - // const logger = this.newLogSession("[findBadLinks]"); - const data: { patchedNodes: Array; deletedLinks: number[] } = { - patchedNodes: [], - deletedLinks: [], - }; - - /** - * Internal patch node. We keep track of changes in patchedNodeSlots in case we're in a dry run. - */ - async function patchNodeSlot( - node: SerializedNode | LGraphNode, - ioDir: IoDirection, - slot: number, - linkId: number, - op: "ADD" | "REMOVE", - ) { - patchedNodeSlots[node.id] = patchedNodeSlots[node.id] || {}; - const patchedNode = patchedNodeSlots[node.id]!; - if (ioDir == IoDirection.INPUT) { - patchedNode["inputs"] = patchedNode["inputs"] || {}; - // We can set to null (delete), so undefined means we haven't set it at all. - if (patchedNode["inputs"]![slot] !== undefined) { - !silent && - logger.log( - ` > Already set ${node.id}.inputs[${slot}] to ${patchedNode["inputs"]![ - slot - ]!} Skipping.`, - ); - return false; - } - let linkIdToSet = op === "REMOVE" ? null : linkId; - patchedNode["inputs"]![slot] = linkIdToSet; - if (fix) { - // node.inputs[slot]!.link = linkIdToSet; - } - } else { - patchedNode["outputs"] = patchedNode["outputs"] || {}; - patchedNode["outputs"]![slot] = patchedNode["outputs"]![slot] || { - links: [...(node.outputs?.[slot]?.links || [])], - changes: {}, - }; - if (patchedNode["outputs"]![slot]!["changes"]![linkId] !== undefined) { - !silent && - logger.log( - ` > Already set ${node.id}.outputs[${slot}] to ${ - patchedNode["inputs"]![slot] - }! Skipping.`, - ); - return false; - } - patchedNode["outputs"]![slot]!["changes"]![linkId] = op; - if (op === "ADD") { - let linkIdIndex = patchedNode["outputs"]![slot]!["links"].indexOf(linkId); - if (linkIdIndex !== -1) { - !silent && logger.log(` > Hmmm.. asked to add ${linkId} but it is already in list...`); - return false; - } - patchedNode["outputs"]![slot]!["links"].push(linkId); - if (fix) { - node.outputs = node.outputs || []; - node.outputs[slot] = node.outputs[slot] || ({} as any); - node.outputs[slot]!.links = node.outputs[slot]!.links || []; - node.outputs[slot]!.links!.push(linkId); - } - } else { - let linkIdIndex = patchedNode["outputs"]![slot]!["links"].indexOf(linkId); - if (linkIdIndex === -1) { - !silent && logger.log(` > Hmmm.. asked to remove ${linkId} but it doesn't exist...`); - return false; - } - patchedNode["outputs"]![slot]!["links"].splice(linkIdIndex, 1); - if (fix) { - node.outputs?.[slot]!.links!.splice(linkIdIndex, 1); - } - } - } - data.patchedNodes.push(node); - return true; - } - - /** - * Internal to check if a node (or patched data) has a linkId. - */ - function nodeHasLinkId( - node: SerializedNode | LGraphNode, - ioDir: IoDirection, - slot: number, - linkId: number, - ) { - // Patched data should be canonical. We can double check if fixing too. - let has = false; - if (ioDir === IoDirection.INPUT) { - let nodeHasIt = node.inputs?.[slot]?.link === linkId; - if (patchedNodeSlots[node.id]?.["inputs"]) { - let patchedHasIt = patchedNodeSlots[node.id]!["inputs"]![slot] === linkId; - // If we're fixing, double check that node matches. - if (fix && nodeHasIt !== patchedHasIt) { - throw Error("Error. Expected node to match patched data."); - } - has = patchedHasIt; - } else { - has = !!nodeHasIt; - } - } else { - let nodeHasIt = node.outputs?.[slot]?.links?.includes(linkId); - if (patchedNodeSlots[node.id]?.["outputs"]?.[slot]?.["changes"][linkId]) { - let patchedHasIt = patchedNodeSlots[node.id]!["outputs"]![slot]?.links.includes(linkId); - // If we're fixing, double check that node matches. - if (fix && nodeHasIt !== patchedHasIt) { - throw Error("Error. Expected node to match patched data."); - } - has = !!patchedHasIt; - } else { - has = !!nodeHasIt; - } - } - return has; - } - - /** - * Internal to check if a node (or patched data) has a linkId. - */ - function nodeHasAnyLink(node: SerializedNode | LGraphNode, ioDir: IoDirection, slot: number) { - // Patched data should be canonical. We can double check if fixing too. - let hasAny = false; - if (ioDir === IoDirection.INPUT) { - let nodeHasAny = node.inputs?.[slot]?.link != null; - if (patchedNodeSlots[node.id]?.["inputs"]) { - let patchedHasAny = patchedNodeSlots[node.id]!["inputs"]![slot] != null; - // If we're fixing, double check that node matches. - if (fix && nodeHasAny !== patchedHasAny) { - throw Error("Error. Expected node to match patched data."); - } - hasAny = patchedHasAny; - } else { - hasAny = !!nodeHasAny; - } - } else { - let nodeHasAny = node.outputs?.[slot]?.links?.length; - if (patchedNodeSlots[node.id]?.["outputs"]?.[slot]?.["changes"]) { - let patchedHasAny = patchedNodeSlots[node.id]!["outputs"]![slot]?.links.length; - // If we're fixing, double check that node matches. - if (fix && nodeHasAny !== patchedHasAny) { - throw Error("Error. Expected node to match patched data."); - } - hasAny = !!patchedHasAny; - } else { - hasAny = !!nodeHasAny; - } - } - return hasAny; - } - - let links: Array = []; - if (!Array.isArray(graph.links)) { - Object.values(graph.links).reduce((acc, v) => { - acc[v.id] = v; - return acc; - }, links); - } else { - links = graph.links; - } - - const linksReverse = [...links]; - linksReverse.reverse(); - for (let l of linksReverse) { - if (!l) continue; - const link = (l as LLink).origin_slot != null ? (l as LLink) : extendLink(l as SerializedLink); - - const originNode = getNodeById(graph, link.origin_id); - const originHasLink = () => - nodeHasLinkId(originNode!, IoDirection.OUTPUT, link.origin_slot, link.id); - const patchOrigin = (op: "ADD" | "REMOVE", id = link.id) => - patchNodeSlot(originNode!, IoDirection.OUTPUT, link.origin_slot, id, op); - - const targetNode = getNodeById(graph, link.target_id); - const targetHasLink = () => - nodeHasLinkId(targetNode!, IoDirection.INPUT, link.target_slot, link.id); - const targetHasAnyLink = () => nodeHasAnyLink(targetNode!, IoDirection.INPUT, link.target_slot); - const patchTarget = (op: "ADD" | "REMOVE", id = link.id) => - patchNodeSlot(targetNode!, IoDirection.INPUT, link.target_slot, id, op); - - const originLog = `origin(${link.origin_id}).outputs[${link.origin_slot}].links`; - const targetLog = `target(${link.target_id}).inputs[${link.target_slot}].link`; - - if (!originNode || !targetNode) { - if (!originNode && !targetNode) { - !silent && - logger.log( - `Link ${link.id} is invalid, ` + - `both origin ${link.origin_id} and target ${link.target_id} do not exist`, - ); - } else if (!originNode) { - !silent && - logger.log( - `Link ${link.id} is funky... ` + - `origin ${link.origin_id} does not exist, but target ${link.target_id} does.`, - ); - if (targetHasLink()) { - !silent && - logger.log( - ` > [PATCH] ${targetLog} does have link, will remove the inputs' link first.`, - ); - patchTarget("REMOVE", -1); - } - } else if (!targetNode) { - !silent && - logger.log( - `Link ${link.id} is funky... ` + - `target ${link.target_id} does not exist, but origin ${link.origin_id} does.`, - ); - if (originHasLink()) { - !silent && - logger.log(` > [PATCH] Origin's links' has ${link.id}; will remove the link first.`); - patchOrigin("REMOVE"); - } - } - continue; - } - - if (targetHasLink() || originHasLink()) { - if (!originHasLink()) { - !silent && - logger.log( - `${link.id} is funky... ${originLog} does NOT contain it, but ${targetLog} does.`, - ); - !silent && - logger.log(` > [PATCH] Attempt a fix by adding this ${link.id} to ${originLog}.`); - patchOrigin("ADD"); - } else if (!targetHasLink()) { - !silent && - logger.log( - `${link.id} is funky... ${targetLog} is NOT correct (is ${targetNode.inputs?.[ - link.target_slot - ]?.link}), but ${originLog} contains it`, - ); - if (!targetHasAnyLink()) { - !silent && logger.log(` > [PATCH] ${targetLog} is not defined, will set to ${link.id}.`); - let patched = patchTarget("ADD"); - if (!patched) { - !silent && - logger.log( - ` > [PATCH] Nvm, ${targetLog} already patched. Removing ${link.id} from ${originLog}.`, - ); - patched = patchOrigin("REMOVE"); - } - } else { - !silent && - logger.log( - ` > [PATCH] ${targetLog} is defined, removing ${link.id} from ${originLog}.`, - ); - patchOrigin("REMOVE"); - } - } - } - } - - // Now that we've cleaned up the inputs, outputs, run through it looking for dangling links., - for (let l of linksReverse) { - if (!l) continue; - const link = (l as LLink).origin_slot != null ? (l as LLink) : extendLink(l as SerializedLink); - const originNode = getNodeById(graph, link.origin_id); - const targetNode = getNodeById(graph, link.target_id); - // Now that we've manipulated the linking, check again if they both exist. - if ( - (!originNode || !nodeHasLinkId(originNode, IoDirection.OUTPUT, link.origin_slot, link.id)) && - (!targetNode || !nodeHasLinkId(targetNode, IoDirection.INPUT, link.target_slot, link.id)) - ) { - !silent && - logger.log( - `${link.id} is def invalid; BOTH origin node ${link.origin_id} ${ - !originNode ? "is removed" : `doesn\'t have ${link.id}` - } and ${link.origin_id} target node ${ - !targetNode ? "is removed" : `doesn\'t have ${link.id}` - }.`, - ); - data.deletedLinks.push(link.id); - continue; - } - } - - // If we're fixing, then we've been patching along the way. Now go through and actually delete - // the zombie links from `app.graph.links` - if (fix) { - for (let i = data.deletedLinks.length - 1; i >= 0; i--) { - !silent && logger.log(`Deleting link #${data.deletedLinks[i]}.`); - if ((graph as LGraph).getNodeById) { - delete graph.links[data.deletedLinks[i]!]; - } else { - graph = graph as SerializedGraph; - // Sometimes we got objects for links if passed after ComfyUI's loadGraphData modifies the - // data. We make a copy now, but can handle the bastardized objects just in case. - const idx = graph.links.findIndex( - (l) => l && (l[0] === data.deletedLinks[i] || (l as any).id === data.deletedLinks[i]), - ); - if (idx === -1) { - logger.log(`INDEX NOT FOUND for #${data.deletedLinks[i]}`); - } - logger.log(`splicing ${idx} from links`); - graph.links.splice(idx, 1); - } - } - // If we're a serialized graph, we can filter out the links because it's just an array. - if (!(graph as LGraph).getNodeById) { - graph.links = (graph as SerializedGraph).links.filter((l) => !!l); - } - } - if (!data.patchedNodes.length && !data.deletedLinks.length) { - return { - hasBadLinks: false, - fixed: false, - graph, - patched: data.patchedNodes.length, - deleted: data.deletedLinks.length, - }; - } - !silent && - logger.log( - `${fix ? "Made" : "Would make"} ${data.patchedNodes.length || "no"} node link patches, and ${ - data.deletedLinks.length || "no" - } stale link removals.`, - ); - - let hasBadLinks: boolean = !!(data.patchedNodes.length || data.deletedLinks.length); - // If we're fixing, then let's run it again to see if there are no more bad links. - if (fix && !silent) { - const rerun = fixBadLinks(graph, false, true); - hasBadLinks = rerun.hasBadLinks; - } - - return { - hasBadLinks, - fixed: !!hasBadLinks && fix, - graph, - patched: data.patchedNodes.length, - deleted: data.deletedLinks.length, - }; -} +import type { BadLinksData, SerializedGraph, SerializedLink, SerializedNode } from "typings/index.js"; +import type { LGraph, LGraphNode, LLink, serializedLGraph } from "typings/litegraph.js"; + +enum IoDirection { + INPUT, + OUTPUT, +} + +function getNodeById(graph: SerializedGraph | LGraph | serializedLGraph, id: number) { + if ((graph as LGraph).getNodeById) { + return (graph as LGraph).getNodeById(id); + } + graph = graph as SerializedGraph; + return graph.nodes.find((n) => n.id === id)!; +} + +function extendLink(link: SerializedLink) { + return { + link: link, + id: link[0], + origin_id: link[1], + origin_slot: link[2], + target_id: link[3], + target_slot: link[4], + type: link[5], + }; +} + +/** + * Takes a SerializedGraph or live LGraph and inspects the links and nodes to ensure the linking + * makes logical sense. Can apply fixes when passed the `fix` argument as true. + * + * Note that fixes are a best-effort attempt. Seems to get it correct in most cases, but there is a + * chance it correct an anomoly that results in placing an incorrect link (say, if there were two + * links in the data). Users should take care to not overwrite work until manually checking the + * result. + */ +export function fixBadLinks( + graph: SerializedGraph | LGraph, + fix = false, + silent = false, + logger: { log: (...args: any[]) => void } = console, +): BadLinksData { + const patchedNodeSlots: { + [nodeId: string]: { + inputs?: { [slot: number]: number | null }; + outputs?: { + [slots: number]: { + links: number[]; + changes: { [linkId: number]: "ADD" | "REMOVE" }; + }; + }; + }; + } = {}; + // const logger = this.newLogSession("[findBadLinks]"); + const data: { patchedNodes: Array; deletedLinks: number[] } = { + patchedNodes: [], + deletedLinks: [], + }; + + /** + * Internal patch node. We keep track of changes in patchedNodeSlots in case we're in a dry run. + */ + async function patchNodeSlot( + node: SerializedNode | LGraphNode, + ioDir: IoDirection, + slot: number, + linkId: number, + op: "ADD" | "REMOVE", + ) { + patchedNodeSlots[node.id] = patchedNodeSlots[node.id] || {}; + const patchedNode = patchedNodeSlots[node.id]!; + if (ioDir == IoDirection.INPUT) { + patchedNode["inputs"] = patchedNode["inputs"] || {}; + // We can set to null (delete), so undefined means we haven't set it at all. + if (patchedNode["inputs"]![slot] !== undefined) { + !silent && + logger.log( + ` > Already set ${node.id}.inputs[${slot}] to ${patchedNode["inputs"]![ + slot + ]!} Skipping.`, + ); + return false; + } + let linkIdToSet = op === "REMOVE" ? null : linkId; + patchedNode["inputs"]![slot] = linkIdToSet; + if (fix) { + // node.inputs[slot]!.link = linkIdToSet; + } + } else { + patchedNode["outputs"] = patchedNode["outputs"] || {}; + patchedNode["outputs"]![slot] = patchedNode["outputs"]![slot] || { + links: [...(node.outputs?.[slot]?.links || [])], + changes: {}, + }; + if (patchedNode["outputs"]![slot]!["changes"]![linkId] !== undefined) { + !silent && + logger.log( + ` > Already set ${node.id}.outputs[${slot}] to ${ + patchedNode["inputs"]![slot] + }! Skipping.`, + ); + return false; + } + patchedNode["outputs"]![slot]!["changes"]![linkId] = op; + if (op === "ADD") { + let linkIdIndex = patchedNode["outputs"]![slot]!["links"].indexOf(linkId); + if (linkIdIndex !== -1) { + !silent && logger.log(` > Hmmm.. asked to add ${linkId} but it is already in list...`); + return false; + } + patchedNode["outputs"]![slot]!["links"].push(linkId); + if (fix) { + node.outputs = node.outputs || []; + node.outputs[slot] = node.outputs[slot] || ({} as any); + node.outputs[slot]!.links = node.outputs[slot]!.links || []; + node.outputs[slot]!.links!.push(linkId); + } + } else { + let linkIdIndex = patchedNode["outputs"]![slot]!["links"].indexOf(linkId); + if (linkIdIndex === -1) { + !silent && logger.log(` > Hmmm.. asked to remove ${linkId} but it doesn't exist...`); + return false; + } + patchedNode["outputs"]![slot]!["links"].splice(linkIdIndex, 1); + if (fix) { + node.outputs?.[slot]!.links!.splice(linkIdIndex, 1); + } + } + } + data.patchedNodes.push(node); + return true; + } + + /** + * Internal to check if a node (or patched data) has a linkId. + */ + function nodeHasLinkId( + node: SerializedNode | LGraphNode, + ioDir: IoDirection, + slot: number, + linkId: number, + ) { + // Patched data should be canonical. We can double check if fixing too. + let has = false; + if (ioDir === IoDirection.INPUT) { + let nodeHasIt = node.inputs?.[slot]?.link === linkId; + if (patchedNodeSlots[node.id]?.["inputs"]) { + let patchedHasIt = patchedNodeSlots[node.id]!["inputs"]![slot] === linkId; + // If we're fixing, double check that node matches. + if (fix && nodeHasIt !== patchedHasIt) { + throw Error("Error. Expected node to match patched data."); + } + has = patchedHasIt; + } else { + has = !!nodeHasIt; + } + } else { + let nodeHasIt = node.outputs?.[slot]?.links?.includes(linkId); + if (patchedNodeSlots[node.id]?.["outputs"]?.[slot]?.["changes"][linkId]) { + let patchedHasIt = patchedNodeSlots[node.id]!["outputs"]![slot]?.links.includes(linkId); + // If we're fixing, double check that node matches. + if (fix && nodeHasIt !== patchedHasIt) { + throw Error("Error. Expected node to match patched data."); + } + has = !!patchedHasIt; + } else { + has = !!nodeHasIt; + } + } + return has; + } + + /** + * Internal to check if a node (or patched data) has a linkId. + */ + function nodeHasAnyLink(node: SerializedNode | LGraphNode, ioDir: IoDirection, slot: number) { + // Patched data should be canonical. We can double check if fixing too. + let hasAny = false; + if (ioDir === IoDirection.INPUT) { + let nodeHasAny = node.inputs?.[slot]?.link != null; + if (patchedNodeSlots[node.id]?.["inputs"]) { + let patchedHasAny = patchedNodeSlots[node.id]!["inputs"]![slot] != null; + // If we're fixing, double check that node matches. + if (fix && nodeHasAny !== patchedHasAny) { + throw Error("Error. Expected node to match patched data."); + } + hasAny = patchedHasAny; + } else { + hasAny = !!nodeHasAny; + } + } else { + let nodeHasAny = node.outputs?.[slot]?.links?.length; + if (patchedNodeSlots[node.id]?.["outputs"]?.[slot]?.["changes"]) { + let patchedHasAny = patchedNodeSlots[node.id]!["outputs"]![slot]?.links.length; + // If we're fixing, double check that node matches. + if (fix && nodeHasAny !== patchedHasAny) { + throw Error("Error. Expected node to match patched data."); + } + hasAny = !!patchedHasAny; + } else { + hasAny = !!nodeHasAny; + } + } + return hasAny; + } + + let links: Array = []; + if (!Array.isArray(graph.links)) { + Object.values(graph.links).reduce((acc, v) => { + acc[v.id] = v; + return acc; + }, links); + } else { + links = graph.links; + } + + const linksReverse = [...links]; + linksReverse.reverse(); + for (let l of linksReverse) { + if (!l) continue; + const link = (l as LLink).origin_slot != null ? (l as LLink) : extendLink(l as SerializedLink); + + const originNode = getNodeById(graph, link.origin_id); + const originHasLink = () => + nodeHasLinkId(originNode!, IoDirection.OUTPUT, link.origin_slot, link.id); + const patchOrigin = (op: "ADD" | "REMOVE", id = link.id) => + patchNodeSlot(originNode!, IoDirection.OUTPUT, link.origin_slot, id, op); + + const targetNode = getNodeById(graph, link.target_id); + const targetHasLink = () => + nodeHasLinkId(targetNode!, IoDirection.INPUT, link.target_slot, link.id); + const targetHasAnyLink = () => nodeHasAnyLink(targetNode!, IoDirection.INPUT, link.target_slot); + const patchTarget = (op: "ADD" | "REMOVE", id = link.id) => + patchNodeSlot(targetNode!, IoDirection.INPUT, link.target_slot, id, op); + + const originLog = `origin(${link.origin_id}).outputs[${link.origin_slot}].links`; + const targetLog = `target(${link.target_id}).inputs[${link.target_slot}].link`; + + if (!originNode || !targetNode) { + if (!originNode && !targetNode) { + !silent && + logger.log( + `Link ${link.id} is invalid, ` + + `both origin ${link.origin_id} and target ${link.target_id} do not exist`, + ); + } else if (!originNode) { + !silent && + logger.log( + `Link ${link.id} is funky... ` + + `origin ${link.origin_id} does not exist, but target ${link.target_id} does.`, + ); + if (targetHasLink()) { + !silent && + logger.log( + ` > [PATCH] ${targetLog} does have link, will remove the inputs' link first.`, + ); + patchTarget("REMOVE", -1); + } + } else if (!targetNode) { + !silent && + logger.log( + `Link ${link.id} is funky... ` + + `target ${link.target_id} does not exist, but origin ${link.origin_id} does.`, + ); + if (originHasLink()) { + !silent && + logger.log(` > [PATCH] Origin's links' has ${link.id}; will remove the link first.`); + patchOrigin("REMOVE"); + } + } + continue; + } + + if (targetHasLink() || originHasLink()) { + if (!originHasLink()) { + !silent && + logger.log( + `${link.id} is funky... ${originLog} does NOT contain it, but ${targetLog} does.`, + ); + !silent && + logger.log(` > [PATCH] Attempt a fix by adding this ${link.id} to ${originLog}.`); + patchOrigin("ADD"); + } else if (!targetHasLink()) { + !silent && + logger.log( + `${link.id} is funky... ${targetLog} is NOT correct (is ${targetNode.inputs?.[ + link.target_slot + ]?.link}), but ${originLog} contains it`, + ); + if (!targetHasAnyLink()) { + !silent && logger.log(` > [PATCH] ${targetLog} is not defined, will set to ${link.id}.`); + let patched = patchTarget("ADD"); + if (!patched) { + !silent && + logger.log( + ` > [PATCH] Nvm, ${targetLog} already patched. Removing ${link.id} from ${originLog}.`, + ); + patched = patchOrigin("REMOVE"); + } + } else { + !silent && + logger.log( + ` > [PATCH] ${targetLog} is defined, removing ${link.id} from ${originLog}.`, + ); + patchOrigin("REMOVE"); + } + } + } + } + + // Now that we've cleaned up the inputs, outputs, run through it looking for dangling links., + for (let l of linksReverse) { + if (!l) continue; + const link = (l as LLink).origin_slot != null ? (l as LLink) : extendLink(l as SerializedLink); + const originNode = getNodeById(graph, link.origin_id); + const targetNode = getNodeById(graph, link.target_id); + // Now that we've manipulated the linking, check again if they both exist. + if ( + (!originNode || !nodeHasLinkId(originNode, IoDirection.OUTPUT, link.origin_slot, link.id)) && + (!targetNode || !nodeHasLinkId(targetNode, IoDirection.INPUT, link.target_slot, link.id)) + ) { + !silent && + logger.log( + `${link.id} is def invalid; BOTH origin node ${link.origin_id} ${ + !originNode ? "is removed" : `doesn\'t have ${link.id}` + } and ${link.origin_id} target node ${ + !targetNode ? "is removed" : `doesn\'t have ${link.id}` + }.`, + ); + data.deletedLinks.push(link.id); + continue; + } + } + + // If we're fixing, then we've been patching along the way. Now go through and actually delete + // the zombie links from `app.graph.links` + if (fix) { + for (let i = data.deletedLinks.length - 1; i >= 0; i--) { + !silent && logger.log(`Deleting link #${data.deletedLinks[i]}.`); + if ((graph as LGraph).getNodeById) { + delete graph.links[data.deletedLinks[i]!]; + } else { + graph = graph as SerializedGraph; + // Sometimes we got objects for links if passed after ComfyUI's loadGraphData modifies the + // data. We make a copy now, but can handle the bastardized objects just in case. + const idx = graph.links.findIndex( + (l) => l && (l[0] === data.deletedLinks[i] || (l as any).id === data.deletedLinks[i]), + ); + if (idx === -1) { + logger.log(`INDEX NOT FOUND for #${data.deletedLinks[i]}`); + } + logger.log(`splicing ${idx} from links`); + graph.links.splice(idx, 1); + } + } + // If we're a serialized graph, we can filter out the links because it's just an array. + if (!(graph as LGraph).getNodeById) { + graph.links = (graph as SerializedGraph).links.filter((l) => !!l); + } + } + if (!data.patchedNodes.length && !data.deletedLinks.length) { + return { + hasBadLinks: false, + fixed: false, + graph, + patched: data.patchedNodes.length, + deleted: data.deletedLinks.length, + }; + } + !silent && + logger.log( + `${fix ? "Made" : "Would make"} ${data.patchedNodes.length || "no"} node link patches, and ${ + data.deletedLinks.length || "no" + } stale link removals.`, + ); + + let hasBadLinks: boolean = !!(data.patchedNodes.length || data.deletedLinks.length); + // If we're fixing, then let's run it again to see if there are no more bad links. + if (fix && !silent) { + const rerun = fixBadLinks(graph, false, true); + hasBadLinks = rerun.hasBadLinks; + } + + return { + hasBadLinks, + fixed: !!hasBadLinks && fix, + graph, + patched: data.patchedNodes.length, + deleted: data.deletedLinks.length, + }; +} diff --git a/src_web/common/model_info_service.ts b/src_web/common/model_info_service.ts index bb0f6ba..306297e 100644 --- a/src_web/common/model_info_service.ts +++ b/src_web/common/model_info_service.ts @@ -1,74 +1,74 @@ -import type { RgthreeModelInfo } from "typings/rgthree"; -import { rgthreeApi } from "./rgthree_api.js"; -import { api } from "scripts/api.js"; - -/** - * A singleton service to fetch and cache model infos from rgthree-comfy. - */ -class ModelInfoService extends EventTarget { - private readonly loraToInfo = new Map(); - - constructor() { - super(); - api.addEventListener( - "rgthree-refreshed-lora-info", - this.handleLoraAsyncUpdate.bind(this) as EventListener, - ); - } - - /** - * Single point to set data into the info cache, and fire an event. Note, this doesn't determine - * if the data is actually different. - */ - private setFreshLoraData(file: string, info: RgthreeModelInfo) { - this.loraToInfo.set(file, info); - this.dispatchEvent( - new CustomEvent("rgthree-model-service-lora-details", { detail: { lora: info } }), - ); - } - - async getLora(file: string, refresh = false, light = false) { - if (this.loraToInfo.has(file) && !refresh) { - return this.loraToInfo.get(file)!; - } - return this.fetchLora(file, refresh, light); - } - - async fetchLora(file: string, refresh = false, light = false) { - let info = null; - if (!refresh) { - info = await rgthreeApi.getLorasInfo(file, light); - } else { - info = await rgthreeApi.refreshLorasInfo(file); - } - if (!light) { - this.loraToInfo.set(file, info); - } - return info; - } - - async refreshLora(file: string) { - return this.fetchLora(file, true); - } - - async clearLoraFetchedData(file: string) { - await rgthreeApi.clearLorasInfo(file); - this.loraToInfo.delete(file); - return null; - } - - async saveLoraPartial(file: string, data: Partial) { - let info = await rgthreeApi.saveLoraInfo(file, data); - this.loraToInfo.set(file, info); - return info; - } - - private handleLoraAsyncUpdate(event: CustomEvent<{ data: RgthreeModelInfo }>) { - const info = event.detail?.data as RgthreeModelInfo; - if (info?.file) { - this.setFreshLoraData(info.file, info); - } - } -} - -export const SERVICE = new ModelInfoService(); +import type { RgthreeModelInfo } from "typings/rgthree.js"; +import { rgthreeApi } from "./rgthree_api.js"; +import { api } from "scripts/api.js"; + +/** + * A singleton service to fetch and cache model infos from rgthree-comfy. + */ +class ModelInfoService extends EventTarget { + private readonly loraToInfo = new Map(); + + constructor() { + super(); + api.addEventListener( + "rgthree-refreshed-lora-info", + this.handleLoraAsyncUpdate.bind(this) as EventListener, + ); + } + + /** + * Single point to set data into the info cache, and fire an event. Note, this doesn't determine + * if the data is actually different. + */ + private setFreshLoraData(file: string, info: RgthreeModelInfo) { + this.loraToInfo.set(file, info); + this.dispatchEvent( + new CustomEvent("rgthree-model-service-lora-details", { detail: { lora: info } }), + ); + } + + async getLora(file: string, refresh = false, light = false) { + if (this.loraToInfo.has(file) && !refresh) { + return this.loraToInfo.get(file)!; + } + return this.fetchLora(file, refresh, light); + } + + async fetchLora(file: string, refresh = false, light = false) { + let info = null; + if (!refresh) { + info = await rgthreeApi.getLorasInfo(file, light); + } else { + info = await rgthreeApi.refreshLorasInfo(file); + } + if (!light) { + this.loraToInfo.set(file, info); + } + return info; + } + + async refreshLora(file: string) { + return this.fetchLora(file, true); + } + + async clearLoraFetchedData(file: string) { + await rgthreeApi.clearLorasInfo(file); + this.loraToInfo.delete(file); + return null; + } + + async saveLoraPartial(file: string, data: Partial) { + let info = await rgthreeApi.saveLoraInfo(file, data); + this.loraToInfo.set(file, info); + return info; + } + + private handleLoraAsyncUpdate(event: CustomEvent<{ data: RgthreeModelInfo }>) { + const info = event.detail?.data as RgthreeModelInfo; + if (info?.file) { + this.setFreshLoraData(info.file, info); + } + } +} + +export const SERVICE = new ModelInfoService(); diff --git a/src_web/common/rgthree_api.ts b/src_web/common/rgthree_api.ts index 994e43a..c9b59b8 100644 --- a/src_web/common/rgthree_api.ts +++ b/src_web/common/rgthree_api.ts @@ -1,99 +1,99 @@ -import type { RgthreeModelInfo } from "typings/rgthree"; - -class RgthreeApi { - private baseUrl: string; - getCheckpointsPromise: Promise | null = null; - getSamplersPromise: Promise | null = null; - getSchedulersPromise: Promise | null = null; - getLorasPromise: Promise | null = null; - getWorkflowsPromise: Promise | null = null; - - constructor(baseUrl?: string) { - this.baseUrl = baseUrl || "./rgthree/api"; - } - - apiURL(route: string) { - return `${this.baseUrl}${route}`; - } - - fetchApi(route: string, options?: RequestInit) { - return fetch(this.apiURL(route), options); - } - - async fetchJson(route: string, options?: RequestInit) { - const r = await this.fetchApi(route, options); - return await r.json(); - } - - async postJson(route: string, json: any) { - const body = new FormData(); - body.append("json", JSON.stringify(json)); - return await rgthreeApi.fetchJson(route, { method: "POST", body }); - } - - getLoras(force = false) { - if (!this.getLorasPromise || force) { - this.getLorasPromise = this.fetchJson("/loras", { cache: "no-store" }); - } - return this.getLorasPromise; - } - - async fetchApiJsonOrNull(route: string, options?: RequestInit) { - const response = await this.fetchJson(route, options); - if (response.status === 200 && response.data) { - return (response.data as T) || null; - } - return null; - } - - /** - * Fetches the lora information. - * - * @param light Whether or not to generate a json file if there isn't one. This isn't necessary if - * we're just checking for values, but is more necessary when opening an info dialog. - */ - async getLorasInfo(lora: string, light?: boolean): Promise; - async getLorasInfo(light?: boolean): Promise; - async getLorasInfo(...args: any) { - const params = new URLSearchParams(); - const isSingleLora = typeof args[0] == 'string'; - if (isSingleLora) { - params.set("file", args[0]); - } - params.set("light", (isSingleLora ? args[1] : args[0]) === false ? '0' : '1'); - const path = `/loras/info?` + params.toString(); - return await this.fetchApiJsonOrNull(path); - } - - async refreshLorasInfo(file: string): Promise; - async refreshLorasInfo(): Promise; - async refreshLorasInfo(file?: string) { - const path = `/loras/info/refresh` + (file ? `?file=${encodeURIComponent(file)}` : ''); - const infos = await this.fetchApiJsonOrNull(path); - return infos; - } - - async clearLorasInfo(file?: string): Promise { - const path = `/loras/info/clear` + (file ? `?file=${encodeURIComponent(file)}` : ''); - await this.fetchApiJsonOrNull(path); - return; - } - - /** - * Saves partial data sending it to the backend.. - */ - async saveLoraInfo( - lora: string, - data: Partial, - ): Promise { - const body = new FormData(); - body.append("json", JSON.stringify(data)); - return await this.fetchApiJsonOrNull( - `/loras/info?file=${encodeURIComponent(lora)}`, - { cache: "no-store", method: "POST", body }, - ); - } - -} - -export const rgthreeApi = new RgthreeApi(); +import type { RgthreeModelInfo } from "typings/rgthree.js"; + +class RgthreeApi { + private baseUrl: string; + getCheckpointsPromise: Promise | null = null; + getSamplersPromise: Promise | null = null; + getSchedulersPromise: Promise | null = null; + getLorasPromise: Promise | null = null; + getWorkflowsPromise: Promise | null = null; + + constructor(baseUrl?: string) { + this.baseUrl = baseUrl || "./rgthree/api"; + } + + apiURL(route: string) { + return `${this.baseUrl}${route}`; + } + + fetchApi(route: string, options?: RequestInit) { + return fetch(this.apiURL(route), options); + } + + async fetchJson(route: string, options?: RequestInit) { + const r = await this.fetchApi(route, options); + return await r.json(); + } + + async postJson(route: string, json: any) { + const body = new FormData(); + body.append("json", JSON.stringify(json)); + return await rgthreeApi.fetchJson(route, { method: "POST", body }); + } + + getLoras(force = false) { + if (!this.getLorasPromise || force) { + this.getLorasPromise = this.fetchJson("/loras", { cache: "no-store" }); + } + return this.getLorasPromise; + } + + async fetchApiJsonOrNull(route: string, options?: RequestInit) { + const response = await this.fetchJson(route, options); + if (response.status === 200 && response.data) { + return (response.data as T) || null; + } + return null; + } + + /** + * Fetches the lora information. + * + * @param light Whether or not to generate a json file if there isn't one. This isn't necessary if + * we're just checking for values, but is more necessary when opening an info dialog. + */ + async getLorasInfo(lora: string, light?: boolean): Promise; + async getLorasInfo(light?: boolean): Promise; + async getLorasInfo(...args: any) { + const params = new URLSearchParams(); + const isSingleLora = typeof args[0] == 'string'; + if (isSingleLora) { + params.set("file", args[0]); + } + params.set("light", (isSingleLora ? args[1] : args[0]) === false ? '0' : '1'); + const path = `/loras/info?` + params.toString(); + return await this.fetchApiJsonOrNull(path); + } + + async refreshLorasInfo(file: string): Promise; + async refreshLorasInfo(): Promise; + async refreshLorasInfo(file?: string) { + const path = `/loras/info/refresh` + (file ? `?file=${encodeURIComponent(file)}` : ''); + const infos = await this.fetchApiJsonOrNull(path); + return infos; + } + + async clearLorasInfo(file?: string): Promise { + const path = `/loras/info/clear` + (file ? `?file=${encodeURIComponent(file)}` : ''); + await this.fetchApiJsonOrNull(path); + return; + } + + /** + * Saves partial data sending it to the backend.. + */ + async saveLoraInfo( + lora: string, + data: Partial, + ): Promise { + const body = new FormData(); + body.append("json", JSON.stringify(data)); + return await this.fetchApiJsonOrNull( + `/loras/info?file=${encodeURIComponent(lora)}`, + { cache: "no-store", method: "POST", body }, + ); + } + +} + +export const rgthreeApi = new RgthreeApi(); diff --git a/src_web/link_fixer/link_page.ts b/src_web/link_fixer/link_page.ts index d16f35b..d266ef2 100644 --- a/src_web/link_fixer/link_page.ts +++ b/src_web/link_fixer/link_page.ts @@ -1,235 +1,235 @@ -import type { SerializedGraph, BadLinksData } from "typings/index"; -import { fixBadLinks } from "../common/link_fixer.js"; -import { getPngMetadata } from "scripts/pnginfo.js"; - -function wait(ms = 16, value?: any) { - return new Promise((resolve) => { - setTimeout(() => { - resolve(value); - }, ms); - }); -} - -const logger = { - logTo: console as Console | HTMLElement, - log: (...args: any[]) => { - logger.logTo === console - ? console.log(...args) - : ((logger.logTo as HTMLElement).innerText += args.join(",") + "\n"); - }, -}; - -const findBadLinksLogger = { - log: async (...args: any[]) => { - logger.log(...args); - // await wait(48); - }, -}; - -export class LinkPage { - private containerEl: HTMLDivElement; - private figcaptionEl: HTMLElement; - private btnFix: HTMLButtonElement; - private outputeMessageEl: HTMLDivElement; - private outputImageEl: HTMLImageElement; - - private file?: File | Blob; - private graph?: SerializedGraph; - private graphResults?: BadLinksData; - private graphFinalResults?: BadLinksData; - - constructor() { - // const consoleEl = document.getElementById("console")!; - this.containerEl = document.querySelector(".box")!; - this.figcaptionEl = document.querySelector("figcaption")!; - this.outputeMessageEl = document.querySelector(".output")!; - this.outputImageEl = document.querySelector(".output-image")!; - this.btnFix = document.querySelector(".btn-fix")!; - - // Need to prevent on dragover to allow drop... - document.addEventListener( - "dragover", - (e) => { - e.preventDefault(); - }, - false, - ); - document.addEventListener("drop", (e) => { - this.onDrop(e); - }); - this.btnFix.addEventListener("click", (e) => { - this.onFixClick(e); - }); - } - - private async onFixClick(e: MouseEvent) { - if (!this.graphResults || !this.graph) { - this.updateUi("โ›” Fix button click without results."); - return; - } - // Fix - let graphFinalResults = fixBadLinks(this.graph, true); - // Confirm - graphFinalResults = fixBadLinks(graphFinalResults.graph, true); - // This should have happened, but try to run it through again if there's till an issue. - if (graphFinalResults.patched || graphFinalResults.deleted) { - graphFinalResults = fixBadLinks(graphFinalResults.graph, true); - } - this.graphFinalResults = graphFinalResults; - - await this.saveFixedWorkflow(); - - if (graphFinalResults.hasBadLinks) { - this.updateUi( - "โ›” Hmm... Still detecting bad links. Can you file an issue at https://github.com/rgthree/rgthree-comfy/issues with your image/workflow.", - ); - } else { - this.updateUi( - "โœ… Workflow fixed.

Please load new saved workflow json and double check linking and execution.", - ); - } - } - - private async onDrop(event: DragEvent) { - if (!event.dataTransfer) { - return; - } - this.reset(); - - event.preventDefault(); - event.stopPropagation(); - - // Dragging from Chrome->Firefox there is a file but its a bmp, so ignore that - if (event.dataTransfer.files.length && event.dataTransfer.files?.[0]?.type !== "image/bmp") { - await this.handleFile(event.dataTransfer.files[0]!); - return; - } - - // Try loading the first URI in the transfer list - const validTypes = ["text/uri-list", "text/x-moz-url"]; - const match = [...event.dataTransfer.types].find((t) => validTypes.find((v) => t === v)); - if (match) { - const uri = event.dataTransfer.getData(match)?.split("\n")?.[0]; - if (uri) { - await this.handleFile(await (await fetch(uri)).blob()); - } - } - } - - reset() { - this.file = undefined; - this.graph = undefined; - this.graphResults = undefined; - this.graphFinalResults = undefined; - this.updateUi(); - } - - private updateUi(msg?: string) { - this.outputeMessageEl.innerHTML = ""; - if (this.file && !this.containerEl.classList.contains("-has-file")) { - this.containerEl.classList.add("-has-file"); - this.figcaptionEl.innerHTML = (this.file as File).name || this.file.type; - if (this.file.type === "application/json") { - this.outputImageEl.src = "icon_file_json.png"; - } else { - const reader = new FileReader(); - reader.onload = () => (this.outputImageEl.src = reader.result as string); - reader.readAsDataURL(this.file); - } - } else if (!this.file && this.containerEl.classList.contains("-has-file")) { - this.containerEl.classList.remove("-has-file"); - this.outputImageEl.src = ""; - this.outputImageEl.removeAttribute("src"); - } - - if (this.graphResults) { - this.containerEl.classList.add("-has-results"); - if (!this.graphResults.patched && !this.graphResults.deleted) { - this.outputeMessageEl.innerHTML = "โœ… No bad links detected in the workflow."; - } else { - this.containerEl.classList.add("-has-fixable-results"); - this.outputeMessageEl.innerHTML = `โš ๏ธ Found ${this.graphResults.patched} links to fix, and ${this.graphResults.deleted} to be removed.`; - } - } else { - this.containerEl.classList.remove("-has-results"); - this.containerEl.classList.remove("-has-fixable-results"); - } - - if (msg) { - this.outputeMessageEl.innerHTML = msg; - } - } - - private async handleFile(file: File | Blob) { - this.file = file; - this.updateUi(); - - let workflow: string | undefined | null = null; - if (file.type.startsWith("image/")) { - const pngInfo = await getPngMetadata(file); - workflow = pngInfo?.workflow; - } else if ( - file.type === "application/json" || - (file instanceof File && file.name.endsWith(".json")) - ) { - workflow = await new Promise((resolve) => { - const reader = new FileReader(); - reader.onload = () => { - resolve(reader.result as string); - }; - reader.readAsText(file); - }); - } - if (!workflow) { - this.updateUi("โ›” No workflow found in dropped item."); - } else { - try { - this.graph = JSON.parse(workflow); - } catch (e) { - this.graph = undefined; - } - if (!this.graph) { - this.updateUi("โ›” Invalid workflow found in dropped item."); - } else { - this.loadGraphData(this.graph); - } - } - } - - private async loadGraphData(graphData: SerializedGraph) { - this.graphResults = await fixBadLinks(graphData); - this.updateUi(); - } - - private async saveFixedWorkflow() { - if (!this.graphFinalResults) { - this.updateUi("โ›” Save w/o final graph patched."); - return false; - } - - let filename: string | null = (this.file as File).name || "workflow.json"; - let filenames = filename.split("."); - filenames.pop(); - filename = filenames.join("."); - filename += "_fixed.json"; - filename = prompt("Save workflow as:", filename); - if (!filename) return false; - if (!filename.toLowerCase().endsWith(".json")) { - filename += ".json"; - } - const json = JSON.stringify(this.graphFinalResults.graph, null, 2); - const blob = new Blob([json], { type: "application/json" }); - const url = URL.createObjectURL(blob); - const anchor = document.createElement("a"); - anchor.download = filename; - anchor.href = url; - anchor.style.display = "none"; - document.body.appendChild(anchor); - await wait(); - anchor.click(); - await wait(); - anchor.remove(); - window.URL.revokeObjectURL(url); - return true; - } -} +import type { SerializedGraph, BadLinksData } from "typings/index.js"; +import { fixBadLinks } from "../common/link_fixer.js"; +import { getPngMetadata } from "scripts/pnginfo.js"; + +function wait(ms = 16, value?: any) { + return new Promise((resolve) => { + setTimeout(() => { + resolve(value); + }, ms); + }); +} + +const logger = { + logTo: console as Console | HTMLElement, + log: (...args: any[]) => { + logger.logTo === console + ? console.log(...args) + : ((logger.logTo as HTMLElement).innerText += args.join(",") + "\n"); + }, +}; + +const findBadLinksLogger = { + log: async (...args: any[]) => { + logger.log(...args); + // await wait(48); + }, +}; + +export class LinkPage { + private containerEl: HTMLDivElement; + private figcaptionEl: HTMLElement; + private btnFix: HTMLButtonElement; + private outputeMessageEl: HTMLDivElement; + private outputImageEl: HTMLImageElement; + + private file?: File | Blob; + private graph?: SerializedGraph; + private graphResults?: BadLinksData; + private graphFinalResults?: BadLinksData; + + constructor() { + // const consoleEl = document.getElementById("console")!; + this.containerEl = document.querySelector(".box")!; + this.figcaptionEl = document.querySelector("figcaption")!; + this.outputeMessageEl = document.querySelector(".output")!; + this.outputImageEl = document.querySelector(".output-image")!; + this.btnFix = document.querySelector(".btn-fix")!; + + // Need to prevent on dragover to allow drop... + document.addEventListener( + "dragover", + (e) => { + e.preventDefault(); + }, + false, + ); + document.addEventListener("drop", (e) => { + this.onDrop(e); + }); + this.btnFix.addEventListener("click", (e) => { + this.onFixClick(e); + }); + } + + private async onFixClick(e: MouseEvent) { + if (!this.graphResults || !this.graph) { + this.updateUi("โ›” Fix button click without results."); + return; + } + // Fix + let graphFinalResults = fixBadLinks(this.graph, true); + // Confirm + graphFinalResults = fixBadLinks(graphFinalResults.graph, true); + // This should have happened, but try to run it through again if there's till an issue. + if (graphFinalResults.patched || graphFinalResults.deleted) { + graphFinalResults = fixBadLinks(graphFinalResults.graph, true); + } + this.graphFinalResults = graphFinalResults; + + await this.saveFixedWorkflow(); + + if (graphFinalResults.hasBadLinks) { + this.updateUi( + "โ›” Hmm... Still detecting bad links. Can you file an issue at https://github.com/rgthree/rgthree-comfy/issues with your image/workflow.", + ); + } else { + this.updateUi( + "โœ… Workflow fixed.

Please load new saved workflow json and double check linking and execution.", + ); + } + } + + private async onDrop(event: DragEvent) { + if (!event.dataTransfer) { + return; + } + this.reset(); + + event.preventDefault(); + event.stopPropagation(); + + // Dragging from Chrome->Firefox there is a file but its a bmp, so ignore that + if (event.dataTransfer.files.length && event.dataTransfer.files?.[0]?.type !== "image/bmp") { + await this.handleFile(event.dataTransfer.files[0]!); + return; + } + + // Try loading the first URI in the transfer list + const validTypes = ["text/uri-list", "text/x-moz-url"]; + const match = [...event.dataTransfer.types].find((t) => validTypes.find((v) => t === v)); + if (match) { + const uri = event.dataTransfer.getData(match)?.split("\n")?.[0]; + if (uri) { + await this.handleFile(await (await fetch(uri)).blob()); + } + } + } + + reset() { + this.file = undefined; + this.graph = undefined; + this.graphResults = undefined; + this.graphFinalResults = undefined; + this.updateUi(); + } + + private updateUi(msg?: string) { + this.outputeMessageEl.innerHTML = ""; + if (this.file && !this.containerEl.classList.contains("-has-file")) { + this.containerEl.classList.add("-has-file"); + this.figcaptionEl.innerHTML = (this.file as File).name || this.file.type; + if (this.file.type === "application/json") { + this.outputImageEl.src = "icon_file_json.png"; + } else { + const reader = new FileReader(); + reader.onload = () => (this.outputImageEl.src = reader.result as string); + reader.readAsDataURL(this.file); + } + } else if (!this.file && this.containerEl.classList.contains("-has-file")) { + this.containerEl.classList.remove("-has-file"); + this.outputImageEl.src = ""; + this.outputImageEl.removeAttribute("src"); + } + + if (this.graphResults) { + this.containerEl.classList.add("-has-results"); + if (!this.graphResults.patched && !this.graphResults.deleted) { + this.outputeMessageEl.innerHTML = "โœ… No bad links detected in the workflow."; + } else { + this.containerEl.classList.add("-has-fixable-results"); + this.outputeMessageEl.innerHTML = `โš ๏ธ Found ${this.graphResults.patched} links to fix, and ${this.graphResults.deleted} to be removed.`; + } + } else { + this.containerEl.classList.remove("-has-results"); + this.containerEl.classList.remove("-has-fixable-results"); + } + + if (msg) { + this.outputeMessageEl.innerHTML = msg; + } + } + + private async handleFile(file: File | Blob) { + this.file = file; + this.updateUi(); + + let workflow: string | undefined | null = null; + if (file.type.startsWith("image/")) { + const pngInfo = await getPngMetadata(file); + workflow = pngInfo?.workflow; + } else if ( + file.type === "application/json" || + (file instanceof File && file.name.endsWith(".json")) + ) { + workflow = await new Promise((resolve) => { + const reader = new FileReader(); + reader.onload = () => { + resolve(reader.result as string); + }; + reader.readAsText(file); + }); + } + if (!workflow) { + this.updateUi("โ›” No workflow found in dropped item."); + } else { + try { + this.graph = JSON.parse(workflow); + } catch (e) { + this.graph = undefined; + } + if (!this.graph) { + this.updateUi("โ›” Invalid workflow found in dropped item."); + } else { + this.loadGraphData(this.graph); + } + } + } + + private async loadGraphData(graphData: SerializedGraph) { + this.graphResults = await fixBadLinks(graphData); + this.updateUi(); + } + + private async saveFixedWorkflow() { + if (!this.graphFinalResults) { + this.updateUi("โ›” Save w/o final graph patched."); + return false; + } + + let filename: string | null = (this.file as File).name || "workflow.json"; + let filenames = filename.split("."); + filenames.pop(); + filename = filenames.join("."); + filename += "_fixed.json"; + filename = prompt("Save workflow as:", filename); + if (!filename) return false; + if (!filename.toLowerCase().endsWith(".json")) { + filename += ".json"; + } + const json = JSON.stringify(this.graphFinalResults.graph, null, 2); + const blob = new Blob([json], { type: "application/json" }); + const url = URL.createObjectURL(blob); + const anchor = document.createElement("a"); + anchor.download = filename; + anchor.href = url; + anchor.style.display = "none"; + document.body.appendChild(anchor); + await wait(); + anchor.click(); + await wait(); + anchor.remove(); + window.URL.revokeObjectURL(url); + return true; + } +} diff --git a/src_web/scripts_comfy/app.ts b/src_web/scripts_comfy/app.ts index f829090..8ef5f52 100644 --- a/src_web/scripts_comfy/app.ts +++ b/src_web/scripts_comfy/app.ts @@ -1,7 +1,7 @@ -import { ComfyApp } from "../typings/comfy"; - -/** - * A dummy ComfyApp that we can import from our code, which we'll rewrite later to the comfyui - * hosted app.js - */ -export declare const app: ComfyApp; +import { ComfyApp } from "../typings/comfy.js"; + +/** + * A dummy ComfyApp that we can import from our code, which we'll rewrite later to the comfyui + * hosted app.js + */ +export declare const app: ComfyApp; diff --git a/src_web/scripts_comfy/ui/components/button.ts b/src_web/scripts_comfy/ui/components/button.ts index 1287828..865eeba 100644 --- a/src_web/scripts_comfy/ui/components/button.ts +++ b/src_web/scripts_comfy/ui/components/button.ts @@ -1,23 +1,23 @@ -import type { ComfyApp } from "typings/comfy.ts"; - -type ComfyButtonProps = { - icon?: string; - overIcon?: string; - iconSize?: number; - content?: string | HTMLElement; - tooltip?: string; - enabled?: boolean; - action?: (e: Event, btn: ComfyButton) => void; - classList?: string; - visibilitySetting?: { id: string, showValue: any }; - app?: ComfyApp; -} - -export declare class ComfyButton { - element: HTMLElement; - iconElement: HTMLElement; - contentElement: HTMLElement; - constructor(props: ComfyButtonProps); - updateIcon(): void; - withPopup(popup: any, mode: "click"|"hover"): this; -}; +import type { ComfyApp } from "typings/comfy.js"; + +type ComfyButtonProps = { + icon?: string; + overIcon?: string; + iconSize?: number; + content?: string | HTMLElement; + tooltip?: string; + enabled?: boolean; + action?: (e: Event, btn: ComfyButton) => void; + classList?: string; + visibilitySetting?: { id: string, showValue: any }; + app?: ComfyApp; +} + +export declare class ComfyButton { + element: HTMLElement; + iconElement: HTMLElement; + contentElement: HTMLElement; + constructor(props: ComfyButtonProps); + updateIcon(): void; + withPopup(popup: any, mode: "click"|"hover"): this; +}; diff --git a/src_web/scripts_comfy/ui/components/buttonGroup.ts b/src_web/scripts_comfy/ui/components/buttonGroup.ts index 977e90a..af8954a 100644 --- a/src_web/scripts_comfy/ui/components/buttonGroup.ts +++ b/src_web/scripts_comfy/ui/components/buttonGroup.ts @@ -1,10 +1,10 @@ -import type {ComfyButton} from "scripts/ui/components/button.ts"; - -export declare class ComfyButtonGroup { - element: HTMLElement; - constructor(...buttons: Array); - insert(button: ComfyButton, index: number): void; - append(button: ComfyButton): void; - remove(indexOrButton: ComfyButton|number): ComfyButton|HTMLElement|void; - update(): void; -}; +import type {ComfyButton} from "scripts/ui/components/button.js"; + +export declare class ComfyButtonGroup { + element: HTMLElement; + constructor(...buttons: Array); + insert(button: ComfyButton, index: number): void; + append(button: ComfyButton): void; + remove(indexOrButton: ComfyButton|number): ComfyButton|HTMLElement|void; + update(): void; +}; diff --git a/src_web/scripts_comfy/widgets.ts b/src_web/scripts_comfy/widgets.ts index c533302..ee8e9c6 100644 --- a/src_web/scripts_comfy/widgets.ts +++ b/src_web/scripts_comfy/widgets.ts @@ -1,19 +1,19 @@ -import type { LGraphNode } from "typings/litegraph"; -import type { ComfyApp, ComfyWidget } from "../typings/comfy"; - -type ComfyWidgetFn = ( - node: LGraphNode, - inputName: string, - inputData: any, - app: ComfyApp, -) => { widget: ComfyWidget }; - -/** - * A dummy ComfyWidgets that we can import from our code, which we'll rewrite later to the comfyui - * hosted widgets.js - */ -export declare const ComfyWidgets: { - COMBO: ComfyWidgetFn; - STRING: ComfyWidgetFn; - [key: string]: ComfyWidgetFn; -}; +import type { LGraphNode } from "typings/litegraph.js"; +import type { ComfyApp, ComfyWidget } from "../typings/comfy.js"; + +type ComfyWidgetFn = ( + node: LGraphNode, + inputName: string, + inputData: any, + app: ComfyApp, +) => { widget: ComfyWidget }; + +/** + * A dummy ComfyWidgets that we can import from our code, which we'll rewrite later to the comfyui + * hosted widgets.js + */ +export declare const ComfyWidgets: { + COMBO: ComfyWidgetFn; + STRING: ComfyWidgetFn; + [key: string]: ComfyWidgetFn; +}; diff --git a/src_web/typings/comfy.d.ts b/src_web/typings/comfy.d.ts index 2d8a2f2..bbecb80 100644 --- a/src_web/typings/comfy.d.ts +++ b/src_web/typings/comfy.d.ts @@ -1,227 +1,227 @@ -import type { LGraphGroup as TLGraphGroup, LGraphNode as TLGraphNode, IWidget, SerializedLGraphNode, LGraph as TLGraph, LGraphCanvas as TLGraphCanvas, LiteGraph as TLiteGraph } from "./litegraph"; -import type {Constructor, SerializedGraph} from './index'; - -declare global { - const LiteGraph: typeof TLiteGraph; - const LGraph: typeof TLGraph; - const LGraphNode: typeof TLGraphNode; - const LGraphCanvas: typeof TLGraphCanvas; - const LGraphGroup: typeof TLGraphGroup; -} - -// @rgthree: Types on ComfyApp as needed. -export interface ComfyApp { - extensions: ComfyExtension[]; - async queuePrompt(number?: number, batchCount = 1): Promise; - graph: TLGraph; - canvas: TLGraphCanvas; - clean() : void; - registerExtension(extension: ComfyExtension): void; - getPreviewFormatParam(): string; - getRandParam(): string; - loadApiJson(apiData: {}, fileName: string): void; - async graphToPrompt(graph?: TLGraph, clean?: boolean): Promise; - // workflow: ComfyWorkflowInstance ??? - async loadGraphData(graphData: {}, clean?: boolean, restore_view?: boolean, workflow?: any|null): Promise - ui: { - settings: { - addSetting(config: {id: string, name: string, type: () => HTMLElement}) : void; - } - } - // Just marking as any for now. - menu?: any; -} - -export interface ComfyWidget extends IWidget { - // https://github.com/comfyanonymous/ComfyUI/issues/2193 Changes from SerializedLGraphNode to - // LGraphNode... - serializeValue(nodeType: TLGraphNode, index: number): Promise; - afterQueued(): void; - inputEl?: HTMLTextAreaElement; - width: number; -} - -export interface ComfyGraphNode extends TLGraphNode { - getExtraMenuOptions: (node: TLGraphNode, options: ContextMenuItem[]) => void; - onExecuted(message: any): void; -} - -export interface ComfyNode extends TLGraphNode { - comfyClass: string; -} - -// @rgthree -export interface ComfyNodeConstructor extends Constructor { - static title: string; - static type?: string; - static comfyClass: string; -} - -export type NodeMode = 0|1|2|3|4|undefined; - - -export interface ComfyExtension { - /** - * The name of the extension - */ - name: string; - /** - * Allows any initialisation, e.g. loading resources. Called after the canvas is created but before nodes are added - * @param app The ComfyUI app instance - */ - init?(app: ComfyApp): Promise; - /** - * Allows any additonal setup, called after the application is fully set up and running - * @param app The ComfyUI app instance - */ - setup?(app: ComfyApp): Promise; - /** - * Called before nodes are registered with the graph - * @param defs The collection of node definitions, add custom ones or edit existing ones - * @param app The ComfyUI app instance - */ - addCustomNodeDefs?(defs: Record, app: ComfyApp): Promise; - /** - * Allows the extension to add custom widgets - * @param app The ComfyUI app instance - * @returns An array of {[widget name]: widget data} - */ - getCustomWidgets?( - app: ComfyApp - ): Promise< - Record { widget?: IWidget; minWidth?: number; minHeight?: number }> - >; - /** - * Allows the extension to add additional handling to the node before it is registered with LGraph - * @rgthree changed nodeType from `typeof LGraphNode` to `ComfyNodeConstructor` - * @param nodeType The node class (not an instance) - * @param nodeData The original node object info config object - * @param app The ComfyUI app instance - */ - beforeRegisterNodeDef?(nodeType: ComfyNodeConstructor, nodeData: ComfyObjectInfo, app: ComfyApp): Promise; - /** - * Allows the extension to register additional nodes with LGraph after standard nodes are added - * @param app The ComfyUI app instance - */ - // @rgthree - add void for non async - registerCustomNodes?(app: ComfyApp): void|Promise; - /** - * Allows the extension to modify a node that has been reloaded onto the graph. - * If you break something in the backend and want to patch workflows in the frontend - * This is the place to do this - * @param node The node that has been loaded - * @param app The ComfyUI app instance - */ - loadedGraphNode?(node: TLGraphNode, app: ComfyApp); - /** - * Allows the extension to run code after the constructor of the node - * @param node The node that has been created - * @param app The ComfyUI app instance - */ - nodeCreated?(node: TLGraphNode, app: ComfyApp); -} - -export type ComfyObjectInfo = { - name: string; - display_name?: string; - description?: string; - category: string; - input?: { - required?: Record; - optional?: Record; - hidden?: Record; - }; - output?: string[]; - output_name: string[]; - // @rgthree - output_node?: boolean; -}; - -export type ComfyObjectInfoConfig = [string | any[]] | [string | any[], any]; - -// @rgthree -type ComfyApiInputLink = [ - /** The id string of the connected node. */ - string, - /** The output index. */ - number, -] - -// @rgthree -export type ComfyApiFormatNode = { - "inputs": { - [input_name: string]: string|number|boolean|ComfyApiInputLink, - }, - "class_type": string, - "_meta": { - "title": string, - } -} - -// @rgthree -export type ComfyApiFormat = { - [node_id: string]: ComfyApiFormatNode -} - -// @rgthree -export type ComfyApiPrompt = { - workflow: SerializedGraph, - output: ComfyApiFormat, -} - -// @rgthree -export type ComfyApiEventDetailStatus = { - exec_info: { - queue_remaining: number; - }; -}; - -// @rgthree -export type ComfyApiEventDetailExecutionStart = { - prompt_id: string; -}; - -// @rgthree -export type ComfyApiEventDetailExecuting = null | string; - -// @rgthree -export type ComfyApiEventDetailProgress = { - node: string; - prompt_id: string; - max: number; - value: number; -}; - -// @rgthree -export type ComfyApiEventDetailExecuted = { - node: string; - prompt_id: string; - output: any; -}; - -// @rgthree -export type ComfyApiEventDetailCached = { - nodes: string[]; - prompt_id: string; -}; - -// @rgthree -export type ComfyApiEventDetailExecuted = { - prompt_id: string; - node: string; - output: any; -}; - -// @rgthree -export type ComfyApiEventDetailError = { - prompt_id: string; - exception_type: string; - exception_message: string; - node_id: string; - node_type: string; - node_id: string; - traceback: string; - executed: any[]; - current_inputs: {[key: string]: (number[]|string[])}; - current_outputs: {[key: string]: (number[]|string[])}; -} +import type { LGraphGroup as TLGraphGroup, LGraphNode as TLGraphNode, IWidget, SerializedLGraphNode, LGraph as TLGraph, LGraphCanvas as TLGraphCanvas, LiteGraph as TLiteGraph } from "./litegraph.js"; +import type {Constructor, SerializedGraph} from './index.js'; + +declare global { + const LiteGraph: typeof TLiteGraph; + const LGraph: typeof TLGraph; + const LGraphNode: typeof TLGraphNode; + const LGraphCanvas: typeof TLGraphCanvas; + const LGraphGroup: typeof TLGraphGroup; +} + +// @rgthree: Types on ComfyApp as needed. +export interface ComfyApp { + extensions: ComfyExtension[]; + async queuePrompt(number?: number, batchCount = 1): Promise; + graph: TLGraph; + canvas: TLGraphCanvas; + clean() : void; + registerExtension(extension: ComfyExtension): void; + getPreviewFormatParam(): string; + getRandParam(): string; + loadApiJson(apiData: {}, fileName: string): void; + async graphToPrompt(graph?: TLGraph, clean?: boolean): Promise; + // workflow: ComfyWorkflowInstance ??? + async loadGraphData(graphData: {}, clean?: boolean, restore_view?: boolean, workflow?: any|null): Promise + ui: { + settings: { + addSetting(config: {id: string, name: string, type: () => HTMLElement}) : void; + } + } + // Just marking as any for now. + menu?: any; +} + +export interface ComfyWidget extends IWidget { + // https://github.com/comfyanonymous/ComfyUI/issues/2193 Changes from SerializedLGraphNode to + // LGraphNode... + serializeValue(nodeType: TLGraphNode, index: number): Promise; + afterQueued(): void; + inputEl?: HTMLTextAreaElement; + width: number; +} + +export interface ComfyGraphNode extends TLGraphNode { + getExtraMenuOptions: (node: TLGraphNode, options: ContextMenuItem[]) => void; + onExecuted(message: any): void; +} + +export interface ComfyNode extends TLGraphNode { + comfyClass: string; +} + +// @rgthree +export interface ComfyNodeConstructor extends Constructor { + static title: string; + static type?: string; + static comfyClass: string; +} + +export type NodeMode = 0|1|2|3|4|undefined; + + +export interface ComfyExtension { + /** + * The name of the extension + */ + name: string; + /** + * Allows any initialisation, e.g. loading resources. Called after the canvas is created but before nodes are added + * @param app The ComfyUI app instance + */ + init?(app: ComfyApp): Promise; + /** + * Allows any additonal setup, called after the application is fully set up and running + * @param app The ComfyUI app instance + */ + setup?(app: ComfyApp): Promise; + /** + * Called before nodes are registered with the graph + * @param defs The collection of node definitions, add custom ones or edit existing ones + * @param app The ComfyUI app instance + */ + addCustomNodeDefs?(defs: Record, app: ComfyApp): Promise; + /** + * Allows the extension to add custom widgets + * @param app The ComfyUI app instance + * @returns An array of {[widget name]: widget data} + */ + getCustomWidgets?( + app: ComfyApp + ): Promise< + Record { widget?: IWidget; minWidth?: number; minHeight?: number }> + >; + /** + * Allows the extension to add additional handling to the node before it is registered with LGraph + * @rgthree changed nodeType from `typeof LGraphNode` to `ComfyNodeConstructor` + * @param nodeType The node class (not an instance) + * @param nodeData The original node object info config object + * @param app The ComfyUI app instance + */ + beforeRegisterNodeDef?(nodeType: ComfyNodeConstructor, nodeData: ComfyObjectInfo, app: ComfyApp): Promise; + /** + * Allows the extension to register additional nodes with LGraph after standard nodes are added + * @param app The ComfyUI app instance + */ + // @rgthree - add void for non async + registerCustomNodes?(app: ComfyApp): void|Promise; + /** + * Allows the extension to modify a node that has been reloaded onto the graph. + * If you break something in the backend and want to patch workflows in the frontend + * This is the place to do this + * @param node The node that has been loaded + * @param app The ComfyUI app instance + */ + loadedGraphNode?(node: TLGraphNode, app: ComfyApp); + /** + * Allows the extension to run code after the constructor of the node + * @param node The node that has been created + * @param app The ComfyUI app instance + */ + nodeCreated?(node: TLGraphNode, app: ComfyApp); +} + +export type ComfyObjectInfo = { + name: string; + display_name?: string; + description?: string; + category: string; + input?: { + required?: Record; + optional?: Record; + hidden?: Record; + }; + output?: string[]; + output_name: string[]; + // @rgthree + output_node?: boolean; +}; + +export type ComfyObjectInfoConfig = [string | any[]] | [string | any[], any]; + +// @rgthree +type ComfyApiInputLink = [ + /** The id string of the connected node. */ + string, + /** The output index. */ + number, +] + +// @rgthree +export type ComfyApiFormatNode = { + "inputs": { + [input_name: string]: string|number|boolean|ComfyApiInputLink, + }, + "class_type": string, + "_meta": { + "title": string, + } +} + +// @rgthree +export type ComfyApiFormat = { + [node_id: string]: ComfyApiFormatNode +} + +// @rgthree +export type ComfyApiPrompt = { + workflow: SerializedGraph, + output: ComfyApiFormat, +} + +// @rgthree +export type ComfyApiEventDetailStatus = { + exec_info: { + queue_remaining: number; + }; +}; + +// @rgthree +export type ComfyApiEventDetailExecutionStart = { + prompt_id: string; +}; + +// @rgthree +export type ComfyApiEventDetailExecuting = null | string; + +// @rgthree +export type ComfyApiEventDetailProgress = { + node: string; + prompt_id: string; + max: number; + value: number; +}; + +// @rgthree +export type ComfyApiEventDetailExecuted = { + node: string; + prompt_id: string; + output: any; +}; + +// @rgthree +export type ComfyApiEventDetailCached = { + nodes: string[]; + prompt_id: string; +}; + +// @rgthree +export type ComfyApiEventDetailExecuted = { + prompt_id: string; + node: string; + output: any; +}; + +// @rgthree +export type ComfyApiEventDetailError = { + prompt_id: string; + exception_type: string; + exception_message: string; + node_id: string; + node_type: string; + node_id: string; + traceback: string; + executed: any[]; + current_inputs: {[key: string]: (number[]|string[])}; + current_outputs: {[key: string]: (number[]|string[])}; +} diff --git a/src_web/typings/index.d.ts b/src_web/typings/index.d.ts index 9a3cc79..a574726 100644 --- a/src_web/typings/index.d.ts +++ b/src_web/typings/index.d.ts @@ -1,55 +1,55 @@ -import { LGraph } from "litegraph"; - -export type Constructor = new(...args: any[]) => T; - -export type SerializedLink = [ - number, // this.id, - number, // this.origin_id, - number, // this.origin_slot, - number, // this.target_id, - number, // this.target_slot, - string, // this.type -]; - -export interface SerializedNodeInput { - name: string; - type: string; - link: number; -} -export interface SerializedNodeOutput { - name: string; - type: string; - link: number; - slot_index: number; - links: number[]; -} -export interface SerializedNode { - id: number; - inputs: SerializedNodeInput[]; - outputs: SerializedNodeOutput[]; - mode: number; - order: number; - pos: [number, number]; - properties: any; - size: [number, number]; - type: string; - widgets_values: Array; -} - -export interface SerializedGraph { - config: any; - extra: any; - groups: any; - last_link_id: number; - last_node_id: number; - links: SerializedLink[]; - nodes: SerializedNode[]; -} - -export interface BadLinksData { - hasBadLinks: boolean; - fixed: boolean; - graph: T; - patched: number; - deleted: number; -} +import { LGraph } from "./litegraph.js"; + +export type Constructor = new(...args: any[]) => T; + +export type SerializedLink = [ + number, // this.id, + number, // this.origin_id, + number, // this.origin_slot, + number, // this.target_id, + number, // this.target_slot, + string, // this.type +]; + +export interface SerializedNodeInput { + name: string; + type: string; + link: number; +} +export interface SerializedNodeOutput { + name: string; + type: string; + link: number; + slot_index: number; + links: number[]; +} +export interface SerializedNode { + id: number; + inputs: SerializedNodeInput[]; + outputs: SerializedNodeOutput[]; + mode: number; + order: number; + pos: [number, number]; + properties: any; + size: [number, number]; + type: string; + widgets_values: Array; +} + +export interface SerializedGraph { + config: any; + extra: any; + groups: any; + last_link_id: number; + last_node_id: number; + links: SerializedLink[]; + nodes: SerializedNode[]; +} + +export interface BadLinksData { + hasBadLinks: boolean; + fixed: boolean; + graph: T; + patched: number; + deleted: number; +} diff --git a/src_web/typings/litegraph.d.ts b/src_web/typings/litegraph.d.ts index c0a94a2..6e09657 100644 --- a/src_web/typings/litegraph.d.ts +++ b/src_web/typings/litegraph.d.ts @@ -40,6 +40,8 @@ export interface INodeSlot { disabled?: boolean; // @rgthree - Found this checked in getSlotMenuOptions default. removable?: boolean; + // @rgthree - A status we put on some nodes so we can draw things around it. + rgthree_status?: 'WARN' | 'ERROR'; } export interface INodeInputSlot extends INodeSlot { @@ -1172,7 +1174,8 @@ export declare class LGraphNode { slotIndex: number, isConnected: boolean, link: LLink, - ioSlot: (INodeOutputSlot | INodeInputSlot) + // @rgthree - Make it INodeSlot instead of union + ioSlot: INodeSlot ): void; /** @@ -1269,6 +1272,13 @@ export declare class DragAndScale { reset(): void; } +// @rgthree. +interface CanvasDivDialog extends HTMLDivElement { + close: () => void; + modified: () => void; + is_modified: boolean; +} + /** * This class is in charge of rendering one graph inside a canvas. And provides all the interaction required. * Valid callbacks are: onNodeSelected, onNodeDeselected, onShowNodePanel, onNodeDblClicked @@ -1662,7 +1672,10 @@ export declare class LGraphCanvas { createDialog( html: string, options?: { position?: Vector2; event?: MouseEvent } - ): void; + // @rgthree - Fix return type from void (added above) + ): CanvasDivDialog; + + convertOffsetToCanvas: DragAndScale["convertOffsetToCanvas"]; convertCanvasToOffset: DragAndScale["convertCanvasToOffset"]; diff --git a/src_web/typings/rgthree.d.ts b/src_web/typings/rgthree.d.ts index 03e247a..bf11fa2 100644 --- a/src_web/typings/rgthree.d.ts +++ b/src_web/typings/rgthree.d.ts @@ -1,67 +1,67 @@ -import type { AdjustedMouseEvent, LGraphNode, Vector2 } from "./litegraph"; -import type {Constructor} from "./index"; -import type {RgthreeBaseVirtualNode} from '../comfyui/base_node.js' - -export type AdjustedMouseCustomEvent = CustomEvent<{ originalEvent: AdjustedMouseEvent }>; - - -export interface RgthreeBaseNodeConstructor extends Constructor { - static type: string; - static category: string; - static comfyClass: string; - static exposedActions: string[]; -} - -export interface RgthreeBaseVirtualNodeConstructor extends Constructor { - static type: string; - static category: string; - static _category: string; -} - - -export interface RgthreeBaseServerNodeConstructor extends Constructor { - static nodeType: ComfyNodeConstructor; - static nodeData: ComfyObjectInfo; - static __registeredForOverride__: boolean; - onRegisteredForOverride(comfyClass: any, rgthreeClass: any) : void; -} - - -export type RgthreeModelInfo = { - file?: string; - name?: string; - type?: string; - baseModel?: string; - baseModelFile?: string; - links?: string[]; - strengthMin?: number; - strengthMax?: number; - triggerWords?: string[]; - trainedWords?: { - word: string; - count?: number; - civitai?: boolean - user?: boolean - }[]; - description?: string; - sha256?: string; - path?: string; - images?: { - url: string; - civitaiUrl?: string; - steps?: string|number; - cfg?: string|number; - type?: 'image'|'video'; - sampler?: string; - model?: string; - seed?: string; - negative?: string; - positive?: string; - resources?: {name?: string, type?: string, weight?: string|number}[]; - }[] - userTags?: string[]; - userNote?: string; - raw?: any; - // This one is just on the client. - filterDir?: string; -} +import type { AdjustedMouseEvent, LGraphNode, Vector2 } from "./litegraph.js"; +import type {Constructor} from "./index.js"; +import type {RgthreeBaseVirtualNode} from '../comfyui/base_node.js' + +export type AdjustedMouseCustomEvent = CustomEvent<{ originalEvent: AdjustedMouseEvent }>; + + +export interface RgthreeBaseNodeConstructor extends Constructor { + static type: string; + static category: string; + static comfyClass: string; + static exposedActions: string[]; +} + +export interface RgthreeBaseVirtualNodeConstructor extends Constructor { + static type: string; + static category: string; + static _category: string; +} + + +export interface RgthreeBaseServerNodeConstructor extends Constructor { + static nodeType: ComfyNodeConstructor; + static nodeData: ComfyObjectInfo; + static __registeredForOverride__: boolean; + onRegisteredForOverride(comfyClass: any, rgthreeClass: any) : void; +} + + +export type RgthreeModelInfo = { + file?: string; + name?: string; + type?: string; + baseModel?: string; + baseModelFile?: string; + links?: string[]; + strengthMin?: number; + strengthMax?: number; + triggerWords?: string[]; + trainedWords?: { + word: string; + count?: number; + civitai?: boolean + user?: boolean + }[]; + description?: string; + sha256?: string; + path?: string; + images?: { + url: string; + civitaiUrl?: string; + steps?: string|number; + cfg?: string|number; + type?: 'image'|'video'; + sampler?: string; + model?: string; + seed?: string; + negative?: string; + positive?: string; + resources?: {name?: string, type?: string, weight?: string|number}[]; + }[] + userTags?: string[]; + userNote?: string; + raw?: any; + // This one is just on the client. + filterDir?: string; +} diff --git a/web/comfyui/constants.js b/web/comfyui/constants.js index 4332e77..cc506af 100644 --- a/web/comfyui/constants.js +++ b/web/comfyui/constants.js @@ -1,3 +1,4 @@ +import { SERVICE as CONFIG_SERVICE } from "./services/config_service.js"; export function addRgthree(str) { return str + " (rgthree)"; } @@ -12,6 +13,8 @@ export const NodeTypesString = { CONTEXT_SWITCH_BIG: addRgthree("Context Switch Big"), CONTEXT_MERGE: addRgthree("Context Merge"), CONTEXT_MERGE_BIG: addRgthree("Context Merge Big"), + DYNAMIC_CONTEXT: addRgthree("Dynamic Context"), + DYNAMIC_CONTEXT_SWITCH: addRgthree("Dynamic Context Switch"), DISPLAY_ANY: addRgthree("Display Any"), NODE_MODE_RELAY: addRgthree("Mute / Bypass Relay"), NODE_MODE_REPEATER: addRgthree("Mute / Bypass Repeater"), @@ -39,5 +42,12 @@ export const NodeTypesString = { export function getNodeTypeStrings() { return Object.values(NodeTypesString) .map((i) => stripRgthree(i)) + .filter((i) => { + if (i.startsWith("Dynamic Context") && + !CONFIG_SERVICE.getConfigValue("unreleased.dynamic_context.enabled")) { + return false; + } + return true; + }) .sort(); } diff --git a/web/comfyui/context.js b/web/comfyui/context.js index a651bec..3df222e 100644 --- a/web/comfyui/context.js +++ b/web/comfyui/context.js @@ -50,7 +50,7 @@ function findMatchingIndexByTypeOrName(otherNode, otherSlot, ctxSlots) { } return ctxSlotIndex; } -class BaseContextNode extends RgthreeBaseServerNode { +export class BaseContextNode extends RgthreeBaseServerNode { constructor(title) { super(title); this.___collapsed_width = 0; diff --git a/web/comfyui/dynamic_context.js b/web/comfyui/dynamic_context.js new file mode 100644 index 0000000..5a2bc16 --- /dev/null +++ b/web/comfyui/dynamic_context.js @@ -0,0 +1,253 @@ +import { app } from "../../scripts/app.js"; +import { IoDirection, followConnectionUntilType, getConnectedInputInfosAndFilterPassThroughs, } from "./utils.js"; +import { rgthree } from "./rgthree.js"; +import { SERVICE as CONTEXT_SERVICE, InputMutationOperation, } from "./services/context_service.js"; +import { NodeTypesString } from "./constants.js"; +import { removeUnusedInputsFromEnd } from "./utils_inputs_outputs.js"; +import { DynamicContextNodeBase } from "./dynamic_context_base.js"; +import { SERVICE as CONFIG_SERVICE } from "./services/config_service.js"; +const OWNED_PREFIX = "+"; +const REGEX_OWNED_PREFIX = /^\+\s*/; +const REGEX_EMPTY_INPUT = /^\+\s*$/; +export class DynamicContextNode extends DynamicContextNodeBase { + constructor(title = DynamicContextNode.title) { + super(title); + } + onNodeCreated() { + this.addInput("base_ctx", "RGTHREE_DYNAMIC_CONTEXT"); + this.ensureOneRemainingNewInputSlot(); + super.onNodeCreated(); + } + onConnectionsChange(type, slotIndex, isConnected, link, ioSlot) { + var _a; + (_a = super.onConnectionsChange) === null || _a === void 0 ? void 0 : _a.call(this, type, slotIndex, isConnected, link, ioSlot); + if (this.configuring) { + return; + } + if (type === LiteGraph.INPUT) { + if (isConnected) { + this.handleInputConnected(slotIndex); + } + else { + this.handleInputDisconnected(slotIndex); + } + } + } + onConnectInput(inputIndex, outputType, outputSlot, outputNode, outputIndex) { + var _a; + let canConnect = true; + if (super.onConnectInput) { + canConnect = super.onConnectInput.apply(this, [...arguments]); + } + if (canConnect && + outputNode instanceof DynamicContextNode && + outputIndex === 0 && + inputIndex !== 0) { + const [n, v] = rgthree.logger.warnParts("Currently, you can only connect a context node in the first slot."); + (_a = console[n]) === null || _a === void 0 ? void 0 : _a.call(console, ...v); + canConnect = false; + } + return canConnect; + } + handleInputConnected(slotIndex) { + const ioSlot = this.inputs[slotIndex]; + const connectedIndexes = []; + if (slotIndex === 0) { + let baseNodeInfos = getConnectedInputInfosAndFilterPassThroughs(this, this, 0); + const baseNodes = baseNodeInfos.map((n) => n.node); + const baseNodesDynamicCtx = baseNodes[0]; + if (baseNodesDynamicCtx === null || baseNodesDynamicCtx === void 0 ? void 0 : baseNodesDynamicCtx.provideInputsData) { + const inputsData = CONTEXT_SERVICE.getDynamicContextInputsData(baseNodesDynamicCtx); + console.log("inputsData", inputsData); + for (const input of baseNodesDynamicCtx.provideInputsData()) { + if (input.name === "base_ctx" || input.name === "+") { + continue; + } + this.addContextInput(input.name, input.type, input.index); + this.stabilizeNames(); + } + } + } + else if (this.isInputSlotForNewInput(slotIndex)) { + this.handleNewInputConnected(slotIndex); + } + } + isInputSlotForNewInput(slotIndex) { + const ioSlot = this.inputs[slotIndex]; + return ioSlot && ioSlot.name === "+" && ioSlot.type === "*"; + } + handleNewInputConnected(slotIndex) { + if (!this.isInputSlotForNewInput(slotIndex)) { + throw new Error('Expected the incoming slot index to be the "new input" input.'); + } + const ioSlot = this.inputs[slotIndex]; + let cxn = null; + if (ioSlot.link != null) { + cxn = followConnectionUntilType(this, IoDirection.INPUT, slotIndex, true); + } + if ((cxn === null || cxn === void 0 ? void 0 : cxn.type) && (cxn === null || cxn === void 0 ? void 0 : cxn.name)) { + let name = this.addOwnedPrefix(this.getNextUniqueNameForThisNode(cxn.name)); + if (name.match(/^\+\s*[A-Z_]+(\.\d+)?$/)) { + name = name.toLowerCase(); + } + ioSlot.name = name; + ioSlot.type = cxn.type; + ioSlot.removable = true; + while (!this.outputs[slotIndex]) { + this.addOutput("*", "*"); + } + this.outputs[slotIndex].type = cxn.type; + this.outputs[slotIndex].name = this.stripOwnedPrefix(name).toLocaleUpperCase(); + if (cxn.type === "COMBO" || cxn.type.includes(",") || Array.isArray(cxn.type)) { + this.outputs[slotIndex].widget = true; + } + this.inputsMutated({ + operation: InputMutationOperation.ADDED, + node: this, + slotIndex, + slot: ioSlot, + }); + this.stabilizeNames(); + this.ensureOneRemainingNewInputSlot(); + } + } + handleInputDisconnected(slotIndex) { + var _a, _b; + const inputs = this.getContextInputsList(); + if (slotIndex === 0) { + for (let index = inputs.length - 1; index > 0; index--) { + if (index === 0 || index === inputs.length - 1) { + continue; + } + const input = inputs[index]; + if (!this.isOwnedInput(input.name)) { + if (input.link || ((_b = (_a = this.outputs[index]) === null || _a === void 0 ? void 0 : _a.links) === null || _b === void 0 ? void 0 : _b.length)) { + this.renameContextInput(index, input.name, true); + } + else { + this.removeContextInput(index); + } + } + } + this.setSize(this.computeSize()); + this.setDirtyCanvas(true, true); + } + } + ensureOneRemainingNewInputSlot() { + removeUnusedInputsFromEnd(this, 1, REGEX_EMPTY_INPUT); + this.addInput(OWNED_PREFIX, "*"); + } + getNextUniqueNameForThisNode(desiredName) { + const inputs = this.getContextInputsList(); + const allExistingKeys = inputs.map((i) => this.stripOwnedPrefix(i.name).toLocaleUpperCase()); + desiredName = this.stripOwnedPrefix(desiredName); + let newName = desiredName; + let n = 0; + while (allExistingKeys.includes(newName.toLocaleUpperCase())) { + newName = `${desiredName}.${++n}`; + } + return newName; + } + removeInput(slotIndex) { + const slot = this.inputs[slotIndex]; + super.removeInput(slotIndex); + if (this.outputs[slotIndex]) { + this.removeOutput(slotIndex); + } + this.inputsMutated({ operation: InputMutationOperation.REMOVED, node: this, slotIndex, slot }); + this.stabilizeNames(); + } + stabilizeNames() { + const inputs = this.getContextInputsList(); + const names = []; + for (const [index, input] of inputs.entries()) { + if (index === 0 || index === inputs.length - 1) { + continue; + } + input.label = undefined; + this.outputs[index].label = undefined; + let origName = this.stripOwnedPrefix(input.name).replace(/\.\d+$/, ""); + let name = input.name; + if (!this.isOwnedInput(name)) { + names.push(name.toLocaleUpperCase()); + } + else { + let n = 0; + name = this.addOwnedPrefix(origName); + while (names.includes(this.stripOwnedPrefix(name).toLocaleUpperCase())) { + name = `${this.addOwnedPrefix(origName)}.${++n}`; + } + names.push(this.stripOwnedPrefix(name).toLocaleUpperCase()); + if (input.name !== name) { + this.renameContextInput(index, name); + } + } + } + } + getSlotMenuOptions(slot) { + const editable = this.isOwnedInput(slot.input.name) && this.type !== "*"; + return [ + { + content: "โœ๏ธ Rename Input", + disabled: !editable, + callback: () => { + var dialog = app.canvas.createDialog("Name", {}); + var dialogInput = dialog.querySelector("input"); + if (dialogInput) { + dialogInput.value = this.stripOwnedPrefix(slot.input.name || ""); + } + var inner = () => { + this.handleContextMenuRenameInputDialog(slot.slot, dialogInput.value); + dialog.close(); + }; + dialog.querySelector("button").addEventListener("click", inner); + dialogInput.addEventListener("keydown", (e) => { + var _a; + dialog.is_modified = true; + if (e.keyCode == 27) { + dialog.close(); + } + else if (e.keyCode == 13) { + inner(); + } + else if (e.keyCode != 13 && ((_a = e.target) === null || _a === void 0 ? void 0 : _a.localName) != "textarea") { + return; + } + e.preventDefault(); + e.stopPropagation(); + }); + dialogInput.focus(); + }, + }, + { + content: "๐Ÿ—‘๏ธ Delete Input", + disabled: !editable, + callback: () => { + this.removeInput(slot.slot); + }, + }, + ]; + } + handleContextMenuRenameInputDialog(slotIndex, value) { + app.graph.beforeChange(); + this.renameContextInput(slotIndex, value); + this.stabilizeNames(); + this.setDirtyCanvas(true, true); + app.graph.afterChange(); + } +} +DynamicContextNode.title = NodeTypesString.DYNAMIC_CONTEXT; +DynamicContextNode.type = NodeTypesString.DYNAMIC_CONTEXT; +DynamicContextNode.comfyClass = NodeTypesString.DYNAMIC_CONTEXT; +const contextDynamicNodes = [DynamicContextNode]; +app.registerExtension({ + name: "rgthree.DynamicContext", + async beforeRegisterNodeDef(nodeType, nodeData) { + if (!CONFIG_SERVICE.getConfigValue("unreleased.dynamic_context.enabled")) { + return; + } + if (nodeData.name === DynamicContextNode.type) { + DynamicContextNode.setUp(nodeType, nodeData); + } + }, +}); diff --git a/web/comfyui/dynamic_context_base.js b/web/comfyui/dynamic_context_base.js new file mode 100644 index 0000000..864aca9 --- /dev/null +++ b/web/comfyui/dynamic_context_base.js @@ -0,0 +1,189 @@ +import { BaseContextNode } from "./context.js"; +import { RgthreeBaseServerNode } from "./base_node.js"; +import { moveArrayItem, wait } from "../../rgthree/common/shared_utils.js"; +import { RgthreeInvisibleWidget } from "./utils_widgets.js"; +import { getContextOutputName, InputMutationOperation, } from "./services/context_service.js"; +import { app } from "../../scripts/app.js"; +import { SERVICE as CONTEXT_SERVICE } from "./services/context_service.js"; +const OWNED_PREFIX = "+"; +const REGEX_OWNED_PREFIX = /^\+\s*/; +const REGEX_EMPTY_INPUT = /^\+\s*$/; +export class DynamicContextNodeBase extends BaseContextNode { + constructor() { + super(...arguments); + this.hasShadowInputs = false; + } + getContextInputsList() { + return this.inputs; + } + provideInputsData() { + const inputs = this.getContextInputsList(); + return inputs + .map((input, index) => ({ + name: this.stripOwnedPrefix(input.name), + type: String(input.type), + index, + })) + .filter((i) => i.type !== "*"); + } + addOwnedPrefix(name) { + return `+ ${this.stripOwnedPrefix(name)}`; + } + isOwnedInput(inputOrName) { + const name = typeof inputOrName == "string" ? inputOrName : (inputOrName === null || inputOrName === void 0 ? void 0 : inputOrName.name) || ""; + return REGEX_OWNED_PREFIX.test(name); + } + stripOwnedPrefix(name) { + return name.replace(REGEX_OWNED_PREFIX, ""); + } + handleUpstreamMutation(mutation) { + console.log(`[node ${this.id}] handleUpstreamMutation`, mutation); + if (mutation.operation === InputMutationOperation.ADDED) { + const slot = mutation.slot; + if (!slot) { + throw new Error("Cannot have an ADDED mutation without a provided slot data."); + } + this.addContextInput(this.stripOwnedPrefix(slot.name), slot.type, mutation.slotIndex); + return; + } + if (mutation.operation === InputMutationOperation.REMOVED) { + const slot = mutation.slot; + if (!slot) { + throw new Error("Cannot have an REMOVED mutation without a provided slot data."); + } + this.removeContextInput(mutation.slotIndex); + return; + } + if (mutation.operation === InputMutationOperation.RENAMED) { + const slot = mutation.slot; + if (!slot) { + throw new Error("Cannot have an RENAMED mutation without a provided slot data."); + } + this.renameContextInput(mutation.slotIndex, slot.name); + return; + } + } + clone() { + const cloned = super.clone(); + while (cloned.inputs.length > 1) { + cloned.removeInput(cloned.inputs.length - 1); + } + while (cloned.widgets.length > 1) { + cloned.removeWidget(cloned.widgets.length - 1); + } + while (cloned.outputs.length > 1) { + cloned.removeOutput(cloned.outputs.length - 1); + } + return cloned; + } + onNodeCreated() { + const node = this; + this.addCustomWidget(new RgthreeInvisibleWidget("output_keys", "RGTHREE_DYNAMIC_CONTEXT_OUTPUTS", "", () => { + return (node.outputs || []) + .map((o, i) => i > 0 && o.name) + .filter((n) => n !== false) + .join(","); + })); + } + addContextInput(name, type, slot = -1) { + const inputs = this.getContextInputsList(); + if (this.hasShadowInputs) { + inputs.push({ name, type, link: null }); + } + else { + this.addInput(name, type); + } + if (slot > -1) { + moveArrayItem(inputs, inputs.length - 1, slot); + } + else { + slot = inputs.length - 1; + } + if (type !== "*") { + const output = this.addOutput(getContextOutputName(name), type); + if (type === "COMBO" || String(type).includes(",") || Array.isArray(type)) { + output.widget = true; + } + if (slot > -1) { + moveArrayItem(this.outputs, this.outputs.length - 1, slot); + } + } + this.fixInputsOutputsLinkSlots(); + this.inputsMutated({ + operation: InputMutationOperation.ADDED, + node: this, + slotIndex: slot, + slot: inputs[slot], + }); + } + removeContextInput(slotIndex) { + if (this.hasShadowInputs) { + const inputs = this.getContextInputsList(); + const input = inputs.splice(slotIndex, 1)[0]; + if (this.outputs[slotIndex]) { + this.removeOutput(slotIndex); + } + } + else { + this.removeInput(slotIndex); + } + } + renameContextInput(index, newName, forceOwnBool = null) { + const inputs = this.getContextInputsList(); + const input = inputs[index]; + const oldName = input.name; + newName = this.stripOwnedPrefix(newName.trim() || this.getSlotDefaultInputLabel(index)); + if (forceOwnBool === true || (this.isOwnedInput(oldName) && forceOwnBool !== false)) { + newName = this.addOwnedPrefix(newName); + } + if (oldName !== newName) { + input.name = newName; + input.removable = this.isOwnedInput(newName); + this.outputs[index].name = getContextOutputName(inputs[index].name); + this.inputsMutated({ + node: this, + operation: InputMutationOperation.RENAMED, + slotIndex: index, + slot: input, + }); + } + } + getSlotDefaultInputLabel(slotIndex) { + const inputs = this.getContextInputsList(); + const input = inputs[slotIndex]; + let defaultLabel = this.stripOwnedPrefix(input.name).toLowerCase(); + return defaultLabel.toLocaleLowerCase(); + } + inputsMutated(mutation) { + CONTEXT_SERVICE.onInputChanges(this, mutation); + } + fixInputsOutputsLinkSlots() { + if (!this.hasShadowInputs) { + const inputs = this.getContextInputsList(); + for (let index = inputs.length - 1; index > 0; index--) { + const input = inputs[index]; + if ((input === null || input === void 0 ? void 0 : input.link) != null) { + app.graph.links[input.link].target_slot = index; + } + } + } + const outputs = this.outputs; + for (let index = outputs.length - 1; index > 0; index--) { + const output = outputs[index]; + if (output) { + output.nameLocked = true; + for (const link of output.links || []) { + app.graph.links[link].origin_slot = index; + } + } + } + } + static setUp(comfyClass, nodeData) { + RgthreeBaseServerNode.registerForOverride(comfyClass, nodeData, this); + wait(500).then(() => { + LiteGraph.slot_types_default_out["RGTHREE_DYNAMIC_CONTEXT"] = + LiteGraph.slot_types_default_out["RGTHREE_DYNAMIC_CONTEXT"] || []; + LiteGraph.slot_types_default_out["RGTHREE_DYNAMIC_CONTEXT"].push(comfyClass.comfyClass); + }); + } +} diff --git a/web/comfyui/dynamic_context_switch.js b/web/comfyui/dynamic_context_switch.js new file mode 100644 index 0000000..9e5ebcd --- /dev/null +++ b/web/comfyui/dynamic_context_switch.js @@ -0,0 +1,146 @@ +import { app } from "../../scripts/app.js"; +import { DynamicContextNodeBase } from "./dynamic_context_base.js"; +import { NodeTypesString } from "./constants.js"; +import { SERVICE as CONTEXT_SERVICE, getContextOutputName, } from "./services/context_service.js"; +import { getConnectedInputNodesAndFilterPassThroughs } from "./utils.js"; +import { debounce, moveArrayItem } from "../../rgthree/common/shared_utils.js"; +import { measureText } from "./utils_canvas.js"; +import { SERVICE as CONFIG_SERVICE } from "./services/config_service.js"; +class DynamicContextSwitchNode extends DynamicContextNodeBase { + constructor(title = DynamicContextSwitchNode.title) { + super(title); + this.hasShadowInputs = true; + this.lastInputsList = []; + this.shadowInputs = [ + { name: "base_ctx", type: "RGTHREE_DYNAMIC_CONTEXT", link: null, count: 0 }, + ]; + } + getContextInputsList() { + return this.shadowInputs; + } + handleUpstreamMutation(mutation) { + this.scheduleHardRefresh(); + } + onConnectionsChange(type, slotIndex, isConnected, link, ioSlot) { + var _a; + (_a = super.onConnectionsChange) === null || _a === void 0 ? void 0 : _a.call(this, type, slotIndex, isConnected, link, ioSlot); + if (this.configuring) { + return; + } + if (type === LiteGraph.INPUT) { + this.scheduleHardRefresh(); + } + } + scheduleHardRefresh(ms = 64) { + return debounce(() => { + this.refreshInputsAndOutputs(); + }, ms); + } + onNodeCreated() { + this.addInput("ctx_1", "RGTHREE_DYNAMIC_CONTEXT"); + this.addInput("ctx_2", "RGTHREE_DYNAMIC_CONTEXT"); + this.addInput("ctx_3", "RGTHREE_DYNAMIC_CONTEXT"); + this.addInput("ctx_4", "RGTHREE_DYNAMIC_CONTEXT"); + this.addInput("ctx_5", "RGTHREE_DYNAMIC_CONTEXT"); + super.onNodeCreated(); + } + addContextInput(name, type, slot) { } + refreshInputsAndOutputs() { + var _a; + const inputs = [ + { name: "base_ctx", type: "RGTHREE_DYNAMIC_CONTEXT", link: null, count: 0 }, + ]; + let numConnected = 0; + for (let i = 0; i < this.inputs.length; i++) { + const childCtxs = getConnectedInputNodesAndFilterPassThroughs(this, this, i); + if (childCtxs.length > 1) { + throw new Error("How is there more than one input?"); + } + const ctx = childCtxs[0]; + if (!ctx) + continue; + numConnected++; + const slotsData = CONTEXT_SERVICE.getDynamicContextInputsData(ctx); + console.log(slotsData); + for (const slotData of slotsData) { + const found = inputs.find((n) => getContextOutputName(slotData.name) === getContextOutputName(n.name)); + if (found) { + found.count += 1; + continue; + } + inputs.push({ + name: slotData.name, + type: slotData.type, + link: null, + count: 1, + }); + } + } + this.shadowInputs = inputs; + let i = 0; + for (i; i < this.shadowInputs.length; i++) { + const data = this.shadowInputs[i]; + let existing = this.outputs.find((o) => getContextOutputName(o.name) === getContextOutputName(data.name)); + if (!existing) { + existing = this.addOutput(getContextOutputName(data.name), data.type); + } + moveArrayItem(this.outputs, existing, i); + delete existing.rgthree_status; + if (data.count !== numConnected) { + existing.rgthree_status = "WARN"; + } + } + while (this.outputs[i]) { + const output = this.outputs[i]; + if ((_a = output === null || output === void 0 ? void 0 : output.links) === null || _a === void 0 ? void 0 : _a.length) { + output.rgthree_status = "ERROR"; + i++; + } + else { + this.removeOutput(i); + } + } + this.fixInputsOutputsLinkSlots(); + } + onDrawForeground(ctx, canvas) { + var _a, _b; + const low_quality = ((_b = (_a = canvas === null || canvas === void 0 ? void 0 : canvas.ds) === null || _a === void 0 ? void 0 : _a.scale) !== null && _b !== void 0 ? _b : 1) < 0.6; + if (low_quality || this.size[0] <= 10) { + return; + } + let y = LiteGraph.NODE_SLOT_HEIGHT - 1; + const w = this.size[0]; + ctx.save(); + ctx.font = "normal " + LiteGraph.NODE_SUBTEXT_SIZE + "px Arial"; + ctx.textAlign = "right"; + for (const output of this.outputs) { + if (!output.rgthree_status) { + y += LiteGraph.NODE_SLOT_HEIGHT; + continue; + } + const x = w - 20 - measureText(ctx, output.name); + if (output.rgthree_status === "ERROR") { + ctx.fillText("๐Ÿ›‘", x, y); + } + else if (output.rgthree_status === "WARN") { + ctx.fillText("โš ๏ธ", x, y); + } + y += LiteGraph.NODE_SLOT_HEIGHT; + } + ctx.restore(); + } +} +DynamicContextSwitchNode.title = NodeTypesString.DYNAMIC_CONTEXT_SWITCH; +DynamicContextSwitchNode.type = NodeTypesString.DYNAMIC_CONTEXT_SWITCH; +DynamicContextSwitchNode.comfyClass = NodeTypesString.DYNAMIC_CONTEXT_SWITCH; +app.registerExtension({ + name: "rgthree.DynamicContextSwitch", + async beforeRegisterNodeDef(nodeType, nodeData) { + if (!CONFIG_SERVICE.getConfigValue("unreleased.dynamic_context.enabled")) { + return; + } + if (nodeData.name === DynamicContextSwitchNode.type) { + DynamicContextSwitchNode.setUp(nodeType, nodeData); + } + }, +}); diff --git a/web/comfyui/services/context_service.js b/web/comfyui/services/context_service.js new file mode 100644 index 0000000..27e1965 --- /dev/null +++ b/web/comfyui/services/context_service.js @@ -0,0 +1,51 @@ +import { getConnectedOutputNodesAndFilterPassThroughs } from "../utils.js"; +export let SERVICE; +const OWNED_PREFIX = "+"; +const REGEX_PREFIX = /^[\+โš ๏ธ]\s*/; +const REGEX_EMPTY_INPUT = /^\+\s*$/; +export function stripContextInputPrefixes(name) { + return name.replace(REGEX_PREFIX, ""); +} +export function getContextOutputName(inputName) { + if (inputName === "base_ctx") + return "CONTEXT"; + return stripContextInputPrefixes(inputName).toUpperCase(); +} +export var InputMutationOperation; +(function (InputMutationOperation) { + InputMutationOperation[InputMutationOperation["UNKNOWN"] = 0] = "UNKNOWN"; + InputMutationOperation[InputMutationOperation["ADDED"] = 1] = "ADDED"; + InputMutationOperation[InputMutationOperation["REMOVED"] = 2] = "REMOVED"; + InputMutationOperation[InputMutationOperation["RENAMED"] = 3] = "RENAMED"; +})(InputMutationOperation || (InputMutationOperation = {})); +export class ContextService { + constructor() { + if (SERVICE) { + throw new Error("ContextService was already instantiated."); + } + } + onInputChanges(node, mutation) { + const childCtxs = getConnectedOutputNodesAndFilterPassThroughs(node, node, 0); + for (const childCtx of childCtxs) { + childCtx.handleUpstreamMutation(mutation); + } + } + getDynamicContextInputsData(node) { + return node + .getContextInputsList() + .map((input, index) => ({ + name: stripContextInputPrefixes(input.name), + type: String(input.type), + index, + })) + .filter((i) => i.type !== "*"); + } + getDynamicContextOutputsData(node) { + return node.outputs.map((output, index) => ({ + name: stripContextInputPrefixes(output.name), + type: String(output.type), + index, + })); + } +} +SERVICE = new ContextService();