From 8244362d7da4d5aed6469387f60dcbc6abab0adf Mon Sep 17 00:00:00 2001 From: huchenlei Date: Wed, 27 Nov 2024 12:29:47 -0500 Subject: [PATCH] Move code --- comfy_execution/validation.py | 29 ++++++++++++++++++ execution.py | 30 +------------------ .../validate_node_input_test.py | 2 +- 3 files changed, 31 insertions(+), 30 deletions(-) create mode 100644 comfy_execution/validation.py diff --git a/comfy_execution/validation.py b/comfy_execution/validation.py new file mode 100644 index 000000000..aff7d6eba --- /dev/null +++ b/comfy_execution/validation.py @@ -0,0 +1,29 @@ +def validate_node_input( + received_type: str, input_type: str, strict: bool = False +) -> bool: + """ + received_type and input_type are both strings of the form "T1,T2,...". + + If strict is True, the input_type must contain the received_type. + For example, if received_type is "STRING" and input_type is "STRING,INT", + this will return True. But if received_type is "STRING,INT" and input_type is + "INT", this will return False. + + If strict is False, the input_type must have overlap with the received_type. + For example, if received_type is "STRING,BOOLEAN" and input_type is "STRING,INT", + this will return True. + """ + # If the types are exactly the same, we can return immediately + if received_type == input_type: + return True + + # Split the type strings into sets for comparison + received_types = set(t.strip() for t in received_type.split(",")) + input_types = set(t.strip() for t in input_type.split(",")) + + if strict: + # In strict mode, all received types must be in the input types + return received_types.issubset(input_types) + else: + # In non-strict mode, there must be at least one type in common + return len(received_types.intersection(input_types)) > 0 diff --git a/execution.py b/execution.py index 84b370a95..fd4fa74fa 100644 --- a/execution.py +++ b/execution.py @@ -17,6 +17,7 @@ import comfy.model_management from comfy_execution.graph import get_input_info, ExecutionList, DynamicPrompt, ExecutionBlocker from comfy_execution.graph_utils import is_link, GraphBuilder from comfy_execution.caching import HierarchicalCache, LRUCache, CacheKeySetInputSignature, CacheKeySetID +from comfy_execution.validation import validate_node_input from comfy.cli_args import args class ExecutionResult(Enum): @@ -528,35 +529,6 @@ class PromptExecutor: comfy.model_management.unload_all_models() -def validate_node_input(received_type: str, input_type: str, strict: bool = False) -> bool: - """ - received_type and input_type are both strings of the form "T1,T2,...". - - If strict is True, the input_type must contain the received_type. - For example, if received_type is "STRING" and input_type is "STRING,INT", - this will return True. But if received_type is "STRING,INT" and input_type is - "INT", this will return False. - - If strict is False, the input_type must have overlap with the received_type. - For example, if received_type is "STRING,BOOLEAN" and input_type is "STRING,INT", - this will return True. - """ - # If the types are exactly the same, we can return immediately - if received_type == input_type: - return True - - # Split the type strings into sets for comparison - received_types = set(t.strip() for t in received_type.split(',')) - input_types = set(t.strip() for t in input_type.split(',')) - - if strict: - # In strict mode, all received types must be in the input types - return received_types.issubset(input_types) - else: - # In non-strict mode, there must be at least one type in common - return len(received_types.intersection(input_types)) > 0 - - def validate_inputs(prompt, item, validated): unique_id = item if unique_id in validated: diff --git a/tests-unit/execution_test/validate_node_input_test.py b/tests-unit/execution_test/validate_node_input_test.py index bc37bac0e..d6605e97f 100644 --- a/tests-unit/execution_test/validate_node_input_test.py +++ b/tests-unit/execution_test/validate_node_input_test.py @@ -1,5 +1,5 @@ import pytest -from execution import validate_node_input +from comfy_execution.validation import validate_node_input def test_exact_match():