diff --git a/__init__.py b/__init__.py
index ed1b35a..317c3a8 100644
--- a/__init__.py
+++ b/__init__.py
@@ -27,6 +27,7 @@ from .py.context_big import RgthreeBigContext
from .py.ksampler_config import RgthreeKSamplerConfig
from .py.sdxl_power_prompt_postive import RgthreeSDXLPowerPromptPositive
from .py.sdxl_power_prompt_simple import RgthreeSDXLPowerPromptSimple
+from .py.any_switch import RgthreeAnySwitch
NODE_CLASS_MAPPINGS = {
RgthreeBigContext.NAME: RgthreeBigContext,
@@ -44,6 +45,7 @@ NODE_CLASS_MAPPINGS = {
RgthreeSDXLEmptyLatentImage.NAME: RgthreeSDXLEmptyLatentImage,
RgthreeSDXLPowerPromptPositive.NAME: RgthreeSDXLPowerPromptPositive,
RgthreeSDXLPowerPromptSimple.NAME: RgthreeSDXLPowerPromptSimple,
+ RgthreeAnySwitch.NAME: RgthreeAnySwitch,
}
diff --git a/py/any_switch.py b/py/any_switch.py
new file mode 100644
index 0000000..ad9458b
--- /dev/null
+++ b/py/any_switch.py
@@ -0,0 +1,52 @@
+import json
+
+from .context_utils import is_context_empty
+from .constants import get_category, get_name
+from .utils import any_type
+
+
+def is_none(value):
+ """Checks if a value is none. Pulled out in case we want to expand what 'None' means."""
+ if value is not None:
+ if isinstance(value, dict) and 'model' in value and 'clip' in value:
+ return is_context_empty(value)
+ return value is None
+
+
+class RgthreeAnySwitch:
+ """The any switch. """
+
+ NAME = get_name("Any Switch")
+ CATEGORY = get_category()
+
+ @classmethod
+ def INPUT_TYPES(cls): # pylint: disable = invalid-name, missing-function-docstring
+ return {
+ "required": {},
+ "optional": {
+ "any_01": (any_type,),
+ "any_02": (any_type,),
+ "any_03": (any_type,),
+ "any_04": (any_type,),
+ "any_05": (any_type,),
+ },
+ }
+
+ RETURN_TYPES = (any_type,)
+ RETURN_NAMES = ('*',)
+ FUNCTION = "switch"
+
+ def switch(self, any_01=None, any_02=None, any_03=None, any_04=None, any_05=None):
+ """Chooses the first non-empty item to output."""
+ any_value = None
+ if not is_none(any_01):
+ any_value = any_01
+ elif not is_none(any_02):
+ any_value = any_02
+ elif not is_none(any_03):
+ any_value = any_03
+ elif not is_none(any_04):
+ any_value = any_04
+ elif not is_none(any_05):
+ any_value = any_05
+ return (any_value,)
diff --git a/py/utils.py b/py/utils.py
new file mode 100644
index 0000000..ded6dc9
--- /dev/null
+++ b/py/utils.py
@@ -0,0 +1,9 @@
+
+class AnyType(str):
+ """A special class that is always equal in not equal comparisons. Credit to pythongosssss"""
+
+ def __ne__(self, __value: object) -> bool:
+ return False
+
+
+any_type = AnyType("*")
diff --git a/ts/any_switch.ts b/ts/any_switch.ts
new file mode 100644
index 0000000..7e8b155
--- /dev/null
+++ b/ts/any_switch.ts
@@ -0,0 +1,110 @@
+// /
+// @ts-ignore
+import { app } from "../../scripts/app.js";
+// @ts-ignore
+import { ComfyWidgets } from "../../scripts/widgets.js";
+import type {
+ INodeInputSlot,
+ INodeOutputSlot,
+ LGraphNodeConstructor,
+ LLink,
+ SerializedLGraphNode,
+ LGraphNode as TLGraphNode,
+ LiteGraph as TLiteGraph,
+} from "./typings/litegraph.js";
+import type { ComfyApp, ComfyObjectInfo } from "./typings/comfy.js";
+import {
+ IoDirection,
+ addConnectionLayoutSupport,
+ applyMixins,
+ followConnectionUntilType,
+ replaceNode,
+} from "./utils.js";
+import { RgthreeBaseNode } from "./base_node.js";
+
+declare const LiteGraph: typeof TLiteGraph;
+declare const LGraphNode: typeof TLGraphNode;
+
+let hasShownAlertForUpdatingInt = false;
+
+class AnySwitchforMixin extends RgthreeBaseNode {
+ static comfyClass? = "";
+
+ private scheduleStabilizePromise: Promise | null = null;
+ private nodeType: string | string[] | null = null;
+
+ override onConnectionsChange(
+ type: number,
+ slotIndex: number,
+ isConnected: boolean,
+ linkInfo: LLink,
+ ioSlot: INodeOutputSlot | INodeInputSlot,
+ ) {
+ super.onConnectionsChange?.(type, slotIndex, isConnected, linkInfo, ioSlot);
+ this.scheduleStabilize();
+ }
+
+ onConnectionsChainChange() {
+ this.scheduleStabilize();
+ }
+
+ scheduleStabilize(ms = 64) {
+ if (!this.scheduleStabilizePromise) {
+ this.scheduleStabilizePromise = new Promise((resolve) => {
+ setTimeout(() => {
+ this.scheduleStabilizePromise = null;
+ this.stabilize();
+ resolve();
+ }, ms);
+ });
+ }
+ return this.scheduleStabilizePromise;
+ }
+
+ stabilize() {
+ // We prefer the inputs, then the output.
+ let connectedType = followConnectionUntilType(this, IoDirection.INPUT, undefined, true);
+ if (!connectedType) {
+ connectedType = followConnectionUntilType(this, IoDirection.OUTPUT, undefined, true);
+ }
+ // TODO: What this doesn't do is broadcast to other nodes when its type changes. Reroute node
+ // does, but, for now, if this was connected to another Any Switch, say, the second one wouldn't
+ // change its type when the first does. The user would need to change the connections.
+ this.nodeType = connectedType?.type || "*";
+ for (const input of this.inputs) {
+ input.type = this.nodeType as string; // So, types can indeed be arrays,,
+ }
+ for (const output of this.outputs) {
+ output.type = this.nodeType as string; // So, types can indeed be arrays,,
+ output.label =
+ output.type === 'RGTHREE_CONTEXT' ? 'CONTEXT' :
+ Array.isArray(this.nodeType) || this.nodeType.includes(",")
+ ? connectedType?.label || String(this.nodeType)
+ : String(this.nodeType);
+ }
+ }
+
+ static override setUp(nodeType: new(title?: any) => T) {
+ AnySwitchforMixin.title = (nodeType as any).title;
+ AnySwitchforMixin.type = (nodeType as any).type || (nodeType as any).title;
+ AnySwitchforMixin.comfyClass = (nodeType as any).comfyClass;
+ setTimeout(() => {
+ AnySwitchforMixin.category = (nodeType as any).category;
+ });
+ applyMixins(nodeType, [RgthreeBaseNode, AnySwitchforMixin]);
+ addConnectionLayoutSupport(nodeType, app, [["Left"], ["Right"]]);
+ }
+}
+
+app.registerExtension({
+ name: "rgthree.AnySwitch",
+ async beforeRegisterNodeDef(
+ nodeType: LGraphNodeConstructor,
+ nodeData: ComfyObjectInfo,
+ app: ComfyApp,
+ ) {
+ if (nodeData.name === "Any Switch (rgthree)") {
+ AnySwitchforMixin.setUp(nodeType as any);
+ }
+ },
+});
diff --git a/ts/reroute.ts b/ts/reroute.ts
index fc9266b..6e2c156 100644
--- a/ts/reroute.ts
+++ b/ts/reroute.ts
@@ -157,6 +157,7 @@ app.registerExtension({
let updateNodes = [];
let inputType = null;
let inputNode = null;
+ let inputNodeOutputSlot = null;
while (currentNode) {
updateNodes.unshift(currentNode);
const linkId: number | null = currentNode.inputs[0]!.link;
@@ -183,8 +184,9 @@ app.registerExtension({
}
} else {
// We've found the end
- inputNode = currentNode;
- inputType = node.outputs[link.origin_slot]?.type ?? null;
+ inputNode = node;
+ inputNodeOutputSlot = link.origin_slot;
+ inputType = node.outputs[inputNodeOutputSlot]?.type ?? null;
break;
}
} else {
@@ -196,6 +198,7 @@ app.registerExtension({
// Find all outputs
const nodes: TLGraphNode[] = [this];
+ let outputNode = null;
let outputType = null;
while (nodes.length) {
currentNode = nodes.pop()!;
@@ -232,6 +235,7 @@ app.registerExtension({
node.disconnectInput(link.target_slot);
} else {
outputType = nodeOutType;
+ outputNode = node;
}
}
}
@@ -261,12 +265,17 @@ app.registerExtension({
}
}
- if (inputNode) {
- const link = app.graph.links[inputNode.inputs[0]!.link];
- if (link) {
- link.color = color;
+ if (inputNode && inputNodeOutputSlot != null) {
+ const links = inputNode.outputs[inputNodeOutputSlot]!.links;
+ for (const l of links || []) {
+ const link = app.graph.links[l];
+ if (link) {
+ link.color = color;
+ }
}
}
+ (inputNode as any)?.onConnectionsChainChange?.();
+ (outputNode as any)?.onConnectionsChainChange?.();
app.graph.setDirtyCanvas(true, true);
}
diff --git a/ts/typings/litegraph.d.ts b/ts/typings/litegraph.d.ts
index 3a54b6a..08ce681 100644
--- a/ts/typings/litegraph.d.ts
+++ b/ts/typings/litegraph.d.ts
@@ -622,7 +622,7 @@ export declare class LGraphNode {
// end @rgthree added
- static title_color: string;
+ static title_color?: string;
static title: string;
static type: null | string;
static widgets_up: boolean;
@@ -1089,6 +1089,17 @@ export declare class LGraphNode {
export type LGraphNodeConstructor = {
new (): T;
+
+ // @rgthree
+ title_mode?:
+ typeof LiteGraph.NORMAL_TITLE |
+ typeof LiteGraph.TRANSPARENT_TITLE |
+ typeof LiteGraph.AUTOHIDE_TITLE |
+ typeof LiteGraph.NO_TITLE;
+ title: string;
+ category: string;
+ type: string;
+ comfyClass?: string;
};
export type SerializedLGraphGroup = {
diff --git a/ts/utils.ts b/ts/utils.ts
index c1ffd3c..b04c023 100644
--- a/ts/utils.ts
+++ b/ts/utils.ts
@@ -407,15 +407,15 @@ export function getConnectedInputNodes(
currentNode?: TLGraphNode,
slot?: number,
passThroughFollowing = PassThroughFollowing.ALL,
-) {
- return getConnectedNodes(startNode, IoDirection.INPUT, currentNode, slot, passThroughFollowing);
+) : TLGraphNode[] {
+ return getConnectedNodes(startNode, IoDirection.INPUT, currentNode, slot, passThroughFollowing).map(n => n.node);
}
export function getConnectedInputNodesAndFilterPassThroughs(
startNode: TLGraphNode,
currentNode?: TLGraphNode,
slot?: number,
passThroughFollowing = PassThroughFollowing.ALL,
-) {
+ ) : TLGraphNode[] {
return filterOutPassthroughNodes(
getConnectedInputNodes(startNode, currentNode, slot, passThroughFollowing),
passThroughFollowing,
@@ -426,38 +426,39 @@ export function getConnectedOutputNodes(
currentNode?: TLGraphNode,
slot?: number,
passThroughFollowing = PassThroughFollowing.ALL,
-) {
- return getConnectedNodes(startNode, IoDirection.OUTPUT, currentNode, slot, passThroughFollowing);
+) : TLGraphNode[] {
+ return getConnectedNodes(startNode, IoDirection.OUTPUT, currentNode, slot, passThroughFollowing).map(n => n.node);
}
export function getConnectedOutputNodesAndFilterPassThroughs(
startNode: TLGraphNode,
currentNode?: TLGraphNode,
slot?: number,
passThroughFollowing = PassThroughFollowing.ALL,
-) {
+) : TLGraphNode[] {
return filterOutPassthroughNodes(
getConnectedOutputNodes(startNode, currentNode, slot, passThroughFollowing),
passThroughFollowing,
);
}
-function getConnectedNodes(
+
+export function getConnectedNodes(
startNode: TLGraphNode,
dir = IoDirection.INPUT,
currentNode?: TLGraphNode,
slot?: number,
passThroughFollowing = PassThroughFollowing.ALL,
-) {
+) : {node:TLGraphNode, slot: number}[] {
currentNode = currentNode || startNode;
- let rootNodes: TLGraphNode[] = [];
+ let rootNodes: {node:TLGraphNode, slot: number}[] = [];
const slotsToRemove = [];
if (startNode === currentNode || shouldPassThrough(currentNode, passThroughFollowing)) {
// const removeDups = startNode === currentNode;
let linkIds: Array;
if (dir == IoDirection.OUTPUT) {
- linkIds = currentNode.outputs?.flatMap((i) => i.links);
+ linkIds = currentNode.outputs?.flatMap((i) => i.links) || [];
} else {
- linkIds = currentNode.inputs?.map((i) => i.link);
+ linkIds = currentNode.inputs?.map((i) => i.link) || [];
}
if (typeof slot == "number" && slot > -1) {
if (linkIds[slot]) {
@@ -473,12 +474,13 @@ function getConnectedNodes(
continue;
}
const connectedId = dir == IoDirection.OUTPUT ? link.target_id : link.origin_id;
+ const originSlot = dir == IoDirection.OUTPUT ? link.target_slot : link.origin_slot;
const originNode: TLGraphNode = graph.getNodeById(connectedId)!;
if (!link) {
console.error("No connected node found... weird");
continue;
}
- if (rootNodes.includes(originNode)) {
+ if (rootNodes.some((n) => n.node == originNode)) {
console.log(
`${startNode.title} (${startNode.id}) seems to have two links to ${originNode.title} (${
originNode.id
@@ -486,7 +488,7 @@ function getConnectedNodes(
);
} else {
// Add the node and, if it's a pass through, let's collect all its nodes as well.
- rootNodes.push(originNode);
+ rootNodes.push({node: originNode, slot: originSlot});
if (shouldPassThrough(originNode, passThroughFollowing)) {
for (const foundNode of getConnectedNodes(startNode, dir, originNode)) {
if (!rootNodes.includes(foundNode)) {
@@ -500,6 +502,76 @@ function getConnectedNodes(
return rootNodes;
}
+type ConnectionType = { type: string | string[]; 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: TLGraphNode,
+ 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 || 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: TLGraphNode = 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];
+ console.log(connectedSlot);
+ if (connectedSlot?.type != null && connectedSlot?.type != "*") {
+ return {
+ type: connectedSlot.type as string,
+ label: connectedSlot?.label || connectedSlot?.name,
+ };
+ } else if (connectedSlot?.type == "*") {
+ return followConnectionUntilType(connectedNode, dir);
+ }
+ }
+ return null;
+}
+
export async function replaceNode(
existingNode: TLGraphNode,
typeOrNewNode: string | TLGraphNode,
diff --git a/web/any_switch.js b/web/any_switch.js
new file mode 100644
index 0000000..e54e7b6
--- /dev/null
+++ b/web/any_switch.js
@@ -0,0 +1,68 @@
+import { app } from "../../scripts/app.js";
+import { IoDirection, addConnectionLayoutSupport, applyMixins, followConnectionUntilType, } from "./utils.js";
+import { RgthreeBaseNode } from "./base_node.js";
+let hasShownAlertForUpdatingInt = false;
+class AnySwitchforMixin extends RgthreeBaseNode {
+ constructor() {
+ super(...arguments);
+ this.scheduleStabilizePromise = null;
+ this.nodeType = null;
+ }
+ onConnectionsChange(type, slotIndex, isConnected, linkInfo, ioSlot) {
+ var _a;
+ (_a = super.onConnectionsChange) === null || _a === void 0 ? void 0 : _a.call(this, type, slotIndex, isConnected, linkInfo, ioSlot);
+ this.scheduleStabilize();
+ }
+ onConnectionsChainChange() {
+ this.scheduleStabilize();
+ }
+ scheduleStabilize(ms = 64) {
+ if (!this.scheduleStabilizePromise) {
+ this.scheduleStabilizePromise = new Promise((resolve) => {
+ setTimeout(() => {
+ this.scheduleStabilizePromise = null;
+ this.stabilize();
+ resolve();
+ }, ms);
+ });
+ }
+ return this.scheduleStabilizePromise;
+ }
+ stabilize() {
+ let connectedType = followConnectionUntilType(this, IoDirection.INPUT, undefined, true);
+ if (!connectedType) {
+ connectedType = followConnectionUntilType(this, IoDirection.OUTPUT, undefined, true);
+ }
+ this.nodeType = (connectedType === null || connectedType === void 0 ? void 0 : connectedType.type) || "*";
+ for (const input of this.inputs) {
+ input.type = this.nodeType;
+ }
+ for (const output of this.outputs) {
+ output.type = this.nodeType;
+ output.label =
+ output.type === 'RGTHREE_CONTEXT' ? 'CONTEXT' :
+ Array.isArray(this.nodeType) || this.nodeType.includes(",")
+ ? (connectedType === null || connectedType === void 0 ? void 0 : connectedType.label) || String(this.nodeType)
+ : String(this.nodeType);
+ }
+ }
+ static setUp(nodeType) {
+ AnySwitchforMixin.title = nodeType.title;
+ AnySwitchforMixin.type = nodeType.type || nodeType.title;
+ AnySwitchforMixin.comfyClass = nodeType.comfyClass;
+ setTimeout(() => {
+ AnySwitchforMixin.category = nodeType.category;
+ });
+ applyMixins(nodeType, [RgthreeBaseNode, AnySwitchforMixin]);
+ addConnectionLayoutSupport(nodeType, app, [["Left"], ["Right"]]);
+ }
+}
+AnySwitchforMixin.comfyClass = "";
+app.registerExtension({
+ name: "rgthree.AnySwitch",
+ async beforeRegisterNodeDef(nodeType, nodeData, app) {
+ if (nodeData.name === "Any Switch (rgthree)") {
+ AnySwitchforMixin.setUp(nodeType);
+ }
+ },
+});
diff --git a/web/reroute.js b/web/reroute.js
index 6b6f242..859f8bd 100644
--- a/web/reroute.js
+++ b/web/reroute.js
@@ -85,7 +85,7 @@ app.registerExtension({
return this.schedulePromise;
}
stabilize() {
- var _a, _b, _c, _d, _e;
+ var _a, _b, _c, _d, _e, _f, _g;
if (this.configuring) {
return;
}
@@ -93,6 +93,7 @@ app.registerExtension({
let updateNodes = [];
let inputType = null;
let inputNode = null;
+ let inputNodeOutputSlot = null;
while (currentNode) {
updateNodes.unshift(currentNode);
const linkId = currentNode.inputs[0].link;
@@ -115,8 +116,9 @@ app.registerExtension({
}
}
else {
- inputNode = currentNode;
- inputType = (_b = (_a = node.outputs[link.origin_slot]) === null || _a === void 0 ? void 0 : _a.type) !== null && _b !== void 0 ? _b : null;
+ inputNode = node;
+ inputNodeOutputSlot = link.origin_slot;
+ inputType = (_b = (_a = node.outputs[inputNodeOutputSlot]) === null || _a === void 0 ? void 0 : _a.type) !== null && _b !== void 0 ? _b : null;
break;
}
}
@@ -126,6 +128,7 @@ app.registerExtension({
}
}
const nodes = [this];
+ let outputNode = null;
let outputType = null;
while (nodes.length) {
currentNode = nodes.pop();
@@ -157,6 +160,7 @@ app.registerExtension({
}
else {
outputType = nodeOutType;
+ outputNode = node;
}
}
}
@@ -179,12 +183,17 @@ app.registerExtension({
}
}
}
- if (inputNode) {
- const link = app.graph.links[inputNode.inputs[0].link];
- if (link) {
- link.color = color;
+ if (inputNode && inputNodeOutputSlot != null) {
+ const links = inputNode.outputs[inputNodeOutputSlot].links;
+ for (const l of links || []) {
+ const link = app.graph.links[l];
+ if (link) {
+ link.color = color;
+ }
}
}
+ (_f = inputNode === null || inputNode === void 0 ? void 0 : inputNode.onConnectionsChainChange) === null || _f === void 0 ? void 0 : _f.call(inputNode);
+ (_g = outputNode === null || outputNode === void 0 ? void 0 : outputNode.onConnectionsChainChange) === null || _g === void 0 ? void 0 : _g.call(outputNode);
app.graph.setDirtyCanvas(true, true);
}
computeSize(out) {
diff --git a/web/utils.js b/web/utils.js
index eebe9b3..4feff0c 100644
--- a/web/utils.js
+++ b/web/utils.js
@@ -283,18 +283,18 @@ export function filterOutPassthroughNodes(nodes, passThroughFollowing = PassThro
return nodes.filter((n) => !shouldPassThrough(n, passThroughFollowing));
}
export function getConnectedInputNodes(startNode, currentNode, slot, passThroughFollowing = PassThroughFollowing.ALL) {
- return getConnectedNodes(startNode, IoDirection.INPUT, currentNode, slot, passThroughFollowing);
+ return getConnectedNodes(startNode, IoDirection.INPUT, currentNode, slot, passThroughFollowing).map(n => n.node);
}
export function getConnectedInputNodesAndFilterPassThroughs(startNode, currentNode, slot, passThroughFollowing = PassThroughFollowing.ALL) {
return filterOutPassthroughNodes(getConnectedInputNodes(startNode, currentNode, slot, passThroughFollowing), passThroughFollowing);
}
export function getConnectedOutputNodes(startNode, currentNode, slot, passThroughFollowing = PassThroughFollowing.ALL) {
- return getConnectedNodes(startNode, IoDirection.OUTPUT, currentNode, slot, passThroughFollowing);
+ return getConnectedNodes(startNode, IoDirection.OUTPUT, currentNode, slot, passThroughFollowing).map(n => n.node);
}
export function getConnectedOutputNodesAndFilterPassThroughs(startNode, currentNode, slot, passThroughFollowing = PassThroughFollowing.ALL) {
return filterOutPassthroughNodes(getConnectedOutputNodes(startNode, currentNode, slot, passThroughFollowing), passThroughFollowing);
}
-function getConnectedNodes(startNode, dir = IoDirection.INPUT, currentNode, slot, passThroughFollowing = PassThroughFollowing.ALL) {
+export function getConnectedNodes(startNode, dir = IoDirection.INPUT, currentNode, slot, passThroughFollowing = PassThroughFollowing.ALL) {
var _a, _b;
currentNode = currentNode || startNode;
let rootNodes = [];
@@ -302,10 +302,10 @@ function getConnectedNodes(startNode, dir = IoDirection.INPUT, currentNode, slot
if (startNode === currentNode || shouldPassThrough(currentNode, passThroughFollowing)) {
let linkIds;
if (dir == IoDirection.OUTPUT) {
- linkIds = (_a = currentNode.outputs) === null || _a === void 0 ? void 0 : _a.flatMap((i) => i.links);
+ linkIds = ((_a = currentNode.outputs) === null || _a === void 0 ? void 0 : _a.flatMap((i) => i.links)) || [];
}
else {
- linkIds = (_b = currentNode.inputs) === null || _b === void 0 ? void 0 : _b.map((i) => i.link);
+ linkIds = ((_b = currentNode.inputs) === null || _b === void 0 ? void 0 : _b.map((i) => i.link)) || [];
}
if (typeof slot == "number" && slot > -1) {
if (linkIds[slot]) {
@@ -322,16 +322,17 @@ function getConnectedNodes(startNode, dir = IoDirection.INPUT, currentNode, slot
continue;
}
const connectedId = dir == IoDirection.OUTPUT ? link.target_id : link.origin_id;
+ const originSlot = dir == IoDirection.OUTPUT ? link.target_slot : link.origin_slot;
const originNode = graph.getNodeById(connectedId);
if (!link) {
console.error("No connected node found... weird");
continue;
}
- if (rootNodes.includes(originNode)) {
+ 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 {
- rootNodes.push(originNode);
+ rootNodes.push({ node: originNode, slot: originSlot });
if (shouldPassThrough(originNode, passThroughFollowing)) {
for (const foundNode of getConnectedNodes(startNode, dir, originNode)) {
if (!rootNodes.includes(foundNode)) {
@@ -344,6 +345,54 @@ function getConnectedNodes(startNode, dir = IoDirection.INPUT, currentNode, slot
}
return rootNodes;
}
+export function followConnectionUntilType(node, dir, slotNum, skipSelf = false) {
+ const slots = dir === IoDirection.OUTPUT ? node.outputs : node.inputs;
+ if (!slots || !slots.length) {
+ return null;
+ }
+ let type = 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;
+}
+function getTypeFromSlot(slot, dir, skipSelf = false) {
+ let graph = app.graph;
+ let type = slot === null || slot === void 0 ? void 0 : slot.type;
+ if (!skipSelf && type != null && type != "*") {
+ return { type: type, label: (slot === null || slot === void 0 ? void 0 : slot.label) || (slot === null || slot === void 0 ? void 0 : 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 = graph.getNodeById(connectedId);
+ const connectedSlots = dir === IoDirection.OUTPUT ? connectedNode.inputs : connectedNode.outputs;
+ let connectedSlot = connectedSlots[connectedSlotNum];
+ console.log(connectedSlot);
+ if ((connectedSlot === null || connectedSlot === void 0 ? void 0 : connectedSlot.type) != null && (connectedSlot === null || connectedSlot === void 0 ? void 0 : connectedSlot.type) != "*") {
+ return {
+ type: connectedSlot.type,
+ label: (connectedSlot === null || connectedSlot === void 0 ? void 0 : connectedSlot.label) || (connectedSlot === null || connectedSlot === void 0 ? void 0 : connectedSlot.name),
+ };
+ }
+ else if ((connectedSlot === null || connectedSlot === void 0 ? void 0 : connectedSlot.type) == "*") {
+ return followConnectionUntilType(connectedNode, dir);
+ }
+ }
+ return null;
+}
export async function replaceNode(existingNode, typeOrNewNode, inputNameMap) {
const existingCtor = existingNode.constructor;
const newNode = typeof typeOrNewNode === "string" ? LiteGraph.createNode(typeOrNewNode) : typeOrNewNode;