"""Workflow-compatible implementations of the ETK and SWAIN list selectors.""" from __future__ import annotations from collections.abc import Mapping from copy import deepcopy class AnyType(str): """ComfyUI wildcard socket type for a list whose item type is not fixed.""" def __ne__(self, other): return False ANY_TYPE = AnyType("*") class ETKListToComfyList: @classmethod def INPUT_TYPES(cls): return { "required": { "items": ("LIST", {"default": []}), } } RETURN_TYPES = (ANY_TYPE,) RETURN_NAMES = ("items",) OUTPUT_IS_LIST = (True,) FUNCTION = "convert" CATEGORY = "ETK/List" def convert(self, items): if not isinstance(items, list): raise TypeError(f"items must be a LIST, got {type(items).__name__}") return (items,) class ImageFromListIdx: @classmethod def INPUT_TYPES(cls): return { "required": { "images": ("LIST", {"default": []}), "index": ("INT", {"default": 0, "min": 0, "max": 10e10}), } } RETURN_TYPES = ("IMAGE",) FUNCTION = "get_image" CATEGORY = "ETK/Image" def get_image(self, images, index): return (images[index],) class LatentFromListIdx: @classmethod def INPUT_TYPES(cls): return { "required": { "latents": ("LIST", {"default": []}), "index": ("INT", {"default": 0, "min": -10e10, "max": 10e10}), } } RETURN_TYPES = ("LATENT",) FUNCTION = "get_latent" CATEGORY = "ETK/Image" def get_latent(self, latents, index): return (latents[index],) class ConditioningFromListIdx: @classmethod def INPUT_TYPES(cls): return { "required": { "conditionings": ("LIST", {"default": []}), "index": ("INT", {"default": 0, "min": 0, "max": 10e10}), } } RETURN_TYPES = ("CONDITIONING",) FUNCTION = "get_conditioning" CATEGORY = "ETK/Image" def get_conditioning(self, conditionings, index): try: return (conditionings[index],) except IndexError as exc: raise IndexError( f"ConditioningFromListIdx index {index} out of range for list length {len(conditionings)}" ) from exc def _normalize_comfy_audio(audio): import torch if audio is None: raise ValueError("audio input must be provided") if not isinstance(audio, Mapping): raise ValueError("AUDIO input must be a mapping with 'waveform' and 'sample_rate' keys") try: waveform = audio["waveform"] except KeyError as exc: raise ValueError("AUDIO input must contain a 'waveform' key") from exc if not torch.is_tensor(waveform): waveform = torch.tensor(waveform) if waveform.dim() == 1: waveform = waveform.unsqueeze(0).unsqueeze(0) if waveform.dim() == 2: waveform = waveform.unsqueeze(0) if waveform.dim() != 3: raise ValueError(f"AUDIO waveform must have shape [B, C, T], got {tuple(waveform.shape)}") sample_rate = int(audio.get("sample_rate", 44100) or 44100) if sample_rate <= 0: raise ValueError(f"AUDIO sample_rate must be positive, got {sample_rate}") audio_dict = dict(audio) audio_dict["waveform"] = waveform audio_dict["sample_rate"] = sample_rate return audio_dict class AudioFromListIdx: @classmethod def INPUT_TYPES(cls): return { "required": { "audios": ("LIST", {"default": []}), "index": ("INT", {"default": 0, "min": -10e10, "max": 10e10}), } } RETURN_TYPES = ("AUDIO",) FUNCTION = "get_audio" CATEGORY = "ETK/audio" def get_audio(self, audios, index): return (_normalize_comfy_audio(audios[int(index)]),) def _swain_input_types(*, bounded_index: bool): index_config = {"default": 0} if bounded_index: index_config.update({"min": -1000, "max": 1000}) return { "required": {"list": ("LIST", {"default": []})}, "optional": {"int": ("INT", index_config)}, } def _selected(kwargs): values = kwargs.get("list", None) index = kwargs.get("int", None) if values is None: raise ValueError("Must provide an input") if not isinstance(values, list): raise ValueError("input must be a list") if not isinstance(index, int): raise ValueError("index must be an int") if index >= len(values): raise ValueError("index must be in the range of the list") return values, values[index] class DictFromListIdx: @classmethod def INPUT_TYPES(cls): return _swain_input_types(bounded_index=True) RETURN_TYPES = ("DICT",) RETURN_NAMES = ("dictionary",) FUNCTION = "handler" CATEGORY = "SWAIN/text" def handler(self, **kwargs): _, value = _selected(deepcopy(kwargs)) return (value,) class StrFromListIdx: @classmethod def INPUT_TYPES(cls): return _swain_input_types(bounded_index=True) RETURN_TYPES = ("STRING", "LIST") RETURN_NAMES = ("string", "list") FUNCTION = "handler" CATEGORY = "SWAIN/text" def handler(self, **kwargs): values, value = _selected(deepcopy(kwargs)) return value, values class BytesFromListIdx: @classmethod def INPUT_TYPES(cls): return _swain_input_types(bounded_index=False) RETURN_TYPES = ("BYTES",) RETURN_NAMES = ("bytes",) FUNCTION = "handler" CATEGORY = "SWAIN/text" def handler(self, **kwargs): _, value = _selected(deepcopy(kwargs)) if isinstance(value, list): if len(value) > 1: raise ValueError("bytes must be a list of length 1") value = value[0] return (value,) class IntFromListIdx: @classmethod def INPUT_TYPES(cls): return _swain_input_types(bounded_index=False) RETURN_TYPES = ("INT", "LIST") RETURN_NAMES = ("int", "list") FUNCTION = "handler" CATEGORY = "SWAIN/text" def handler(self, **kwargs): values, value = _selected(deepcopy(kwargs)) if not isinstance(value, int): try: value = int(value) except Exception as exc: raise ValueError( "list at idx must be an int or a string that can be converted to an int" ) from exc return value, values class FloatFromListIdx: @classmethod def INPUT_TYPES(cls): return _swain_input_types(bounded_index=False) RETURN_TYPES = ("FLOAT", "LIST") RETURN_NAMES = ("float", "list") FUNCTION = "handler" CATEGORY = "SWAIN/text" def handler(self, **kwargs): values, value = _selected(deepcopy(kwargs)) if not isinstance(value, float): try: value = float(value) except Exception as exc: raise ValueError( "list at idx must be an int or a string that can be converted to an int" ) from exc return value, values