diff --git a/README.md b/README.md index 5036558..942cc46 100644 --- a/README.md +++ b/README.md @@ -44,11 +44,26 @@ That is two 32-bit integers (big endian) with values 1 and 2 followed by the PNG {'type': 'executed', 'data': {'node': '', 'output': {'images': [{'source': 'websocket', 'content-type': 'image/png', 'type': 'output'}, ...]}, 'prompt_id': '}} ``` -### Send Image (HTTP) +### Load Image from Cache + +Loads an image or mask that has been uploaded previously into the workflow. +Uploaded images are temporarily stored in RAM rather than written to disk. This +method has less overhead compared to embedding images as base64 into the prompt, +but is more complex to implement. +* Inputs: id of an image that was uploaded previously +* Outputs: image (RGB) and mask (A of RGBA input, or first channel if no alpha present). + +To upload an image, upload the _bytes_ of a PNG via a HTTP PUT request to +`/api/etn/image/{id}`. JPEG or other formats also work. Choose any `id` which +does not clash with other images you upload, and reference it in the node. The +request returns `201` if the image was uploaded and `200` if it was already +cached. + +### Save Image to Cache Stores an output image in RAM temporarily and allows retrieval over HTTP. This is typically faster than WebSocket, especially for large images. -* Inputs: the image (RGB or RGBA), supports batches +* Inputs: the image (RGB or RGBA). Batches are supported. This node will send a JSON message over WebSocket when an image is ready: ```json @@ -66,8 +81,8 @@ This node will send a JSON message over WebSocket when an image is ready: } ``` -To download the images, send a HTTP GET request to `/api/etn/image/{id}` with the image IDs from the message. -Images will be cached for a few minutes. +To download the images, send a HTTP GET request to `/api/etn/image/{id}` with +the image IDs from the message. Images will be cached for a few minutes. ## Regions diff --git a/__init__.py b/__init__.py index 3689797..8ee7bc2 100644 --- a/__init__.py +++ b/__init__.py @@ -5,10 +5,11 @@ from . import api as api, nodes, tile, region, nsfw, translation, krita class ExternalToolingNodes(ComfyExtension): async def get_node_list(self) -> list[type[io.ComfyNode]]: return [ + nodes.LoadImageCache, + nodes.SaveImageCache, nodes.LoadImageBase64, nodes.LoadMaskBase64, nodes.SendImageWebSocket, - nodes.SendImageHTTP, nodes.ApplyMaskToImage, nodes.ReferenceImage, nodes.ApplyReferenceImages, diff --git a/api.py b/api.py index 5bef9a3..5c3accd 100644 --- a/api.py +++ b/api.py @@ -301,6 +301,23 @@ if _server is not None: except Exception as e: return web.json_response(dict(error=str(e)), status=500) + @_server.routes.put("/api/etn/image/{id}") + async def put_image(request: web.Request): + try: + id = request.match_info.get("id", "") + if id in image_cache: + return web.json_response(dict(status="cached"), status=200) + + content_type = request.headers.get("Content-Type", "application/octet-stream") + data = bytearray() + async for chunk, _ in request.content.iter_chunks(): + data.extend(chunk) + + image_cache.insert(id, bytes(data), content_type) + return web.json_response(dict(status="success"), status=201) + except Exception as e: + return web.json_response(dict(error=str(e)), status=500) + @_server.routes.put("/api/etn/upload/{folder_name}/{filename}") async def upload(request: web.Request): folder_name = request.match_info.get("folder_name", "") diff --git a/nodes.py b/nodes.py index 88ab694..9fde8bd 100644 --- a/nodes.py +++ b/nodes.py @@ -123,20 +123,25 @@ class ImageCache: image.save(output, format=format, quality=95, compress_level=1) image_data = output.getvalue() + self.insert(key, image_data, f"image/{format.lower()}") + return key + + def insert(self, key: str, data: bytes, content_type: str): self.images[key] = ImageCache.Entry( - data=image_data, - content_type=f"image/{format.lower()}", + data=data, + content_type=content_type, timestamp=time.time(), retrieved=0, ) - return key - def get(self, key: str): + def get(self, key: str, extend: bool = False): entry = self.images.get(key) if entry is None: return None, None - self.prune() entry.retrieved += 1 + if extend: + entry.timestamp = time.time() + self.prune() return entry.data, entry.content_type def prune(self): @@ -149,16 +154,55 @@ class ImageCache: for key in keys_to_delete: del self.images[key] + def __contains__(self, key: str): + return key in self.images + image_cache = ImageCache() -class SendImageHTTP(io.ComfyNode): +class LoadImageCache(io.ComfyNode): @classmethod def define_schema(cls): return io.Schema( - node_id="ETN_SendImageHTTP", - display_name="Send Image (HTTP)", + node_id="ETN_LoadImageCache", + display_name="Load Image from Cache", + category="external_tooling", + inputs=[io.String.Input("id", multiline=False)], + outputs=[io.Image.Output(display_name="image"), io.Mask.Output(display_name="mask")], + ) + + @classmethod + def execute(cls, id: str): + image_data, content_type = image_cache.get(id, extend=True) + if image_data is None: + raise ValueError(f"Image with ID {id} not found in cache.") + + img = Image.open(BytesIO(image_data)) + w, h = img.size + c = len(img.getbands()) + normalized = np.array(img).astype(np.float32) / 255.0 + tensor = torch.from_numpy(normalized).reshape(1, h, w, c) + match c: + case 1: + image = tensor.expand(1, h, w, 3) + mask = tensor.reshape(1, h, w) + case 3: + image = tensor + mask = tensor[..., 0] + case 4: + image = tensor[..., :3] + mask = tensor[..., 3] + + return io.NodeOutput(image, mask) + + +class SaveImageCache(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="ETN_SaveImageCache", + display_name="Save Image to Cache", category="external_tooling", inputs=[ io.Image.Input("images"), diff --git a/pyproject.toml b/pyproject.toml index 40b7903..6106704 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "comfyui-tooling-nodes" description = "Provides nodes and server API extensions geared towards using ComfyUI as a backend for external tools." -version = "3.0.0" +version = "2.0.6" license = { file = "LICENSE" } [project.urls]