From 740caecda5351cdd03dc9ac5be483f373ba3f465 Mon Sep 17 00:00:00 2001 From: Christian Byrne Date: Mon, 28 Apr 2025 21:52:45 -0700 Subject: [PATCH] Refactor util functions (#20) --- comfy_api_nodes/nodes_api.py | 122 +++++++++++++++++++++++++++++++---- 1 file changed, 111 insertions(+), 11 deletions(-) diff --git a/comfy_api_nodes/nodes_api.py b/comfy_api_nodes/nodes_api.py index 359c2cb7c..a935465d8 100644 --- a/comfy_api_nodes/nodes_api.py +++ b/comfy_api_nodes/nodes_api.py @@ -1,7 +1,7 @@ import io from inspect import cleandoc from comfy.comfy_types.node_typing import FileLocator -from typing import Literal +from typing import Literal, Optional from comfy.utils import common_upscale from comfy.comfy_types.node_typing import IO, ComfyNodeABC, InputTypeDict from comfy_api_nodes.apis import ( @@ -49,6 +49,7 @@ import os import time import uuid import folder_paths +from io import BytesIO def downscale_input(image, total_pixels=1536*1024): samples = image.movedim(-1,1) @@ -121,15 +122,54 @@ def validate_aspect_ratio(aspect_ratio: str, minimum_ratio: float, maximum_ratio raise Exception(f"Aspect ratio cannot reduce to any greater than {maximum_ratio_str} ({maximum_ratio}), but was {aspect_ratio} ({calculated_ratio}).") return aspect_ratio -def process_image_response(response: requests.Response): - '''Uses content from a Response object and converts it to a torch.Tensor''' - image = Image.open(io.BytesIO(response.content)).convert("RGBA") + +def mimetype_to_extension(mime_type: str) -> str: + """Converts a MIME type to a file extension.""" + return mime_type.split('/')[-1].lower() + + +def download_url_to_bytesio(url: str, timeout: int = None) -> BytesIO: + """Downloads content from a URL using requests and returns it as BytesIO. + + Args: + url: The URL to download. + timeout: Request timeout in seconds. Defaults to None (no timeout). + + Returns: + BytesIO object containing the downloaded content. + """ + response = requests.get(url, stream=True, timeout=timeout) + response.raise_for_status() # Raises HTTPError for bad responses (4XX or 5XX) + return BytesIO(response.content) + + +def bytesio_to_image_tensor(image_bytesio: BytesIO, mode: str = "RGBA") -> torch.Tensor: + """Converts image data from BytesIO to a torch.Tensor. + + Args: + image_bytesio: BytesIO object containing the image data. + mode: The PIL mode to convert the image to (e.g., "RGB", "RGBA"). + + Returns: + A torch.Tensor representing the image (1, H, W, C). + + Raises: + PIL.UnidentifiedImageError: If the image data cannot be identified. + ValueError: If the specified mode is invalid. + """ + image = Image.open(image_bytesio) + image = image.convert(mode) image_array = np.array(image).astype(np.float32) / 255.0 return torch.from_numpy(image_array).unsqueeze(0) -def convert_image_to_bytesio(image: torch.Tensor, name: str=None, allow_alpha=True, total_pixels=2048*2048): - img_binary = None - # only care about first image, if it is a batch + +def process_image_response(response: requests.Response): + '''Uses content from a Response object and converts it to a torch.Tensor''' + return bytesio_to_image_tensor(BytesIO(response.content)) + + +def _tensor_to_pil(image: torch.Tensor, total_pixels: int = 2048*2048) -> Image.Image: + """Converts a single torch.Tensor image [H, W, C] to a PIL Image, optionally downscaling.""" if len(image.shape) > 3: image = image[0] # TODO: remove alpha if not allowed and present @@ -137,13 +177,73 @@ def convert_image_to_bytesio(image: torch.Tensor, name: str=None, allow_alpha=Tr input_tensor = downscale_input(input_tensor.unsqueeze(0), total_pixels=total_pixels).squeeze() image_np = (input_tensor.numpy() * 255).astype(np.uint8) img = Image.fromarray(image_np) + return img + + +def _pil_to_bytesio(img: Image.Image, mime_type: str = 'image/png') -> BytesIO: + """Converts a PIL Image to a BytesIO object.""" img_byte_arr = io.BytesIO() - img.save(img_byte_arr, format='PNG') + # Derive PIL format from MIME type (e.g., 'image/png' -> 'PNG') + pil_format = mime_type.split('/')[-1].upper() + if pil_format == 'JPG': + pil_format = 'JPEG' + img.save(img_byte_arr, format=pil_format) img_byte_arr.seek(0) - img_binary = img_byte_arr - img_binary.name = f"{name if name else uuid.uuid4()}.png" + return img_byte_arr + + +def tensor_to_bytesio(image: torch.Tensor, name: Optional[str] = None, total_pixels: int = 2048*2048, mime_type: str = 'image/png') -> BytesIO: + """Converts a torch.Tensor image to a named BytesIO object. + + Args: + image: Input torch.Tensor image. + name: Optional filename for the BytesIO object. + total_pixels: Maximum total pixels for potential downscaling. + mime_type: Target image MIME type (e.g., 'image/png', 'image/jpeg', 'image/webp', 'video/mp4'). + + Returns: + Named BytesIO object containing the image data. + """ + pil_image = _tensor_to_pil(image, total_pixels=total_pixels) + img_binary = _pil_to_bytesio(pil_image, mime_type=mime_type) + img_binary.name = f"{name if name else uuid.uuid4()}.{mimetype_to_extension(mime_type)}" return img_binary + +def tensor_to_base64_string(image_tensor: torch.Tensor, total_pixels: int = 2048*2048, mime_type: str = 'image/png') -> str: + """Convert [B, H, W, C] or [H, W, C] tensor to a base64 string. + + Args: + image_tensor: Input torch.Tensor image. + total_pixels: Maximum total pixels for potential downscaling. + mime_type: Target image MIME type (e.g., 'image/png', 'image/jpeg', 'image/webp', 'video/mp4'). + + Returns: + Base64 encoded string of the image. + """ + pil_image = _tensor_to_pil(image_tensor, total_pixels=total_pixels) + img_byte_arr = _pil_to_bytesio(pil_image, mime_type=mime_type) + img_bytes = img_byte_arr.getvalue() + # Encode bytes to base64 string + base64_encoded_string = base64.b64encode(img_bytes).decode("utf-8") + return base64_encoded_string + + +def tensor_to_data_uri(image_tensor: torch.Tensor, total_pixels: int = 2048 * 2048, mime_type: str = 'image/png') -> str: + """Converts a tensor image to a Data URI string. + + Args: + image_tensor: Input torch.Tensor image. + total_pixels: Maximum total pixels for potential downscaling. + mime_type: Target image MIME type (e.g., 'image/png', 'image/jpeg', 'image/webp'). + + Returns: + Data URI string (e.g., 'data:image/png;base64,...'). + """ + base64_string = tensor_to_base64_string(image_tensor, total_pixels, mime_type) + return f"data:{mime_type};base64,{base64_string}" + + def upload_images_to_comfyapi(image: torch.Tensor, max_images=8, auth_token=None) -> list[str]: # if batch, try to upload each file if max_images is greater than 0 idx_image = 0 @@ -157,7 +257,7 @@ def upload_images_to_comfyapi(image: torch.Tensor, max_images=8, auth_token=None if len(image.shape) > 3: curr_image = image[idx_image] # get BytesIO version of image - img_binary = convert_image_to_bytesio(curr_image) + img_binary = tensor_to_bytesio(curr_image) # first, request upload/download urls from comfy API operation = SynchronousOperation( endpoint=ApiEndpoint(