diff --git a/.pylintrc b/.pylintrc index 2a66d3dfc..f644719cb 100644 --- a/.pylintrc +++ b/.pylintrc @@ -138,6 +138,7 @@ disable=raw-checker-failed, unnecessary-lambda, consider-using-dict-items, dangerous-default-value, + unnecessary-dunder-call, enable=c-extension-no-member [METHOD_ARGS] diff --git a/extensions-builtin/sd-extension-system-info b/extensions-builtin/sd-extension-system-info index 78e841f97..9e6b6b09b 160000 --- a/extensions-builtin/sd-extension-system-info +++ b/extensions-builtin/sd-extension-system-info @@ -1 +1 @@ -Subproject commit 78e841f97f8b292888f45c5435c015dbb489ab57 +Subproject commit 9e6b6b09bd92d912cc142b307b50a2c31a7eac08 diff --git a/extensions-builtin/sd-webui-controlnet b/extensions-builtin/sd-webui-controlnet index 2f7d73a29..d499bd408 160000 --- a/extensions-builtin/sd-webui-controlnet +++ b/extensions-builtin/sd-webui-controlnet @@ -1 +1 @@ -Subproject commit 2f7d73a299ad10b7c866b3760b09a63812699788 +Subproject commit d499bd4080dc87ef00dc5a7bc2024483bc84b51b diff --git a/javascript/ui.js b/javascript/ui.js index 96ec204b9..cca454040 100644 --- a/javascript/ui.js +++ b/javascript/ui.js @@ -344,7 +344,7 @@ function update_token_counter(button_id) { function restart_reload(){ document.body.innerHTML='

Reloading...

'; - setTimeout(function(){location.reload()},2000) + setTimeout(function(){location.reload()},8000) return [] } diff --git a/modules/api/api.py b/modules/api/api.py index fd93e4bf8..125bc5bf0 100644 --- a/modules/api/api.py +++ b/modules/api/api.py @@ -1,19 +1,13 @@ import io import time import base64 -import datetime from io import BytesIO from typing import List from threading import Lock from secrets import compare_digest -import anyio -import starlette -import fastapi -from fastapi import APIRouter, Depends, FastAPI, Request, Response +from fastapi import APIRouter, Depends, FastAPI from fastapi.security import HTTPBasic, HTTPBasicCredentials from fastapi.exceptions import HTTPException -from fastapi.responses import JSONResponse -from fastapi.encoders import jsonable_encoder from PIL import PngImagePlugin,Image import piexif import piexif.helper @@ -33,15 +27,17 @@ from modules.sd_models_config import find_checkpoint_config_near_filename from modules.realesrgan_model import get_realesrgan_models from modules import devices +errors.install() + def upscaler_to_index(name: str): try: return [x.name.lower() for x in shared.sd_upscalers].index(name.lower()) except: raise HTTPException(status_code=400, detail=f"Invalid upscaler, needs to be one of these: {' , '.join([x.name for x in sd_upscalers])}") -def script_name_to_index(name, scripts): +def script_name_to_index(name, scripts_list): try: - return [script.title().lower() for script in scripts].index(name.lower()) + return [script.title().lower() for script in scripts_list].index(name.lower()) except: raise HTTPException(status_code=422, detail=f"Script '{name}' not found") @@ -96,55 +92,6 @@ def encode_pil_to_base64(image): return base64.b64encode(bytes_data) -def api_middleware(app: FastAPI): - - @app.middleware("http") - async def log_and_time(req: Request, call_next): - ts = time.time() - res: Response = await call_next(req) - duration = str(round(time.time() - ts, 4)) - res.headers["X-Process-Time"] = duration - endpoint = req.scope.get('path', 'err') - if shared.cmd_opts.api_log and endpoint.startswith('/sdapi'): - print('API {t} {code} {prot}/{ver} {method} {endpoint} {cli} {duration}'.format( - t = datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S.%f"), - code = res.status_code, - ver = req.scope.get('http_version', '0.0'), - cli = req.scope.get('client', ('0:0.0.0', 0))[0], - prot = req.scope.get('scheme', 'err'), - method = req.scope.get('method', 'err'), - endpoint = endpoint, - duration = duration, - )) - return res - - def handle_exception(request: Request, e: Exception): - err = { - "error": type(e).__name__, - "detail": vars(e).get('detail', ''), - "body": vars(e).get('body', ''), - "errors": str(e), - } - print(f"API error: {request.method}: {request.url} {err}") - if not isinstance(e, HTTPException): # do not print backtrace on known httpexceptions - errors.display(e, 'http api', [anyio, fastapi, uvicorn, starlette]) - return JSONResponse(status_code=vars(e).get('status_code', 500), content=jsonable_encoder(err)) - - @app.middleware("http") - async def exception_handling(request: Request, call_next): - try: - return await call_next(request) - except Exception as e: - return handle_exception(request, e) - - @app.exception_handler(Exception) - async def fastapi_exception_handler(request: Request, e: Exception): - return handle_exception(request, e) - - @app.exception_handler(HTTPException) - async def http_exception_handler(request: Request, e: HTTPException): - return handle_exception(request, e) - class Api: def __init__(self, app: FastAPI, queue_lock: Lock): @@ -157,7 +104,6 @@ class Api: self.router = APIRouter() self.app = app self.queue_lock = queue_lock - api_middleware(self.app) self.add_api_route("/sdapi/v1/txt2img", self.text2imgapi, methods=["POST"], response_model=TextToImageResponse) self.add_api_route("/sdapi/v1/img2img", self.img2imgapi, methods=["POST"], response_model=ImageToImageResponse) self.add_api_route("/sdapi/v1/extra-single-image", self.extras_single_image_api, methods=["POST"], response_model=ExtrasSingleImageResponse) diff --git a/modules/cmd_args.py b/modules/cmd_args.py index a484410ae..bdc3dd5b6 100644 --- a/modules/cmd_args.py +++ b/modules/cmd_args.py @@ -62,6 +62,7 @@ def compatibility_args(opts, args): parser.add_argument("--clip-models-path", type=str, help=argparse.SUPPRESS, default=opts.clip_models_path) parser.add_argument("--disable-nan-check", default = True, action='store_true', help=argparse.SUPPRESS) parser.add_argument("--disable-extension-access", default = False, action='store_true', help=argparse.SUPPRESS) + parser.add_argument("--xformers", default = (opts.cross_attention_optimization == "xFormers"), action='store_true', help=argparse.SUPPRESS) args = parser.parse_args() if vars(parser)['_option_string_actions'].get('--lora-dir', None) is not None: args.lora_dir = opts.lora_dir diff --git a/modules/lycoris b/modules/lycoris index 53119eb85..ded70eedd 160000 --- a/modules/lycoris +++ b/modules/lycoris @@ -1 +1 @@ -Subproject commit 53119eb852b0450ab9294437d78dcdfa0d58a0dc +Subproject commit ded70eeddc2b26bf3efd25c2b29b6746294bb731 diff --git a/modules/middleware.py b/modules/middleware.py new file mode 100644 index 000000000..fbc50da9b --- /dev/null +++ b/modules/middleware.py @@ -0,0 +1,86 @@ +import time +import datetime +import logging +from asyncio.exceptions import CancelledError +import anyio +import starlette +import uvicorn +import fastapi +from starlette.responses import JSONResponse +from fastapi import FastAPI, Request, Response +from fastapi.exceptions import HTTPException +from fastapi.encoders import jsonable_encoder +import modules.errors as errors + +logging.disable(logging.CRITICAL) +errors.install() + +def setup_middleware(app: FastAPI, cmd_opts): + print('Initializing middleware') + # uvicorn_logger=logging.getLogger("uvicorn.error") + # uvicorn_logger.disabled = True + from fastapi.middleware.cors import CORSMiddleware + from fastapi.middleware.gzip import GZipMiddleware + app.user_middleware = [x for x in app.user_middleware if x.cls.__name__ != 'CORSMiddleware'] + app.middleware_stack = None # reset current middleware to allow modifying user provided list + app.add_middleware(GZipMiddleware, minimum_size=1024) + if cmd_opts.cors_origins and cmd_opts.cors_regex: + app.add_middleware(CORSMiddleware, allow_origins=cmd_opts.cors_origins.split(','), allow_origin_regex=cmd_opts.cors_regex, allow_methods=['*'], allow_credentials=True, allow_headers=['*']) + elif cmd_opts.cors_origins: + app.add_middleware(CORSMiddleware, allow_origins=cmd_opts.cors_origins.split(','), allow_methods=['*'], allow_credentials=True, allow_headers=['*']) + elif cmd_opts.cors_regex: + app.add_middleware(CORSMiddleware, allow_origin_regex=cmd_opts.cors_regex, allow_methods=['*'], allow_credentials=True, allow_headers=['*']) + + @app.middleware("http") + async def log_and_time(req: Request, call_next): + ts = time.time() + res: Response = await call_next(req) + duration = str(round(time.time() - ts, 4)) + res.headers["X-Process-Time"] = duration + endpoint = req.scope.get('path', 'err') + if cmd_opts.api_log and endpoint.startswith('/sdapi'): + print('API {t} {code} {prot}/{ver} {method} {endpoint} {cli} {duration}'.format( # pylint: disable=consider-using-f-string + t = datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S.%f"), + code = res.status_code, + ver = req.scope.get('http_version', '0.0'), + cli = req.scope.get('client', ('0:0.0.0', 0))[0], + prot = req.scope.get('scheme', 'err'), + method = req.scope.get('method', 'err'), + endpoint = endpoint, + duration = duration, + )) + return res + + def handle_exception(req: Request, e: Exception): + err = { + "error": type(e).__name__, + "detail": vars(e).get('detail', ''), + "body": vars(e).get('body', ''), + "errors": str(e), + } + print(f"API error: {req.method}: {req.url} {err}") + if not isinstance(e, HTTPException) and err['error'] != 'TypeError': # do not print backtrace on known httpexceptions + errors.display(e, 'HTTP API', [anyio, fastapi, uvicorn, starlette]) + return JSONResponse(status_code=vars(e).get('status_code', 500), content=jsonable_encoder(err)) + + @app.middleware("http") + async def exception_handling(req: Request, call_next): + try: + return await call_next(req) + except CancelledError: + print('WebSocket closed') + except BaseException as e: + return handle_exception(req, e) + + @app.exception_handler(HTTPException) + async def http_exception_handler(req: Request, e: HTTPException): + return handle_exception(req, e) + + @app.exception_handler(Exception) + async def general_exception_handler(req: Request, e: Exception): + if isinstance(e, TypeError): + return JSONResponse(status_code=500, content=jsonable_encoder(str(e))) + else: + return handle_exception(req, e) + + app.build_middleware_stack() # rebuild middleware stack on-the-fly diff --git a/modules/modelloader.py b/modules/modelloader.py index 2b6653903..597afbdf1 100644 --- a/modules/modelloader.py +++ b/modules/modelloader.py @@ -103,6 +103,9 @@ def cleanup_models(): src_path = os.path.join(root_path, "repositories/latent-diffusion/experiments/pretrained_models/") dest_path = os.path.join(models_path, "LDSR") move_files(src_path, dest_path) + src_path = os.path.join(root_path, "ScuNET") + dest_path = os.path.join(models_path, "ScuNET") + move_files(src_path, dest_path) def move_files(src_path: str, dest_path: str, ext_filter: str = None): @@ -134,11 +137,9 @@ forbidden_upscaler_classes = set() def list_builtin_upscalers(): load_upscalers() - builtin_upscaler_classes.clear() builtin_upscaler_classes.extend(Upscaler.__subclasses__()) - def forbid_loaded_nonbuiltin_upscalers(): for cls in Upscaler.__subclasses__(): if cls not in builtin_upscaler_classes: diff --git a/modules/scripts.py b/modules/scripts.py index 06235344b..2606ff267 100644 --- a/modules/scripts.py +++ b/modules/scripts.py @@ -238,20 +238,12 @@ def load_scripts(): elif issubclass(script_class, scripts_postprocessing.ScriptPostprocessing): postprocessing_scripts_data.append(ScriptClassData(script_class, scriptfile.path, scriptfile.basedir, module)) - def orderby(basedir): - # 1st webui, 2nd extensions-builtin, 3rd extensions - priority = {os.path.join(paths.script_path, "extensions-builtin"):1, paths.script_path:0} - for key in priority: - if basedir.startswith(key): - return priority[key] - return 9999 - - for scriptfile in sorted(scripts_list, key=lambda x: [orderby(x.basedir), x]): + alpha_sort = sorted(scripts_list, key=lambda item: item.path.lower()) + for scriptfile in alpha_sort: try: if scriptfile.basedir != paths.script_path: sys.path = [scriptfile.basedir] + sys.path current_basedir = scriptfile.basedir - script_module = script_loading.load_module(scriptfile.path) register_scripts_from_module(script_module) diff --git a/modules/shared.py b/modules/shared.py index f886ab422..0551d3736 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -632,6 +632,8 @@ def reload_gradio_theme(theme_name=None): print("Theme download error accessing HuggingFace") gradio_theme = gr.themes.Default() print(f'Loading theme: {theme_name}') + if demo is not None: + demo.close() class TotalTQDM: @@ -673,6 +675,19 @@ mem_mon = modules.memmon.MemUsageMonitor("MemMon", device, opts) mem_mon.start() +def restart_server(): + if demo is None: + return + try: + demo.server.should_exit = True + demo.server.force_exit = True + demo.close(verbose=False) + demo.server.close() + except: + pass + print('Server shutdown') + + def listfiles(dirname): filenames = [os.path.join(dirname, x) for x in sorted(os.listdir(dirname), key=str.lower) if not x.startswith(".")] return [file for file in filenames if os.path.isfile(file)] diff --git a/modules/ui.py b/modules/ui.py index a9455f565..791a7cf87 100644 --- a/modules/ui.py +++ b/modules/ui.py @@ -32,8 +32,8 @@ from modules.textual_inversion import textual_inversion from modules.generation_parameters_copypaste import image_from_url_text import modules.extras +errors.install() warnings.filterwarnings("ignore", category=UserWarning) - # this is a fix for Windows users. Without it, javascript files will be served with text/html content-type and the browser will not show any UI mimetypes.init() mimetypes.add_type('application/javascript', '.js') @@ -1387,8 +1387,9 @@ def create_ui(): with gr.Blocks(analytics_enabled=False) as settings_interface: with gr.Row(): - with gr.Column(scale=6): - settings_submit = gr.Button(value="Apply settings", variant='primary', elem_id="settings_submit") + settings_submit = gr.Button(value="Apply settings", variant='primary', elem_id="settings_submit") + restart_submit = gr.Button(value="Restart UI", variant='primary', elem_id="restart_submit") + result = gr.HTML(elem_id="settings_result") @@ -1525,6 +1526,7 @@ def create_ui(): inputs=components, outputs=[text_settings, result], ) + restart_submit.click(fn=shared.restart_server, _js="restart_reload") for i, k, item in quicksettings_list: component = component_dict[k] diff --git a/modules/ui_extensions.py b/modules/ui_extensions.py index 4e316238b..da51af8e8 100644 --- a/modules/ui_extensions.py +++ b/modules/ui_extensions.py @@ -45,6 +45,7 @@ def apply_and_restart(disable_list, update_list, disable_all): shared.state.interrupt() shared.state.need_restart = True + shared.restart_server() def check_updates(_id_task, disable_list): @@ -291,7 +292,7 @@ def create_ui(): with gr.TabItem("Installed"): with gr.Row(elem_id="extensions_installed_top"): - apply = gr.Button(value="Apply (restart required)", variant="primary") + apply = gr.Button(value="Apply & restart UI", variant="primary") check = gr.Button(value="Check for updates") extensions_disable_all = gr.Radio(label="Disable all extensions", choices=["none", "extra", "all"], value=shared.opts.disable_all_extensions, elem_id="extensions_disable_all") extensions_disabled_list = gr.Text(elem_id="extensions_disabled_list", visible=False).style(container=False) diff --git a/webui.py b/webui.py index ef6a88f07..7204c7f73 100644 --- a/webui.py +++ b/webui.py @@ -1,10 +1,10 @@ import os import re +import time import signal import warnings import logging from rich import print # pylint: disable=W0622 -from setup import log from modules import timer, errors startup_timer = timer.Timer() @@ -55,6 +55,7 @@ import modules.ui from modules import modelloader from modules.shared import cmd_opts, opts import modules.hypernetworks.hypernetwork +from modules.middleware import setup_middleware startup_timer.record("libraries") if cmd_opts.server_name: @@ -108,19 +109,19 @@ def initialize(): if cmd_opts.tls_keyfile is not None and cmd_opts.tls_keyfile is not None: try: if not os.path.exists(cmd_opts.tls_keyfile): - log.warning("Invalid path to TLS keyfile given") + print("Invalid path to TLS keyfile given") if not os.path.exists(cmd_opts.tls_certfile): - log.warning(f"Invalid path to TLS certfile: '{cmd_opts.tls_certfile}'") + print(f"Invalid path to TLS certfile: '{cmd_opts.tls_certfile}'") except TypeError: cmd_opts.tls_keyfile = cmd_opts.tls_certfile = None - log.warning("TLS setup invalid, running webui without TLS") + print("TLS setup invalid, running webui without TLS") else: - log.info("Running with TLS") + print("Running with TLS") startup_timer.record("TLS") # make the program just exit at ctrl+c without waiting for anything def sigint_handler(_sig, _frame): - log.info('Exiting') + print('Exiting') os._exit(0) signal.signal(signal.SIGINT, sigint_handler) @@ -133,10 +134,10 @@ def load_model(): modules.sd_models.load_model() except Exception as e: errors.display(e, "loading stable diffusion model") - log.error("Stable diffusion model failed to load") + print("Stable diffusion model failed to load") exit(1) if shared.sd_model is None: - log.error("No stable diffusion model loaded") + print("No stable diffusion model loaded") exit(1) shared.opts.data["sd_model_checkpoint"] = shared.sd_model.sd_checkpoint_info.title shared.opts.onchange("sd_model_checkpoint", wrap_queued_call(lambda: modules.sd_models.reload_model_weights())) @@ -144,20 +145,6 @@ def load_model(): startup_timer.record("checkpoint") -def setup_middleware(app): - from fastapi.middleware.cors import CORSMiddleware - from fastapi.middleware.gzip import GZipMiddleware - app.middleware_stack = None # reset current middleware to allow modifying user provided list - app.add_middleware(GZipMiddleware, minimum_size=1024) - if cmd_opts.cors_origins and cmd_opts.cors_regex: - app.add_middleware(CORSMiddleware, allow_origins=cmd_opts.cors_origins.split(','), allow_origin_regex=cmd_opts.cors_regex, allow_methods=['*'], allow_credentials=True, allow_headers=['*']) - elif cmd_opts.cors_origins: - app.add_middleware(CORSMiddleware, allow_origins=cmd_opts.cors_origins.split(','), allow_methods=['*'], allow_credentials=True, allow_headers=['*']) - elif cmd_opts.cors_regex: - app.add_middleware(CORSMiddleware, allow_origin_regex=cmd_opts.cors_regex, allow_methods=['*'], allow_credentials=True, allow_headers=['*']) - app.build_middleware_stack() # rebuild middleware stack on-the-fly - - def create_api(app): from modules.api.api import Api api = Api(app, queue_lock) @@ -195,20 +182,11 @@ def start_ui(): prevent_thread_lock=True, favicon_path='automatic.ico', ) - # for dep in shared.demo.dependencies: - # dep['show_progress'] = False # disable gradio css animation on component update - # app is instance of FastAPI server - # shared.demo.server is instance of gradio class which inherits from uvicorn.Server - # shared.demo.config is instance of uvicorn.Config - # shared.demo.app is instance of ASGIApp - + setup_middleware(app, cmd_opts) cmd_opts.autolaunch = False startup_timer.record("start") - app.user_middleware = [x for x in app.user_middleware if x.cls.__name__ != 'CORSMiddleware'] - setup_middleware(app) - modules.progress.setup_progress_api(app) create_api(app) ui_extra_networks.add_pages_to_demo(app) @@ -217,25 +195,22 @@ def start_ui(): startup_timer.record("scripts app_started_callback") -def stop_ui(): - try: - shared.demo.server.should_exit = True - shared.demo.server.force_exit = True - shared.demo.server.close() - except: - print('Uvicorn shutdown') - shared.demo.close(verbose=True) - - def webui(): start_ui() - load_model() - log.info(f"Startup time: {startup_timer.summary()}") + print(f"Startup time: {startup_timer.summary()}") - import time while True: - time.sleep(0.1) + try: + alive = shared.demo.server.thread.is_alive() + except: + alive = False + if not alive: + print('Server restart') + startup_timer.reset() + start_ui() + print(f"Startup time: {startup_timer.summary()}") + time.sleep(1) if __name__ == "__main__":