158 lines
4.4 KiB
Python
158 lines
4.4 KiB
Python
import execution
|
|
import nodes
|
|
|
|
from .floodgate import FloodGate
|
|
|
|
NODE_CLASS_MAPPINGS = {"FloodGate": FloodGate}
|
|
NODE_DISPLAY_NAME_MAPPINGS = {"FloodGate": "Flood Gate"}
|
|
|
|
|
|
def find_gate(prompt: dict) -> list:
|
|
"""Find the Unique ID of the Floodgate Node"""
|
|
gate_IDs = []
|
|
|
|
for k, v in prompt.items():
|
|
if v["class_type"] == "FloodGate":
|
|
gate_IDs.append(k)
|
|
|
|
# if len(gate_IDs) > 1:
|
|
# print('[Warning] Multiple Floodgates Detected is still experimental!')
|
|
|
|
return gate_IDs
|
|
|
|
|
|
def block_gate(prompt: dict, gate_ID: str, floodgate_open: bool) -> dict:
|
|
""" "Bypass" the Nodes that should be Blocked"""
|
|
nodes_affected = []
|
|
|
|
try:
|
|
sauce_id, out_index = prompt[gate_ID]["inputs"]["source"]
|
|
except KeyError:
|
|
# Floodgate is not connected; let ComfyUI raise the error
|
|
return prompt
|
|
|
|
sauce_class = nodes.NODE_CLASS_MAPPINGS[prompt[sauce_id]["class_type"]]
|
|
sauce_type = str(sauce_class.RETURN_TYPES[out_index]).lower().strip()
|
|
|
|
for node, data in prompt.items():
|
|
for k, v in data["inputs"].items():
|
|
if not isinstance(v, list):
|
|
continue
|
|
|
|
if gate_ID in v:
|
|
|
|
target_class = nodes.NODE_CLASS_MAPPINGS[data["class_type"]]
|
|
target_type = (
|
|
(target_class.INPUT_TYPES()["required"][k][0]).lower().strip()
|
|
)
|
|
|
|
if sauce_type != target_type:
|
|
raise IOError()
|
|
|
|
if (not floodgate_open) and (v[1] == 1):
|
|
nodes_affected.append(node)
|
|
break
|
|
|
|
if (floodgate_open) and (v[1] == 0):
|
|
nodes_affected.append(node)
|
|
break
|
|
|
|
for key in nodes_affected:
|
|
del prompt[key]
|
|
|
|
if len(nodes_affected) > 0:
|
|
return recursive_block_gate(prompt, nodes_affected)
|
|
else:
|
|
return prompt
|
|
|
|
|
|
def recursive_block_gate(prompt: dict, node_IDs: list) -> dict:
|
|
"""Block the subsequent nodes of which source has been blocked"""
|
|
to_delete = []
|
|
|
|
for node, data in prompt.items():
|
|
for k, v in data["inputs"].items():
|
|
# Connection is always a List
|
|
if not isinstance(v, list):
|
|
continue
|
|
|
|
if any(ID in v for ID in node_IDs):
|
|
to_delete.append(node)
|
|
break
|
|
|
|
for key in to_delete:
|
|
del prompt[key]
|
|
|
|
if len(to_delete) > 0:
|
|
return recursive_block_gate(prompt, to_delete)
|
|
else:
|
|
return prompt
|
|
|
|
|
|
original_validate = execution.validate_prompt
|
|
|
|
|
|
async def hijack_validate(*args):
|
|
for arg in args:
|
|
if isinstance(arg, dict):
|
|
prompt = arg
|
|
break
|
|
|
|
gate_IDs: list = find_gate(prompt)
|
|
|
|
if len(gate_IDs) == 0:
|
|
return await original_validate(*args)
|
|
|
|
for ID in gate_IDs:
|
|
if ID not in prompt.keys():
|
|
continue
|
|
|
|
try:
|
|
gate_open = prompt[ID]["inputs"]["gate_open"]
|
|
|
|
if isinstance(gate_open, (tuple, list)) and len(gate_open) == 0:
|
|
gate_open = bool(gate_open[0])
|
|
if type(gate_open) is bool:
|
|
prompt = block_gate(prompt, ID, gate_open)
|
|
elif type(gate_open) is list:
|
|
sauce_id, conn_id = gate_open
|
|
|
|
gate_open = list(prompt[sauce_id]["inputs"].values())[conn_id]
|
|
if type(gate_open) is bool:
|
|
prompt = block_gate(prompt, ID, gate_open)
|
|
else:
|
|
raise ValueError
|
|
else:
|
|
raise ValueError
|
|
|
|
except IOError:
|
|
return (
|
|
False,
|
|
{
|
|
"type": "floodgate_io_mismatch",
|
|
"message": "Floodgate IO Type Mismatch",
|
|
"details": "source cannot be connected to outputs",
|
|
"extra_info": {},
|
|
},
|
|
[],
|
|
[],
|
|
)
|
|
|
|
except ValueError:
|
|
return (
|
|
False,
|
|
{
|
|
"type": "floodgate_invalid_boolean",
|
|
"message": "Floodgate Unable to Determine Boolean",
|
|
"details": "please use a primitive boolean node",
|
|
"extra_info": {},
|
|
},
|
|
[],
|
|
[],
|
|
)
|
|
|
|
return await original_validate(*args)
|
|
|
|
|
|
execution.validate_prompt = hijack_validate
|