diff --git a/__init__.py b/__init__.py index 6e2ad38..f3dea1e 100644 --- a/__init__.py +++ b/__init__.py @@ -347,6 +347,267 @@ if hasattr(PromptServer, "instance"): # Return JSON for other requests return web.json_response({"message": "Welcome to MTB!"}) + import asyncio + import os + from functools import lru_cache + from io import BytesIO + + from aiohttp import web + from PIL import Image + + # In-memory cache for memoization + @lru_cache(maxsize=256) + def get_cached_image(file_path, preview_params=None, channel=None): + with Image.open(file_path) as img: + if preview_params: + img = process_preview(img, preview_params) + if channel: + img = process_channel(img, channel) + return img + + def process_preview(img: Image, preview_params): + image_format, quality, width = preview_params + quality = int(quality) + + if width: + width = int(width) + img.thumbnail((width, int(width * img.height / img.width))) + + buffer = BytesIO() + img.save(buffer, format=image_format, quality=quality) + buffer.seek(0) + return buffer + + def process_channel(img, channel): + if channel == "rgb": + if img.mode == "RGBA": + r, g, b, _ = img.split() + img = Image.merge("RGB", (r, g, b)) + else: + img = img.convert("RGB") + elif channel == "a": + if img.mode == "RGBA": + _, _, _, a = img.split() + else: + a = Image.new("L", img.size, 255) + img = Image.new("RGBA", img.size) + img.putalpha(a) + + buffer = BytesIO() + img.save(buffer, format="PNG") + buffer.seek(0) + return buffer + + async def get_image_response( + file, filename, preview_info=None, channel=None + ): + img = await asyncio.to_thread( + get_cached_image, file, preview_info, channel + ) + return web.Response( + body=img.read(), + content_type="image/webp" if preview_info else "image/png", + headers={"Content-Disposition": f'filename="{filename}"'}, + ) + + @PromptServer.instance.routes.get("/mtb/view") + async def view_image(request): + import folder_paths + + filename = request.rel_url.query.get("filename") + if not filename: + return web.Response(status=404) + + filename, output_dir = folder_paths.annotated_filepath(filename) + if filename[0] == "/" or ".." in filename: + return web.Response(status=400) + + if output_dir is None: + type = request.rel_url.query.get("type", "output") + output_dir = folder_paths.get_directory_by_type(type) + + if output_dir is None: + return web.Response(status=400) + + if "subfolder" in request.rel_url.query: + full_output_dir = os.path.join( + output_dir, request.rel_url.query["subfolder"] + ) + if ( + os.path.commonpath( + (os.path.abspath(full_output_dir), output_dir) + ) + != output_dir + ): + return web.Response(status=403) + output_dir = full_output_dir + + filename = os.path.basename(filename) + file = os.path.join(output_dir, filename) + + if not os.path.isfile(file): + return web.Response(status=404) + + preview_info = None + if "preview" in request.rel_url.query: + preview_params = request.rel_url.query["preview"].split(";") + image_format = ( + preview_params[0] + if preview_params[0] in ["webp", "jpeg"] + else "webp" + ) + quality = ( + int(preview_params[1]) + if len(preview_params) > 1 and preview_params[1].isdigit() + else 90 + ) + width = request.rel_url.query.get("width") + preview_info = (image_format, quality, width) + + channel = request.rel_url.query.get("channel") + + return await get_image_response(file, filename, preview_info, channel) + + # + # @PromptServer.instance.routes.get("/mtb/view") + # async def view_image(request): + # from io import BytesIO + # + # import folder_paths + # from PIL import Image + # + # if "filename" in request.rel_url.query: + # filename = request.rel_url.query["filename"] + # filename, output_dir = folder_paths.annotated_filepath(filename) + # + # # validation for security: prevent accessing arbitrary path + # if filename[0] == "/" or ".." in filename: + # return web.Response(status=400) + # + # if output_dir is None: + # type = request.rel_url.query.get("type", "output") + # output_dir = folder_paths.get_directory_by_type(type) + # + # if output_dir is None: + # return web.Response(status=400) + # + # if "subfolder" in request.rel_url.query: + # full_output_dir = os.path.join( + # output_dir, request.rel_url.query["subfolder"] + # ) + # if ( + # os.path.commonpath( + # (os.path.abspath(full_output_dir), output_dir) + # ) + # != output_dir + # ): + # return web.Response(status=403) + # output_dir = full_output_dir + # + # filename = os.path.basename(filename) + # file = os.path.join(output_dir, filename) + # + # if os.path.isfile(file): + # if "preview" in request.rel_url.query: + # with Image.open(file) as img: + # preview_info = request.rel_url.query["preview"].split( + # ";" + # ) + # image_format = preview_info[0] + # if image_format not in [ + # "webp", + # "jpeg", + # ] or "a" in request.rel_url.query.get("channel", ""): + # image_format = "webp" + # + # quality = 90 + # if preview_info[-1].isdigit(): + # quality = int(preview_info[-1]) + # + # width = request.rel_url.query.get("width") + # if width is not None: + # width = int(width) + # img.resize( + # ( + # width, + # int(width * img.height / img.width), + # ) + # ) + # + # buffer = BytesIO() + # if ( + # image_format in ["jpeg"] + # or request.rel_url.query.get("channel", "") + # == "rgb" + # ): + # img = img.convert("RGB") + # img.save(buffer, format=image_format, quality=quality) + # buffer.seek(0) + # + # return web.Response( + # body=buffer.read(), + # content_type=f"image/{image_format}", + # headers={ + # "Content-Disposition": f'filename="{filename}"' + # }, + # ) + # + # if "channel" not in request.rel_url.query: + # channel = "rgba" + # else: + # channel = request.rel_url.query["channel"] + # + # if channel == "rgb": + # with Image.open(file) as img: + # if img.mode == "RGBA": + # r, g, b, a = img.split() + # new_img = Image.merge("RGB", (r, g, b)) + # else: + # new_img = img.convert("RGB") + # + # buffer = BytesIO() + # new_img.save(buffer, format="PNG") + # buffer.seek(0) + # + # return web.Response( + # body=buffer.read(), + # content_type="image/png", + # headers={ + # "Content-Disposition": f'filename="{filename}"' + # }, + # ) + # + # elif channel == "a": + # with Image.open(file) as img: + # if img.mode == "RGBA": + # _, _, _, a = img.split() + # else: + # a = Image.new("L", img.size, 255) + # + # # alpha img + # alpha_img = Image.new("RGBA", img.size) + # alpha_img.putalpha(a) + # alpha_buffer = BytesIO() + # alpha_img.save(alpha_buffer, format="PNG") + # alpha_buffer.seek(0) + # + # return web.Response( + # body=alpha_buffer.read(), + # content_type="image/png", + # headers={ + # "Content-Disposition": f'filename="{filename}"' + # }, + # ) + # else: + # return web.FileResponse( + # file, + # headers={ + # "Content-Disposition": f'filename="{filename}"' + # }, + # ) + # + # return web.Response(status=404) + # @PromptServer.instance.routes.get("/mtb/debug") async def get_debug(request): from . import endpoint diff --git a/endpoint.py b/endpoint.py index 70709c8..22263f4 100644 --- a/endpoint.py +++ b/endpoint.py @@ -1,4 +1,6 @@ import csv +import os +import random from aiohttp import web @@ -6,6 +8,7 @@ from .log import mklog from .utils import ( backup_file, import_install, + input_dir, reqs_map, run_command, styles_dir, @@ -50,6 +53,24 @@ def ACTIONS_installDependency(dependency_names=None): # break +def ACTIONS_getInputs(count=200, offset=0): + # TODO: find a better name :s + enabled = "MTB_EXPOSE" in os.environ + if not enabled: + return {"error": "Session not authorized to getInputs"} + + imgs = {} + for i, img in enumerate(input_dir.glob("*.png")): + if i < offset: + continue + imgs[img.stem] = ( + f"/mtb/view?filename={img.name}&width=512&type=input&subfolder=&preview=&rand={random.random()}" + ) + if i >= count: + break + return imgs + + def ACTIONS_getStyles(style_name=None): from .nodes.conditions import MTB_StylesLoader @@ -153,7 +174,7 @@ def csv_editor(): html_out = """

Style Editor

- + """ for current, styles in style_files.items(): current_out = f"

{current}

" @@ -299,7 +320,7 @@ def render_table(table_dict, sort=True, title=None): {table_rows} - +
""" @@ -355,6 +376,6 @@ def render_base_template(title, content): - + """ diff --git a/web/mtb_input_output_sidebar.js b/web/mtb_input_output_sidebar.js new file mode 100644 index 0000000..2bb064a --- /dev/null +++ b/web/mtb_input_output_sidebar.js @@ -0,0 +1,50 @@ +import { app } from '../../scripts/app.js' +import { api } from '../../scripts/api.js' + +import * as shared from './comfy_shared.js' + +if (window?.__COMFYUI_FRONTEND_VERSION__) { + const version = window?.__COMFYUI_FRONTEND_VERSION__ + console.log(`%c ${version}`, 'background: orange; color: white;') + + app.extensionManager.registerSidebarTab({ + id: 'mtb-inputs-outputs', + icon: 'pi pi-images', + title: 'Input & Outputs', + tooltip: 'Browse inputs and outputs directories.', + type: 'custom', + // this is run everytime the tab's diplay is toggled on. + render: async (el) => { + const inputs = await api.fetchApi('/mtb/actions', { + method: 'POST', + body: JSON.stringify({ + name: 'getInputs', + }), + }) + const output = await inputs.json() + const urls = output?.result + if (!urls) return + + const cont = document.createElement('div') + Object.assign(cont, 'style', { + display: 'flex', + flexDirection: 'column', + alignItems: 'flex-start', + justifyContent: 'flex-start', + }) + + el.appendChild(cont) + + for (const [key, url] of Object.entries(urls)) { + const a = document.createElement('img') + a.src = url + a.width = 200 + cont.appendChild(a) + + // cont.appendChild(document.createElement('br')) + } + + // el.innerHTML = inputs.join('\n') // .map((p) => `
${p}
`).join('\n') + }, + }) +} diff --git a/web/mtb_sidebar.js b/web/mtb_sidebar.js new file mode 100644 index 0000000..5dec40b --- /dev/null +++ b/web/mtb_sidebar.js @@ -0,0 +1,25 @@ +import { app } from '../../scripts/app.js' +import { api } from '../../scripts/api.js' + +import * as shared from './comfy_shared.js' +import { createOutliner } from './dist/mtb_inspector.js' + +if (window?.__COMFYUI_FRONTEND_VERSION__) { + const version = window?.__COMFYUI_FRONTEND_VERSION__ + console.log(`%c ${version}`, 'background: orange; color: white;') + + const panel = app.extensionManager.registerSidebarTab({ + id: 'mtb-nodes', + icon: 'pi pi-bolt', + title: 'MTB', + tooltip: 'sidebar for mtb nodes', + type: 'custom', + // this is run everytime the tab's diplay is toggled on. + render: (el) => { + const outliner = createOutliner(el) + const inputs = shared.getAPIInputs() + console.log('INPUTS', inputs) + outliner.$$set({ inputs }) + }, + }) +}