brunvelop/ComfyUI
2
1import os2import sys3import asyncio4import traceback5 6import nodes7import folder_paths8import execution9import uuid10import urllib11import json12import glob13import struct14from PIL import Image, ImageOps15from PIL.PngImagePlugin import PngInfo16from io import BytesIO17 18try:19 import aiohttp20 from aiohttp import web21except ImportError:22 print("Module 'aiohttp' not installed. Please install it via:")23 print("pip install aiohttp")24 print("or")25 print("pip install -r requirements.txt")26 sys.exit()27 28import mimetypes29from comfy.cli_args import args30import comfy.utils31import comfy.model_management32 33 34class BinaryEventTypes:35 PREVIEW_IMAGE = 136 UNENCODED_PREVIEW_IMAGE = 237 38async def send_socket_catch_exception(function, message):39 try:40 await function(message)41 except (aiohttp.ClientError, aiohttp.ClientPayloadError, ConnectionResetError) as err:42 print("send error:", err)43 44@web.middleware45async def cache_control(request: web.Request, handler):46 response: web.Response = await handler(request)47 if request.path.endswith('.js') or request.path.endswith('.css'):48 response.headers.setdefault('Cache-Control', 'no-cache')49 return response50 51def create_cors_middleware(allowed_origin: str):52 @web.middleware53 async def cors_middleware(request: web.Request, handler):54 if request.method == "OPTIONS":55 # Pre-flight request. Reply successfully:56 response = web.Response()57 else:58 response = await handler(request)59 60 response.headers['Access-Control-Allow-Origin'] = allowed_origin61 response.headers['Access-Control-Allow-Methods'] = 'POST, GET, DELETE, PUT, OPTIONS'62 response.headers['Access-Control-Allow-Headers'] = 'Content-Type, Authorization'63 response.headers['Access-Control-Allow-Credentials'] = 'true'64 return response65 66 return cors_middleware67 68class PromptServer():69 def __init__(self, loop):70 PromptServer.instance = self71 72 mimetypes.init()73 mimetypes.types_map['.js'] = 'application/javascript; charset=utf-8'74 75 self.supports = ["custom_nodes_from_web"]76 self.prompt_queue = None77 self.loop = loop78 self.messages = asyncio.Queue()79 self.number = 080 81 middlewares = [cache_control]82 if args.enable_cors_header:83 middlewares.append(create_cors_middleware(args.enable_cors_header))84 85 max_upload_size = round(args.max_upload_size * 1024 * 1024)86 self.app = web.Application(client_max_size=max_upload_size, middlewares=middlewares)87 self.sockets = dict()88 self.web_root = os.path.join(os.path.dirname(89 os.path.realpath(__file__)), "web")90 routes = web.RouteTableDef()91 self.routes = routes92 self.last_node_id = None93 self.client_id = None94 95 self.on_prompt_handlers = []96 97 @routes.get('/ws')98 async def websocket_handler(request):99 ws = web.WebSocketResponse()100 await ws.prepare(request)101 sid = request.rel_url.query.get('clientId', '')102 if sid:103 # Reusing existing session, remove old104 self.sockets.pop(sid, None)105 else:106 sid = uuid.uuid4().hex107 108 self.sockets[sid] = ws109 110 try:111 # Send initial state to the new client112 await self.send("status", { "status": self.get_queue_info(), 'sid': sid }, sid)113 # On reconnect if we are the currently executing client send the current node114 if self.client_id == sid and self.last_node_id is not None:115 await self.send("executing", { "node": self.last_node_id }, sid)116 117 async for msg in ws:118 if msg.type == aiohttp.WSMsgType.ERROR:119 print('ws connection closed with exception %s' % ws.exception())120 finally:121 self.sockets.pop(sid, None)122 return ws123 124 @routes.get("/")125 async def get_root(request):126 return web.FileResponse(os.path.join(self.web_root, "index.html"))127 128 @routes.get("/embeddings")129 def get_embeddings(self):130 embeddings = folder_paths.get_filename_list("embeddings")131 return web.json_response(list(map(lambda a: os.path.splitext(a)[0], embeddings)))132 133 @routes.get("/extensions")134 async def get_extensions(request):135 files = glob.glob(os.path.join(136 glob.escape(self.web_root), 'extensions/**/*.js'), recursive=True)137 138 extensions = list(map(lambda f: "/" + os.path.relpath(f, self.web_root).replace("\\", "/"), files))139 140 for name, dir in nodes.EXTENSION_WEB_DIRS.items():141 files = glob.glob(os.path.join(glob.escape(dir), '**/*.js'), recursive=True)142 extensions.extend(list(map(lambda f: "/extensions/" + urllib.parse.quote(143 name) + "/" + os.path.relpath(f, dir).replace("\\", "/"), files)))144 145 return web.json_response(extensions)146 147 def get_dir_by_type(dir_type):148 if dir_type is None:149 dir_type = "input"150 151 if dir_type == "input":152 type_dir = folder_paths.get_input_directory()153 elif dir_type == "temp":154 type_dir = folder_paths.get_temp_directory()155 elif dir_type == "output":156 type_dir = folder_paths.get_output_directory()157 158 return type_dir, dir_type159 160 def image_upload(post, image_save_function=None):161 image = post.get("image")162 overwrite = post.get("overwrite")163 164 image_upload_type = post.get("type")165 upload_dir, image_upload_type = get_dir_by_type(image_upload_type)166 167 if image and image.file:168 filename = image.filename169 if not filename:170 return web.Response(status=400)171 172 subfolder = post.get("subfolder", "")173 full_output_folder = os.path.join(upload_dir, os.path.normpath(subfolder))174 filepath = os.path.abspath(os.path.join(full_output_folder, filename))175 176 if os.path.commonpath((upload_dir, filepath)) != upload_dir:177 return web.Response(status=400)178 179 if not os.path.exists(full_output_folder):180 os.makedirs(full_output_folder)181 182 split = os.path.splitext(filename)183 184 if overwrite is not None and (overwrite == "true" or overwrite == "1"):185 pass186 else:187 i = 1188 while os.path.exists(filepath):189 filename = f"{split[0]} ({i}){split[1]}"190 filepath = os.path.join(full_output_folder, filename)191 i += 1192 193 if image_save_function is not None:194 image_save_function(image, post, filepath)195 else:196 with open(filepath, "wb") as f:197 f.write(image.file.read())198 199 return web.json_response({"name" : filename, "subfolder": subfolder, "type": image_upload_type})200 else:201 return web.Response(status=400)202 203 @routes.post("/upload/image")204 async def upload_image(request):205 post = await request.post()206 return image_upload(post)207 208 209 @routes.post("/upload/mask")210 async def upload_mask(request):211 post = await request.post()212 213 def image_save_function(image, post, filepath):214 original_ref = json.loads(post.get("original_ref"))215 filename, output_dir = folder_paths.annotated_filepath(original_ref['filename'])216 217 # validation for security: prevent accessing arbitrary path218 if filename[0] == '/' or '..' in filename:219 return web.Response(status=400)220 221 if output_dir is None:222 type = original_ref.get("type", "output")223 output_dir = folder_paths.get_directory_by_type(type)224 225 if output_dir is None:226 return web.Response(status=400)227 228 if original_ref.get("subfolder", "") != "":229 full_output_dir = os.path.join(output_dir, original_ref["subfolder"])230 if os.path.commonpath((os.path.abspath(full_output_dir), output_dir)) != output_dir:231 return web.Response(status=403)232 output_dir = full_output_dir233 234 file = os.path.join(output_dir, filename)235 236 if os.path.isfile(file):237 with Image.open(file) as original_pil:238 metadata = PngInfo()239 if hasattr(original_pil,'text'):240 for key in original_pil.text:241 metadata.add_text(key, original_pil.text[key])242 original_pil = original_pil.convert('RGBA')243 mask_pil = Image.open(image.file).convert('RGBA')244 245 # alpha copy246 new_alpha = mask_pil.getchannel('A')247 original_pil.putalpha(new_alpha)248 original_pil.save(filepath, compress_level=4, pnginfo=metadata)249 250 return image_upload(post, image_save_function)251 252 @routes.get("/view")253 async def view_image(request):254 if "filename" in request.rel_url.query:255 filename = request.rel_url.query["filename"]256 filename,output_dir = folder_paths.annotated_filepath(filename)257 258 # validation for security: prevent accessing arbitrary path259 if filename[0] == '/' or '..' in filename:260 return web.Response(status=400)261 262 if output_dir is None:263 type = request.rel_url.query.get("type", "output")264 output_dir = folder_paths.get_directory_by_type(type)265 266 if output_dir is None:267 return web.Response(status=400)268 269 if "subfolder" in request.rel_url.query:270 full_output_dir = os.path.join(output_dir, request.rel_url.query["subfolder"])271 if os.path.commonpath((os.path.abspath(full_output_dir), output_dir)) != output_dir:272 return web.Response(status=403)273 output_dir = full_output_dir274 275 filename = os.path.basename(filename)276 file = os.path.join(output_dir, filename)277 278 if os.path.isfile(file):279 if 'preview' in request.rel_url.query:280 with Image.open(file) as img:281 preview_info = request.rel_url.query['preview'].split(';')282 image_format = preview_info[0]283 if image_format not in ['webp', 'jpeg'] or 'a' in request.rel_url.query.get('channel', ''):284 image_format = 'webp'285 286 quality = 90287 if preview_info[-1].isdigit():288 quality = int(preview_info[-1])289 290 buffer = BytesIO()291 if image_format in ['jpeg'] or request.rel_url.query.get('channel', '') == 'rgb':292 img = img.convert("RGB")293 img.save(buffer, format=image_format, quality=quality)294 buffer.seek(0)295 296 return web.Response(body=buffer.read(), content_type=f'image/{image_format}',297 headers={"Content-Disposition": f"filename=\"{filename}\""})298 299 if 'channel' not in request.rel_url.query:300 channel = 'rgba'301 else:302 channel = request.rel_url.query["channel"]303 304 if channel == 'rgb':305 with Image.open(file) as img:306 if img.mode == "RGBA":307 r, g, b, a = img.split()308 new_img = Image.merge('RGB', (r, g, b))309 else:310 new_img = img.convert("RGB")311 312 buffer = BytesIO()313 new_img.save(buffer, format='PNG')314 buffer.seek(0)315 316 return web.Response(body=buffer.read(), content_type='image/png',317 headers={"Content-Disposition": f"filename=\"{filename}\""})318 319 elif channel == 'a':320 with Image.open(file) as img:321 if img.mode == "RGBA":322 _, _, _, a = img.split()323 else:324 a = Image.new('L', img.size, 255)325 326 # alpha img327 alpha_img = Image.new('RGBA', img.size)328 alpha_img.putalpha(a)329 alpha_buffer = BytesIO()330 alpha_img.save(alpha_buffer, format='PNG')331 alpha_buffer.seek(0)332 333 return web.Response(body=alpha_buffer.read(), content_type='image/png',334 headers={"Content-Disposition": f"filename=\"{filename}\""})335 else:336 return web.FileResponse(file, headers={"Content-Disposition": f"filename=\"{filename}\""})337 338 return web.Response(status=404)339 340 @routes.get("/view_metadata/{folder_name}")341 async def view_metadata(request):342 folder_name = request.match_info.get("folder_name", None)343 if folder_name is None:344 return web.Response(status=404)345 if not "filename" in request.rel_url.query:346 return web.Response(status=404)347 348 filename = request.rel_url.query["filename"]349 if not filename.endswith(".safetensors"):350 return web.Response(status=404)351 352 safetensors_path = folder_paths.get_full_path(folder_name, filename)353 if safetensors_path is None:354 return web.Response(status=404)355 out = comfy.utils.safetensors_header(safetensors_path, max_size=1024*1024)356 if out is None:357 return web.Response(status=404)358 dt = json.loads(out)359 if not "__metadata__" in dt:360 return web.Response(status=404)361 return web.json_response(dt["__metadata__"])362 363 @routes.get("/system_stats")364 async def get_queue(request):365 device = comfy.model_management.get_torch_device()366 device_name = comfy.model_management.get_torch_device_name(device)367 vram_total, torch_vram_total = comfy.model_management.get_total_memory(device, torch_total_too=True)368 vram_free, torch_vram_free = comfy.model_management.get_free_memory(device, torch_free_too=True)369 system_stats = {370 "system": {371 "os": os.name,372 "python_version": sys.version,373 "embedded_python": os.path.split(os.path.split(sys.executable)[0])[1] == "python_embeded"374 },375 "devices": [376 {377 "name": device_name,378 "type": device.type,379 "index": device.index,380 "vram_total": vram_total,381 "vram_free": vram_free,382 "torch_vram_total": torch_vram_total,383 "torch_vram_free": torch_vram_free,384 }385 ]386 }387 return web.json_response(system_stats)388 389 @routes.get("/prompt")390 async def get_prompt(request):391 return web.json_response(self.get_queue_info())392 393 def node_info(node_class):394 obj_class = nodes.NODE_CLASS_MAPPINGS[node_class]395 info = {}396 info['input'] = obj_class.INPUT_TYPES()397 info['output'] = obj_class.RETURN_TYPES398 info['output_is_list'] = obj_class.OUTPUT_IS_LIST if hasattr(obj_class, 'OUTPUT_IS_LIST') else [False] * len(obj_class.RETURN_TYPES)399 info['output_name'] = obj_class.RETURN_NAMES if hasattr(obj_class, 'RETURN_NAMES') else info['output']400 info['name'] = node_class401 info['display_name'] = nodes.NODE_DISPLAY_NAME_MAPPINGS[node_class] if node_class in nodes.NODE_DISPLAY_NAME_MAPPINGS.keys() else node_class402 info['description'] = obj_class.DESCRIPTION if hasattr(obj_class,'DESCRIPTION') else ''403 info['category'] = 'sd'404 if hasattr(obj_class, 'OUTPUT_NODE') and obj_class.OUTPUT_NODE == True:405 info['output_node'] = True406 else:407 info['output_node'] = False408 409 if hasattr(obj_class, 'CATEGORY'):410 info['category'] = obj_class.CATEGORY411 return info412 413 @routes.get("/object_info")414 async def get_object_info(request):415 out = {}416 for x in nodes.NODE_CLASS_MAPPINGS:417 try:418 out[x] = node_info(x)419 except Exception as e:420 print(f"[ERROR] An error occurred while retrieving information for the '{x}' node.", file=sys.stderr)421 traceback.print_exc()422 return web.json_response(out)423 424 @routes.get("/object_info/{node_class}")425 async def get_object_info_node(request):426 node_class = request.match_info.get("node_class", None)427 out = {}428 if (node_class is not None) and (node_class in nodes.NODE_CLASS_MAPPINGS):429 out[node_class] = node_info(node_class)430 return web.json_response(out)431 432 @routes.get("/history")433 async def get_history(request):434 max_items = request.rel_url.query.get("max_items", None)435 if max_items is not None:436 max_items = int(max_items)437 return web.json_response(self.prompt_queue.get_history(max_items=max_items))438 439 @routes.get("/history/{prompt_id}")440 async def get_history(request):441 prompt_id = request.match_info.get("prompt_id", None)442 return web.json_response(self.prompt_queue.get_history(prompt_id=prompt_id))443 444 @routes.get("/queue")445 async def get_queue(request):446 queue_info = {}447 current_queue = self.prompt_queue.get_current_queue()448 queue_info['queue_running'] = current_queue[0]449 queue_info['queue_pending'] = current_queue[1]450 return web.json_response(queue_info)451 452 @routes.post("/prompt")453 async def post_prompt(request):454 print("got prompt")455 resp_code = 200456 out_string = ""457 json_data = await request.json()458 json_data = self.trigger_on_prompt(json_data)459 460 if "number" in json_data:461 number = float(json_data['number'])462 else:463 number = self.number464 if "front" in json_data:465 if json_data['front']:466 number = -number467 468 self.number += 1469 470 if "prompt" in json_data:471 prompt = json_data["prompt"]472 valid = execution.validate_prompt(prompt)473 extra_data = {}474 if "extra_data" in json_data:475 extra_data = json_data["extra_data"]476 477 if "client_id" in json_data:478 extra_data["client_id"] = json_data["client_id"]479 if valid[0]:480 prompt_id = str(uuid.uuid4())481 outputs_to_execute = valid[2]482 self.prompt_queue.put((number, prompt_id, prompt, extra_data, outputs_to_execute))483 response = {"prompt_id": prompt_id, "number": number, "node_errors": valid[3]}484 return web.json_response(response)485 else:486 print("invalid prompt:", valid[1])487 return web.json_response({"error": valid[1], "node_errors": valid[3]}, status=400)488 else:489 return web.json_response({"error": "no prompt", "node_errors": []}, status=400)490 491 @routes.post("/queue")492 async def post_queue(request):493 json_data = await request.json()494 if "clear" in json_data:495 if json_data["clear"]:496 self.prompt_queue.wipe_queue()497 if "delete" in json_data:498 to_delete = json_data['delete']499 for id_to_delete in to_delete:500 delete_func = lambda a: a[1] == id_to_delete501 self.prompt_queue.delete_queue_item(delete_func)502 503 return web.Response(status=200)504 505 @routes.post("/interrupt")506 async def post_interrupt(request):507 nodes.interrupt_processing()508 return web.Response(status=200)509 510 @routes.post("/history")511 async def post_history(request):512 json_data = await request.json()513 if "clear" in json_data:514 if json_data["clear"]:515 self.prompt_queue.wipe_history()516 if "delete" in json_data:517 to_delete = json_data['delete']518 for id_to_delete in to_delete:519 self.prompt_queue.delete_history_item(id_to_delete)520 521 return web.Response(status=200)522 523 def add_routes(self):524 self.app.add_routes(self.routes)525 526 for name, dir in nodes.EXTENSION_WEB_DIRS.items():527 self.app.add_routes([528 web.static('/extensions/' + urllib.parse.quote(name), dir, follow_symlinks=True),529 ])530 531 self.app.add_routes([532 web.static('/', self.web_root, follow_symlinks=True),533 ])534 535 def get_queue_info(self):536 prompt_info = {}537 exec_info = {}538 exec_info['queue_remaining'] = self.prompt_queue.get_tasks_remaining()539 prompt_info['exec_info'] = exec_info540 return prompt_info541 542 async def send(self, event, data, sid=None):543 if event == BinaryEventTypes.UNENCODED_PREVIEW_IMAGE:544 await self.send_image(data, sid=sid)545 elif isinstance(data, (bytes, bytearray)):546 await self.send_bytes(event, data, sid)547 else:548 await self.send_json(event, data, sid)549 550 def encode_bytes(self, event, data):551 if not isinstance(event, int):552 raise RuntimeError(f"Binary event types must be integers, got {event}")553 554 packed = struct.pack(">I", event)555 message = bytearray(packed)556 message.extend(data)557 return message558 559 async def send_image(self, image_data, sid=None):560 image_type = image_data[0]561 image = image_data[1]562 max_size = image_data[2]563 if max_size is not None:564 if hasattr(Image, 'Resampling'):565 resampling = Image.Resampling.BILINEAR566 else:567 resampling = Image.ANTIALIAS568 569 image = ImageOps.contain(image, (max_size, max_size), resampling)570 type_num = 1571 if image_type == "JPEG":572 type_num = 1573 elif image_type == "PNG":574 type_num = 2575 576 bytesIO = BytesIO()577 header = struct.pack(">I", type_num)578 bytesIO.write(header)579 image.save(bytesIO, format=image_type, quality=95, compress_level=4)580 preview_bytes = bytesIO.getvalue()581 await self.send_bytes(BinaryEventTypes.PREVIEW_IMAGE, preview_bytes, sid=sid)582 583 async def send_bytes(self, event, data, sid=None):584 message = self.encode_bytes(event, data)585 586 if sid is None:587 for ws in self.sockets.values():588 await send_socket_catch_exception(ws.send_bytes, message)589 elif sid in self.sockets:590 await send_socket_catch_exception(self.sockets[sid].send_bytes, message)591 592 async def send_json(self, event, data, sid=None):593 message = {"type": event, "data": data}594 595 if sid is None:596 for ws in self.sockets.values():597 await send_socket_catch_exception(ws.send_json, message)598 elif sid in self.sockets:599 await send_socket_catch_exception(self.sockets[sid].send_json, message)600 601 def send_sync(self, event, data, sid=None):602 self.loop.call_soon_threadsafe(603 self.messages.put_nowait, (event, data, sid))604 605 def queue_updated(self):606 self.send_sync("status", { "status": self.get_queue_info() })607 608 async def publish_loop(self):609 while True:610 msg = await self.messages.get()611 await self.send(*msg)612 613 async def start(self, address, port, verbose=True, call_on_start=None):614 runner = web.AppRunner(self.app, access_log=None)615 await runner.setup()616 site = web.TCPSite(runner, address, port)617 await site.start()618 619 if address == '':620 address = '0.0.0.0'621 if verbose:622 print("Starting server\n")623 print("To see the GUI go to: http://{}:{}".format(address, port))624 if call_on_start is not None:625 call_on_start(address, port)626 627 def add_on_prompt_handler(self, handler):628 self.on_prompt_handlers.append(handler)629 630 def trigger_on_prompt(self, json_data):631 for handler in self.on_prompt_handlers:632 try:633 json_data = handler(json_data)634 except Exception as e:635 print(f"[ERROR] An error occurred during the on_prompt_handler processing")636 traceback.print_exc()637 638 return json_data639 