CoolFace
Apppublic

nef7/my-comfyui-workflow

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
execution.py995 linesDownload Raw Back to root
1import sys2import copy3import logging4import threading5import heapq6import time7import traceback8from enum import Enum9import inspect10from typing import List, Literal, NamedTuple, Optional11 12import torch13import nodes14 15import comfy.model_management16from comfy_execution.graph import get_input_info, ExecutionList, DynamicPrompt, ExecutionBlocker17from comfy_execution.graph_utils import is_link, GraphBuilder18from comfy_execution.caching import HierarchicalCache, LRUCache, CacheKeySetInputSignature, CacheKeySetID19from comfy_execution.validation import validate_node_input20 21class ExecutionResult(Enum):22    SUCCESS = 023    FAILURE = 124    PENDING = 225 26class DuplicateNodeError(Exception):27    pass28 29class IsChangedCache:30    def __init__(self, dynprompt, outputs_cache):31        self.dynprompt = dynprompt32        self.outputs_cache = outputs_cache33        self.is_changed = {}34 35    def get(self, node_id):36        if node_id in self.is_changed:37            return self.is_changed[node_id]38 39        node = self.dynprompt.get_node(node_id)40        class_type = node["class_type"]41        class_def = nodes.NODE_CLASS_MAPPINGS[class_type]42        if not hasattr(class_def, "IS_CHANGED"):43            self.is_changed[node_id] = False44            return self.is_changed[node_id]45 46        if "is_changed" in node:47            self.is_changed[node_id] = node["is_changed"]48            return self.is_changed[node_id]49 50        # Intentionally do not use cached outputs here. We only want constants in IS_CHANGED51        input_data_all, _ = get_input_data(node["inputs"], class_def, node_id, None)52        try:53            is_changed = _map_node_over_list(class_def, input_data_all, "IS_CHANGED")54            node["is_changed"] = [None if isinstance(x, ExecutionBlocker) else x for x in is_changed]55        except Exception as e:56            logging.warning("WARNING: {}".format(e))57            node["is_changed"] = float("NaN")58        finally:59            self.is_changed[node_id] = node["is_changed"]60        return self.is_changed[node_id]61 62class CacheSet:63    def __init__(self, lru_size=None):64        if lru_size is None or lru_size == 0:65            self.init_classic_cache()66        else:67            self.init_lru_cache(lru_size)68        self.all = [self.outputs, self.ui, self.objects]69 70    # Useful for those with ample RAM/VRAM -- allows experimenting without71    # blowing away the cache every time72    def init_lru_cache(self, cache_size):73        self.outputs = LRUCache(CacheKeySetInputSignature, max_size=cache_size)74        self.ui = LRUCache(CacheKeySetInputSignature, max_size=cache_size)75        self.objects = HierarchicalCache(CacheKeySetID)76 77    # Performs like the old cache -- dump data ASAP78    def init_classic_cache(self):79        self.outputs = HierarchicalCache(CacheKeySetInputSignature)80        self.ui = HierarchicalCache(CacheKeySetInputSignature)81        self.objects = HierarchicalCache(CacheKeySetID)82 83    def recursive_debug_dump(self):84        result = {85            "outputs": self.outputs.recursive_debug_dump(),86            "ui": self.ui.recursive_debug_dump(),87        }88        return result89 90def get_input_data(inputs, class_def, unique_id, outputs=None, dynprompt=None, extra_data={}):91    valid_inputs = class_def.INPUT_TYPES()92    input_data_all = {}93    missing_keys = {}94    for x in inputs:95        input_data = inputs[x]96        input_type, input_category, input_info = get_input_info(class_def, x, valid_inputs)97        def mark_missing():98            missing_keys[x] = True99            input_data_all[x] = (None,)100        if is_link(input_data) and (not input_info or not input_info.get("rawLink", False)):101            input_unique_id = input_data[0]102            output_index = input_data[1]103            if outputs is None:104                mark_missing()105                continue # This might be a lazily-evaluated input106            cached_output = outputs.get(input_unique_id)107            if cached_output is None:108                mark_missing()109                continue110            if output_index >= len(cached_output):111                mark_missing()112                continue113            obj = cached_output[output_index]114            input_data_all[x] = obj115        elif input_category is not None:116            input_data_all[x] = [input_data]117 118    if "hidden" in valid_inputs:119        h = valid_inputs["hidden"]120        for x in h:121            if h[x] == "PROMPT":122                input_data_all[x] = [dynprompt.get_original_prompt() if dynprompt is not None else {}]123            if h[x] == "DYNPROMPT":124                input_data_all[x] = [dynprompt]125            if h[x] == "EXTRA_PNGINFO":126                input_data_all[x] = [extra_data.get('extra_pnginfo', None)]127            if h[x] == "UNIQUE_ID":128                input_data_all[x] = [unique_id]129    return input_data_all, missing_keys130 131map_node_over_list = None #Don't hook this please132 133def _map_node_over_list(obj, input_data_all, func, allow_interrupt=False, execution_block_cb=None, pre_execute_cb=None):134    # check if node wants the lists135    input_is_list = getattr(obj, "INPUT_IS_LIST", False)136 137    if len(input_data_all) == 0:138        max_len_input = 0139    else:140        max_len_input = max(len(x) for x in input_data_all.values())141 142    # get a slice of inputs, repeat last input when list isn't long enough143    def slice_dict(d, i):144        return {k: v[i if len(v) > i else -1] for k, v in d.items()}145 146    results = []147    def process_inputs(inputs, index=None, input_is_list=False):148        if allow_interrupt:149            nodes.before_node_execution()150        execution_block = None151        for k, v in inputs.items():152            if input_is_list:153                for e in v:154                    if isinstance(e, ExecutionBlocker):155                        v = e156                        break157            if isinstance(v, ExecutionBlocker):158                execution_block = execution_block_cb(v) if execution_block_cb else v159                break160        if execution_block is None:161            if pre_execute_cb is not None and index is not None:162                pre_execute_cb(index)163            results.append(getattr(obj, func)(**inputs))164        else:165            results.append(execution_block)166 167    if input_is_list:168        process_inputs(input_data_all, 0, input_is_list=input_is_list)169    elif max_len_input == 0:170        process_inputs({})171    else:172        for i in range(max_len_input):173            input_dict = slice_dict(input_data_all, i)174            process_inputs(input_dict, i)175    return results176 177def merge_result_data(results, obj):178    # check which outputs need concatenating179    output = []180    output_is_list = [False] * len(results[0])181    if hasattr(obj, "OUTPUT_IS_LIST"):182        output_is_list = obj.OUTPUT_IS_LIST183 184    # merge node execution results185    for i, is_list in zip(range(len(results[0])), output_is_list):186        if is_list:187            value = []188            for o in results:189                if isinstance(o[i], ExecutionBlocker):190                    value.append(o[i])191                else:192                    value.extend(o[i])193            output.append(value)194        else:195            output.append([o[i] for o in results])196    return output197 198def get_output_data(obj, input_data_all, execution_block_cb=None, pre_execute_cb=None):199    results = []200    uis = []201    subgraph_results = []202    return_values = _map_node_over_list(obj, input_data_all, obj.FUNCTION, allow_interrupt=True, execution_block_cb=execution_block_cb, pre_execute_cb=pre_execute_cb)203    has_subgraph = False204    for i in range(len(return_values)):205        r = return_values[i]206        if isinstance(r, dict):207            if 'ui' in r:208                uis.append(r['ui'])209            if 'expand' in r:210                # Perform an expansion, but do not append results211                has_subgraph = True212                new_graph = r['expand']213                result = r.get("result", None)214                if isinstance(result, ExecutionBlocker):215                    result = tuple([result] * len(obj.RETURN_TYPES))216                subgraph_results.append((new_graph, result))217            elif 'result' in r:218                result = r.get("result", None)219                if isinstance(result, ExecutionBlocker):220                    result = tuple([result] * len(obj.RETURN_TYPES))221                results.append(result)222                subgraph_results.append((None, result))223        else:224            if isinstance(r, ExecutionBlocker):225                r = tuple([r] * len(obj.RETURN_TYPES))226            results.append(r)227            subgraph_results.append((None, r))228 229    if has_subgraph:230        output = subgraph_results231    elif len(results) > 0:232        output = merge_result_data(results, obj)233    else:234        output = []235    ui = dict()236    if len(uis) > 0:237        ui = {k: [y for x in uis for y in x[k]] for k in uis[0].keys()}238    return output, ui, has_subgraph239 240def format_value(x):241    if x is None:242        return None243    elif isinstance(x, (int, float, bool, str)):244        return x245    else:246        return str(x)247 248def execute(server, dynprompt, caches, current_item, extra_data, executed, prompt_id, execution_list, pending_subgraph_results):249    unique_id = current_item250    real_node_id = dynprompt.get_real_node_id(unique_id)251    display_node_id = dynprompt.get_display_node_id(unique_id)252    parent_node_id = dynprompt.get_parent_node_id(unique_id)253    inputs = dynprompt.get_node(unique_id)['inputs']254    class_type = dynprompt.get_node(unique_id)['class_type']255    class_def = nodes.NODE_CLASS_MAPPINGS[class_type]256    if caches.outputs.get(unique_id) is not None:257        if server.client_id is not None:258            cached_output = caches.ui.get(unique_id) or {}259            server.send_sync("executed", { "node": unique_id, "display_node": display_node_id, "output": cached_output.get("output",None), "prompt_id": prompt_id }, server.client_id)260        return (ExecutionResult.SUCCESS, None, None)261 262    input_data_all = None263    try:264        if unique_id in pending_subgraph_results:265            cached_results = pending_subgraph_results[unique_id]266            resolved_outputs = []267            for is_subgraph, result in cached_results:268                if not is_subgraph:269                    resolved_outputs.append(result)270                else:271                    resolved_output = []272                    for r in result:273                        if is_link(r):274                            source_node, source_output = r[0], r[1]275                            node_output = caches.outputs.get(source_node)[source_output]276                            for o in node_output:277                                resolved_output.append(o)278 279                        else:280                            resolved_output.append(r)281                    resolved_outputs.append(tuple(resolved_output))282            output_data = merge_result_data(resolved_outputs, class_def)283            output_ui = []284            has_subgraph = False285        else:286            input_data_all, missing_keys = get_input_data(inputs, class_def, unique_id, caches.outputs, dynprompt, extra_data)287            if server.client_id is not None:288                server.last_node_id = display_node_id289                server.send_sync("executing", { "node": unique_id, "display_node": display_node_id, "prompt_id": prompt_id }, server.client_id)290 291            obj = caches.objects.get(unique_id)292            if obj is None:293                obj = class_def()294                caches.objects.set(unique_id, obj)295 296            if hasattr(obj, "check_lazy_status"):297                required_inputs = _map_node_over_list(obj, input_data_all, "check_lazy_status", allow_interrupt=True)298                required_inputs = set(sum([r for r in required_inputs if isinstance(r,list)], []))299                required_inputs = [x for x in required_inputs if isinstance(x,str) and (300                    x not in input_data_all or x in missing_keys301                )]302                if len(required_inputs) > 0:303                    for i in required_inputs:304                        execution_list.make_input_strong_link(unique_id, i)305                    return (ExecutionResult.PENDING, None, None)306 307            def execution_block_cb(block):308                if block.message is not None:309                    mes = {310                        "prompt_id": prompt_id,311                        "node_id": unique_id,312                        "node_type": class_type,313                        "executed": list(executed),314 315                        "exception_message": f"Execution Blocked: {block.message}",316                        "exception_type": "ExecutionBlocked",317                        "traceback": [],318                        "current_inputs": [],319                        "current_outputs": [],320                    }321                    server.send_sync("execution_error", mes, server.client_id)322                    return ExecutionBlocker(None)323                else:324                    return block325            def pre_execute_cb(call_index):326                GraphBuilder.set_default_prefix(unique_id, call_index, 0)327            output_data, output_ui, has_subgraph = get_output_data(obj, input_data_all, execution_block_cb=execution_block_cb, pre_execute_cb=pre_execute_cb)328        if len(output_ui) > 0:329            caches.ui.set(unique_id, {330                "meta": {331                    "node_id": unique_id,332                    "display_node": display_node_id,333                    "parent_node": parent_node_id,334                    "real_node_id": real_node_id,335                },336                "output": output_ui337            })338            if server.client_id is not None:339                server.send_sync("executed", { "node": unique_id, "display_node": display_node_id, "output": output_ui, "prompt_id": prompt_id }, server.client_id)340        if has_subgraph:341            cached_outputs = []342            new_node_ids = []343            new_output_ids = []344            new_output_links = []345            for i in range(len(output_data)):346                new_graph, node_outputs = output_data[i]347                if new_graph is None:348                    cached_outputs.append((False, node_outputs))349                else:350                    # Check for conflicts351                    for node_id in new_graph.keys():352                        if dynprompt.has_node(node_id):353                            raise DuplicateNodeError(f"Attempt to add duplicate node {node_id}. Ensure node ids are unique and deterministic or use graph_utils.GraphBuilder.")354                    for node_id, node_info in new_graph.items():355                        new_node_ids.append(node_id)356                        display_id = node_info.get("override_display_id", unique_id)357                        dynprompt.add_ephemeral_node(node_id, node_info, unique_id, display_id)358                        # Figure out if the newly created node is an output node359                        class_type = node_info["class_type"]360                        class_def = nodes.NODE_CLASS_MAPPINGS[class_type]361                        if hasattr(class_def, 'OUTPUT_NODE') and class_def.OUTPUT_NODE == True:362                            new_output_ids.append(node_id)363                    for i in range(len(node_outputs)):364                        if is_link(node_outputs[i]):365                            from_node_id, from_socket = node_outputs[i][0], node_outputs[i][1]366                            new_output_links.append((from_node_id, from_socket))367                    cached_outputs.append((True, node_outputs))368            new_node_ids = set(new_node_ids)369            for cache in caches.all:370                cache.ensure_subcache_for(unique_id, new_node_ids).clean_unused()371            for node_id in new_output_ids:372                execution_list.add_node(node_id)373            for link in new_output_links:374                execution_list.add_strong_link(link[0], link[1], unique_id)375            pending_subgraph_results[unique_id] = cached_outputs376            return (ExecutionResult.PENDING, None, None)377        caches.outputs.set(unique_id, output_data)378    except comfy.model_management.InterruptProcessingException as iex:379        logging.info("Processing interrupted")380 381        # skip formatting inputs/outputs382        error_details = {383            "node_id": real_node_id,384        }385 386        return (ExecutionResult.FAILURE, error_details, iex)387    except Exception as ex:388        typ, _, tb = sys.exc_info()389        exception_type = full_type_name(typ)390        input_data_formatted = {}391        if input_data_all is not None:392            input_data_formatted = {}393            for name, inputs in input_data_all.items():394                input_data_formatted[name] = [format_value(x) for x in inputs]395 396        logging.error(f"!!! Exception during processing !!! {ex}")397        logging.error(traceback.format_exc())398 399        error_details = {400            "node_id": real_node_id,401            "exception_message": str(ex),402            "exception_type": exception_type,403            "traceback": traceback.format_tb(tb),404            "current_inputs": input_data_formatted405        }406        if isinstance(ex, comfy.model_management.OOM_EXCEPTION):407            logging.error("Got an OOM, unloading all loaded models.")408            comfy.model_management.unload_all_models()409 410        return (ExecutionResult.FAILURE, error_details, ex)411 412    executed.add(unique_id)413 414    return (ExecutionResult.SUCCESS, None, None)415 416class PromptExecutor:417    def __init__(self, server, lru_size=None):418        self.lru_size = lru_size419        self.server = server420        self.reset()421 422    def reset(self):423        self.caches = CacheSet(self.lru_size)424        self.status_messages = []425        self.success = True426 427    def add_message(self, event, data: dict, broadcast: bool):428        data = {429            **data,430            "timestamp": int(time.time() * 1000),431        }432        self.status_messages.append((event, data))433        if self.server.client_id is not None or broadcast:434            self.server.send_sync(event, data, self.server.client_id)435 436    def handle_execution_error(self, prompt_id, prompt, current_outputs, executed, error, ex):437        node_id = error["node_id"]438        class_type = prompt[node_id]["class_type"]439 440        # First, send back the status to the frontend depending441        # on the exception type442        if isinstance(ex, comfy.model_management.InterruptProcessingException):443            mes = {444                "prompt_id": prompt_id,445                "node_id": node_id,446                "node_type": class_type,447                "executed": list(executed),448            }449            self.add_message("execution_interrupted", mes, broadcast=True)450        else:451            mes = {452                "prompt_id": prompt_id,453                "node_id": node_id,454                "node_type": class_type,455                "executed": list(executed),456                "exception_message": error["exception_message"],457                "exception_type": error["exception_type"],458                "traceback": error["traceback"],459                "current_inputs": error["current_inputs"],460                "current_outputs": list(current_outputs),461            }462            self.add_message("execution_error", mes, broadcast=False)463 464    def execute(self, prompt, prompt_id, extra_data={}, execute_outputs=[]):465        nodes.interrupt_processing(False)466 467        if "client_id" in extra_data:468            self.server.client_id = extra_data["client_id"]469        else:470            self.server.client_id = None471 472        self.status_messages = []473        self.add_message("execution_start", { "prompt_id": prompt_id}, broadcast=False)474 475        with torch.inference_mode():476            dynamic_prompt = DynamicPrompt(prompt)477            is_changed_cache = IsChangedCache(dynamic_prompt, self.caches.outputs)478            for cache in self.caches.all:479                cache.set_prompt(dynamic_prompt, prompt.keys(), is_changed_cache)480                cache.clean_unused()481 482            cached_nodes = []483            for node_id in prompt:484                if self.caches.outputs.get(node_id) is not None:485                    cached_nodes.append(node_id)486 487            comfy.model_management.cleanup_models_gc()488            self.add_message("execution_cached",489                          { "nodes": cached_nodes, "prompt_id": prompt_id},490                          broadcast=False)491            pending_subgraph_results = {}492            executed = set()493            execution_list = ExecutionList(dynamic_prompt, self.caches.outputs)494            current_outputs = self.caches.outputs.all_node_ids()495            for node_id in list(execute_outputs):496                execution_list.add_node(node_id)497 498            while not execution_list.is_empty():499                node_id, error, ex = execution_list.stage_node_execution()500                if error is not None:501                    self.handle_execution_error(prompt_id, dynamic_prompt.original_prompt, current_outputs, executed, error, ex)502                    break503 504                result, error, ex = execute(self.server, dynamic_prompt, self.caches, node_id, extra_data, executed, prompt_id, execution_list, pending_subgraph_results)505                self.success = result != ExecutionResult.FAILURE506                if result == ExecutionResult.FAILURE:507                    self.handle_execution_error(prompt_id, dynamic_prompt.original_prompt, current_outputs, executed, error, ex)508                    break509                elif result == ExecutionResult.PENDING:510                    execution_list.unstage_node_execution()511                else: # result == ExecutionResult.SUCCESS:512                    execution_list.complete_node_execution()513            else:514                # Only execute when the while-loop ends without break515                self.add_message("execution_success", { "prompt_id": prompt_id }, broadcast=False)516 517            ui_outputs = {}518            meta_outputs = {}519            all_node_ids = self.caches.ui.all_node_ids()520            for node_id in all_node_ids:521                ui_info = self.caches.ui.get(node_id)522                if ui_info is not None:523                    ui_outputs[node_id] = ui_info["output"]524                    meta_outputs[node_id] = ui_info["meta"]525            self.history_result = {526                "outputs": ui_outputs,527                "meta": meta_outputs,528            }529            self.server.last_node_id = None530            if comfy.model_management.DISABLE_SMART_MEMORY:531                comfy.model_management.unload_all_models()532 533 534def validate_inputs(prompt, item, validated):535    unique_id = item536    if unique_id in validated:537        return validated[unique_id]538 539    inputs = prompt[unique_id]['inputs']540    class_type = prompt[unique_id]['class_type']541    obj_class = nodes.NODE_CLASS_MAPPINGS[class_type]542 543    class_inputs = obj_class.INPUT_TYPES()544    valid_inputs = set(class_inputs.get('required',{})).union(set(class_inputs.get('optional',{})))545 546    errors = []547    valid = True548 549    validate_function_inputs = []550    validate_has_kwargs = False551    if hasattr(obj_class, "VALIDATE_INPUTS"):552        argspec = inspect.getfullargspec(obj_class.VALIDATE_INPUTS)553        validate_function_inputs = argspec.args554        validate_has_kwargs = argspec.varkw is not None555    received_types = {}556 557    for x in valid_inputs:558        type_input, input_category, extra_info = get_input_info(obj_class, x, class_inputs)559        assert extra_info is not None560        if x not in inputs:561            if input_category == "required":562                error = {563                    "type": "required_input_missing",564                    "message": "Required input is missing",565                    "details": f"{x}",566                    "extra_info": {567                        "input_name": x568                    }569                }570                errors.append(error)571            continue572 573        val = inputs[x]574        info = (type_input, extra_info)575        if isinstance(val, list):576            if len(val) != 2:577                error = {578                    "type": "bad_linked_input",579                    "message": "Bad linked input, must be a length-2 list of [node_id, slot_index]",580                    "details": f"{x}",581                    "extra_info": {582                        "input_name": x,583                        "input_config": info,584                        "received_value": val585                    }586                }587                errors.append(error)588                continue589 590            o_id = val[0]591            o_class_type = prompt[o_id]['class_type']592            r = nodes.NODE_CLASS_MAPPINGS[o_class_type].RETURN_TYPES593            received_type = r[val[1]]594            received_types[x] = received_type595            if 'input_types' not in validate_function_inputs and not validate_node_input(received_type, type_input):596                details = f"{x}, received_type({received_type}) mismatch input_type({type_input})"597                error = {598                    "type": "return_type_mismatch",599                    "message": "Return type mismatch between linked nodes",600                    "details": details,601                    "extra_info": {602                        "input_name": x,603                        "input_config": info,604                        "received_type": received_type,605                        "linked_node": val606                    }607                }608                errors.append(error)609                continue610            try:611                r = validate_inputs(prompt, o_id, validated)612                if r[0] is False:613                    # `r` will be set in `validated[o_id]` already614                    valid = False615                    continue616            except Exception as ex:617                typ, _, tb = sys.exc_info()618                valid = False619                exception_type = full_type_name(typ)620                reasons = [{621                    "type": "exception_during_inner_validation",622                    "message": "Exception when validating inner node",623                    "details": str(ex),624                    "extra_info": {625                        "input_name": x,626                        "input_config": info,627                        "exception_message": str(ex),628                        "exception_type": exception_type,629                        "traceback": traceback.format_tb(tb),630                        "linked_node": val631                    }632                }]633                validated[o_id] = (False, reasons, o_id)634                continue635        else:636            try:637                if type_input == "INT":638                    val = int(val)639                    inputs[x] = val640                if type_input == "FLOAT":641                    val = float(val)642                    inputs[x] = val643                if type_input == "STRING":644                    val = str(val)645                    inputs[x] = val646                if type_input == "BOOLEAN":647                    val = bool(val)648                    inputs[x] = val649            except Exception as ex:650                error = {651                    "type": "invalid_input_type",652                    "message": f"Failed to convert an input value to a {type_input} value",653                    "details": f"{x}, {val}, {ex}",654                    "extra_info": {655                        "input_name": x,656                        "input_config": info,657                        "received_value": val,658                        "exception_message": str(ex)659                    }660                }661                errors.append(error)662                continue663 664            if x not in validate_function_inputs and not validate_has_kwargs:665                if "min" in extra_info and val < extra_info["min"]:666                    error = {667                        "type": "value_smaller_than_min",668                        "message": "Value {} smaller than min of {}".format(val, extra_info["min"]),669                        "details": f"{x}",670                        "extra_info": {671                            "input_name": x,672                            "input_config": info,673                            "received_value": val,674                        }675                    }676                    errors.append(error)677                    continue678                if "max" in extra_info and val > extra_info["max"]:679                    error = {680                        "type": "value_bigger_than_max",681                        "message": "Value {} bigger than max of {}".format(val, extra_info["max"]),682                        "details": f"{x}",683                        "extra_info": {684                            "input_name": x,685                            "input_config": info,686                            "received_value": val,687                        }688                    }689                    errors.append(error)690                    continue691 692                if isinstance(type_input, list):693                    if val not in type_input:694                        input_config = info695                        list_info = ""696 697                        # Don't send back gigantic lists like if they're lots of698                        # scanned model filepaths699                        if len(type_input) > 20:700                            list_info = f"(list of length {len(type_input)})"701                            input_config = None702                        else:703                            list_info = str(type_input)704 705                        error = {706                            "type": "value_not_in_list",707                            "message": "Value not in list",708                            "details": f"{x}: '{val}' not in {list_info}",709                            "extra_info": {710                                "input_name": x,711                                "input_config": input_config,712                                "received_value": val,713                            }714                        }715                        errors.append(error)716                        continue717 718    if len(validate_function_inputs) > 0 or validate_has_kwargs:719        input_data_all, _ = get_input_data(inputs, obj_class, unique_id)720        input_filtered = {}721        for x in input_data_all:722            if x in validate_function_inputs or validate_has_kwargs:723                input_filtered[x] = input_data_all[x]724        if 'input_types' in validate_function_inputs:725            input_filtered['input_types'] = [received_types]726 727        #ret = obj_class.VALIDATE_INPUTS(**input_filtered)728        ret = _map_node_over_list(obj_class, input_filtered, "VALIDATE_INPUTS")729        for x in input_filtered:730            for i, r in enumerate(ret):731                if r is not True and not isinstance(r, ExecutionBlocker):732                    details = f"{x}"733                    if r is not False:734                        details += f" - {str(r)}"735 736                    error = {737                        "type": "custom_validation_failed",738                        "message": "Custom validation failed for node",739                        "details": details,740                        "extra_info": {741                            "input_name": x,742                        }743                    }744                    errors.append(error)745                    continue746 747    if len(errors) > 0 or valid is not True:748        ret = (False, errors, unique_id)749    else:750        ret = (True, [], unique_id)751 752    validated[unique_id] = ret753    return ret754 755def full_type_name(klass):756    module = klass.__module__757    if module == 'builtins':758        return klass.__qualname__759    return module + '.' + klass.__qualname__760 761def validate_prompt(prompt):762    outputs = set()763    for x in prompt:764        if 'class_type' not in prompt[x]:765            error = {766                "type": "invalid_prompt",767                "message": "Cannot execute because a node is missing the class_type property.",768                "details": f"Node ID '#{x}'",769                "extra_info": {}770            }771            return (False, error, [], [])772 773        class_type = prompt[x]['class_type']774        class_ = nodes.NODE_CLASS_MAPPINGS.get(class_type, None)775        if class_ is None:776            error = {777                "type": "invalid_prompt",778                "message": f"Cannot execute because node {class_type} does not exist.",779                "details": f"Node ID '#{x}'",780                "extra_info": {}781            }782            return (False, error, [], [])783 784        if hasattr(class_, 'OUTPUT_NODE') and class_.OUTPUT_NODE is True:785            outputs.add(x)786 787    if len(outputs) == 0:788        error = {789            "type": "prompt_no_outputs",790            "message": "Prompt has no outputs",791            "details": "",792            "extra_info": {}793        }794        return (False, error, [], [])795 796    good_outputs = set()797    errors = []798    node_errors = {}799    validated = {}800    for o in outputs:801        valid = False802        reasons = []803        try:804            m = validate_inputs(prompt, o, validated)805            valid = m[0]806            reasons = m[1]807        except Exception as ex:808            typ, _, tb = sys.exc_info()809            valid = False810            exception_type = full_type_name(typ)811            reasons = [{812                "type": "exception_during_validation",813                "message": "Exception when validating node",814                "details": str(ex),815                "extra_info": {816                    "exception_type": exception_type,817                    "traceback": traceback.format_tb(tb)818                }819            }]820            validated[o] = (False, reasons, o)821 822        if valid is True:823            good_outputs.add(o)824        else:825            logging.error(f"Failed to validate prompt for output {o}:")826            if len(reasons) > 0:827                logging.error("* (prompt):")828                for reason in reasons:829                    logging.error(f"  - {reason['message']}: {reason['details']}")830            errors += [(o, reasons)]831            for node_id, result in validated.items():832                valid = result[0]833                reasons = result[1]834                # If a node upstream has errors, the nodes downstream will also835                # be reported as invalid, but there will be no errors attached.836                # So don't return those nodes as having errors in the response.837                if valid is not True and len(reasons) > 0:838                    if node_id not in node_errors:839                        class_type = prompt[node_id]['class_type']840                        node_errors[node_id] = {841                            "errors": reasons,842                            "dependent_outputs": [],843                            "class_type": class_type844                        }845                        logging.error(f"* {class_type} {node_id}:")846                        for reason in reasons:847                            logging.error(f"  - {reason['message']}: {reason['details']}")848                    node_errors[node_id]["dependent_outputs"].append(o)849            logging.error("Output will be ignored")850 851    if len(good_outputs) == 0:852        errors_list = []853        for o, errors in errors:854            for error in errors:855                errors_list.append(f"{error['message']}: {error['details']}")856        errors_list = "\n".join(errors_list)857 858        error = {859            "type": "prompt_outputs_failed_validation",860            "message": "Prompt outputs failed validation",861            "details": errors_list,862            "extra_info": {}863        }864 865        return (False, error, list(good_outputs), node_errors)866 867    return (True, None, list(good_outputs), node_errors)868 869MAXIMUM_HISTORY_SIZE = 10000870 871class PromptQueue:872    def __init__(self, server):873        self.server = server874        self.mutex = threading.RLock()875        self.not_empty = threading.Condition(self.mutex)876        self.task_counter = 0877        self.queue = []878        self.currently_running = {}879        self.history = {}880        self.flags = {}881        server.prompt_queue = self882 883    def put(self, item):884        with self.mutex:885            heapq.heappush(self.queue, item)886            self.server.queue_updated()887            self.not_empty.notify()888 889    def get(self, timeout=None):890        with self.not_empty:891            while len(self.queue) == 0:892                self.not_empty.wait(timeout=timeout)893                if timeout is not None and len(self.queue) == 0:894                    return None895            item = heapq.heappop(self.queue)896            i = self.task_counter897            self.currently_running[i] = copy.deepcopy(item)898            self.task_counter += 1899            self.server.queue_updated()900            return (item, i)901 902    class ExecutionStatus(NamedTuple):903        status_str: Literal['success', 'error']904        completed: bool905        messages: List[str]906 907    def task_done(self, item_id, history_result,908                  status: Optional['PromptQueue.ExecutionStatus']):909        with self.mutex:910            prompt = self.currently_running.pop(item_id)911            if len(self.history) > MAXIMUM_HISTORY_SIZE:912                self.history.pop(next(iter(self.history)))913 914            status_dict: Optional[dict] = None915            if status is not None:916                status_dict = copy.deepcopy(status._asdict())917 918            self.history[prompt[1]] = {919                "prompt": prompt,920                "outputs": {},921                'status': status_dict,922            }923            self.history[prompt[1]].update(history_result)924            self.server.queue_updated()925 926    def get_current_queue(self):927        with self.mutex:928            out = []929            for x in self.currently_running.values():930                out += [x]931            return (out, copy.deepcopy(self.queue))932 933    def get_tasks_remaining(self):934        with self.mutex:935            return len(self.queue) + len(self.currently_running)936 937    def wipe_queue(self):938        with self.mutex:939            self.queue = []940            self.server.queue_updated()941 942    def delete_queue_item(self, function):943        with self.mutex:944            for x in range(len(self.queue)):945                if function(self.queue[x]):946                    if len(self.queue) == 1:947                        self.wipe_queue()948                    else:949                        self.queue.pop(x)950                        heapq.heapify(self.queue)951                    self.server.queue_updated()952                    return True953        return False954 955    def get_history(self, prompt_id=None, max_items=None, offset=-1):956        with self.mutex:957            if prompt_id is None:958                out = {}959                i = 0960                if offset < 0 and max_items is not None:961                    offset = len(self.history) - max_items962                for k in self.history:963                    if i >= offset:964                        out[k] = self.history[k]965                        if max_items is not None and len(out) >= max_items:966                            break967                    i += 1968                return out969            elif prompt_id in self.history:970                return {prompt_id: copy.deepcopy(self.history[prompt_id])}971            else:972                return {}973 974    def wipe_history(self):975        with self.mutex:976            self.history = {}977 978    def delete_history_item(self, id_to_delete):979        with self.mutex:980            self.history.pop(id_to_delete, None)981 982    def set_flag(self, name, data):983        with self.mutex:984            self.flags[name] = data985            self.not_empty.notify()986 987    def get_flags(self, reset=True):988        with self.mutex:989            if reset:990                ret = self.flags991                self.flags = {}992                return ret993            else:994                return self.flags.copy()995