diff --git a/launch.py b/launch.py index 7d44a75e0..b8d0bfc53 100755 --- a/launch.py +++ b/launch.py @@ -311,16 +311,18 @@ def main(): alive = False requests = 0 t_current = time.time() - if float(args.status) > 0 and (t_current - t_server) > float(args.status): + status_rate = float(args.status) if args.status >= 0 else installer.opts.get('server_status', 120) + monitor_rate = float(args.monitor) if args.monitor >= 0 else installer.opts.get('server_monitor', 0) + if float(status_rate) > 0 and (t_current - t_server) > float(status_rate): s = instance.state.status() if (s.timestamp is None) or (s.step == 0): # dont spam during active job log.trace(f'Server: alive={alive} requests={requests} memory={get_memory_stats()} {s}') t_server = t_current - if float(args.monitor) > 0 and t_current - t_monitor > float(args.monitor): + if float(monitor_rate) > 0 and t_current - t_monitor > float(monitor_rate): log.trace(f'Monitor: {get_memory_stats(detailed=True)}') t_monitor = t_current - from modules.api.validate import get_stats - get_stats() + from modules.api.validate import get_api_stats + get_api_stats() if not alive: if uv is not None and uv.wants_restart: clean_server() diff --git a/modules/api/server.py b/modules/api/server.py index 7c3448a36..b5f209e2e 100644 --- a/modules/api/server.py +++ b/modules/api/server.py @@ -97,7 +97,7 @@ def get_history(req: models.ReqHistory = Depends()): return res def get_progress(req: models.ReqProgress = Depends()): - if shared.state.job_count == 0: # idle state + if shared.state.job_count == 0 and shared.state.sampling_step == 0: # truly idle return models.ResProgress(id=shared.state.id, progress=0, eta_relative=0, state=shared.state.dict(), textinfo=shared.state.textinfo) shared.state.do_set_current_image() current_image = None diff --git a/modules/api/validate.py b/modules/api/validate.py index 2ee26e35e..054d8f18a 100644 --- a/modules/api/validate.py +++ b/modules/api/validate.py @@ -1,10 +1,7 @@ import re -import limits -from fastapi.exceptions import HTTPException from modules.logger import log -requests_summary = {} request_cost = { # value is cost, 0=not rate limited, 1=default, >1 more expensive "/file": 0, "/run/predict": 0, @@ -14,28 +11,50 @@ request_cost = { # value is cost, 0=not rate limited, 1=default, >1 more expens "/sdapi/v1/img2img": 5, "/sdapi/v1/control": 5, } -backend = limits.storage.MemoryStorage() -strategy = limits.strategies.SlidingWindowCounterRateLimiter(backend) -limiter = limits.parse("300/minute") -def get_stats(): - for k, v in requests_summary.items(): - if v > 1: - log.trace(f'API stats: {k}={v}') +class Limiter(): + def __init__(self, limit): + import limits + self.limit = limit + self.backend = limits.storage.MemoryStorage() + self.strategy = limits.strategies.SlidingWindowCounterRateLimiter(self.backend) + self.limiter = limits.parse(f'{self.limit}/minute') + self.summary = {} + log.info(f'API: limit={self.limit} strategy={self.strategy.__class__.__name__} backend={self.backend.__class__.__name__}') + + def stats(self): + for k, v in self.summary.items(): + if v > 1: + log.trace(f'API stats: {k}={v}') + + def check(self, key, quiet: bool = False): + if self.limit <= 0: + return True + cost = request_cost.get(key, 1) + status = self.strategy.hit(self.limiter, key, cost=cost) + if not status and not quiet: + log.warning(f'API: key={key} rate limit exceeded') + from fastapi.exceptions import HTTPException + raise HTTPException(status_code=429, detail=f"{key}: rate limit exceeded") + return status -def rate_limit(key): - cost = request_cost.get(key, 1) - if not strategy.hit(limiter, key, cost=cost): - log.warning(f'API: key={key} rate limit exceeded') - raise HTTPException(status_code=429, detail=f'{key}: rate limit exceeded') +limiter = Limiter(300) + + +def get_api_stats(): + limiter.stats() def validate_request(client, endpoint): + global limiter # pylint: disable=global-statement + from modules.shared import opts + if opts.server_rate_limit != limiter.limit: + limiter = Limiter(opts.server_rate_limit) api = re.match(r"^[^?#&=]+", endpoint).group(0) key = f"{client}:{api}" - if key not in requests_summary: - requests_summary[key] = 0 - requests_summary[key] += 1 - rate_limit(key) + if key not in limiter.summary: + limiter.summary[key] = 0 + limiter.summary[key] += 1 + return limiter.check(key) diff --git a/modules/cmd_args.py b/modules/cmd_args.py index c706a9227..35b3cfc3b 100644 --- a/modules/cmd_args.py +++ b/modules/cmd_args.py @@ -18,12 +18,11 @@ def get_argv(): def add_core_args(p): p.add_argument("--ckpt", type=str, default=os.environ.get("SD_MODEL", None), help="Path to model checkpoint to load immediately, default: %(default)s") p.add_argument("--data-dir", type=str, default=os.environ.get("SD_DATADIR", ''), help="Base path where all user data is stored, default: %(default)s") - p.add_argument("--models-dir", type=str, default=os.environ.get("SD_MODELSDIR", None), help="Base path where all models are stored, default: %(default)s",) - p.add_argument("--embeddings-dir", type=str, default=os.environ.get("SD_EMBEDDINGSDIR", None), help="Base path where all embeddings are stored, default: %(default)s",) - p.add_argument("--hypernetwork-dir", type=str, default=os.environ.get("SD_HYPERNETWORKDIR", None), help="Base path where all hypernetworks are stored, default: %(default)s",) - p.add_argument("--vae-dir", type=str, default=os.environ.get("SD_VAEDIR", None), help="Base path where all VAEs are stored, default: %(default)s",) - p.add_argument("--lora-dir", type=str, default=os.environ.get("SD_LORADIR", None), help="Base path where all LoRAs are stored, default: %(default)s",) - p.add_argument("--extensions-dir", type=str, default=os.environ.get("SD_EXTENSIONSDIR", None), help="Base path where all extensions are stored, default: %(default)s",) + p.add_argument("--models-dir", type=str, default=os.environ.get("SD_MODELSDIR", None), help="Base path where all models are stored, default: %(default)s") + p.add_argument("--embeddings-dir", type=str, default=os.environ.get("SD_EMBEDDINGSDIR", None), help="Base path where all embeddings are stored, default: %(default)s") + p.add_argument("--vae-dir", type=str, default=os.environ.get("SD_VAEDIR", None), help="Base path where all VAEs are stored, default: %(default)s") + p.add_argument("--lora-dir", type=str, default=os.environ.get("SD_LORADIR", None), help="Base path where all LoRAs are stored, default: %(default)s") + p.add_argument("--extensions-dir", type=str, default=os.environ.get("SD_EXTENSIONSDIR", None), help="Base path where all extensions are stored, default: %(default)s") def add_config_arg(p, data_dir): @@ -60,7 +59,6 @@ def add_http_args(p): p.add_argument("--auth", type=str, default=os.environ.get("SD_AUTH", None), help='Set access authentication like "user:pwd,user:pwd""') p.add_argument("--auth-file", type=str, default=os.environ.get("SD_AUTHFILE", None), help='Set access authentication using file, default: %(default)s') p.add_argument("--allowed-paths", nargs='+', default=[], type=str, required=False, help="add additional paths to paths allowed for web access") - p.add_argument("--share", default=env_flag("SD_SHARE", False), action='store_true', help="Enable UI accessible through Gradio site, default: %(default)s") p.add_argument("--insecure", default=env_flag("SD_INSECURE", False), action='store_true', help="Enable extensions tab regardless of other options, default: %(default)s") p.add_argument("--listen", default=env_flag("SD_LISTEN", False), action='store_true', help="Launch web server using public IP address, default: %(default)s") p.add_argument("--port", type=int, default=os.environ.get("SD_PORT", 7860), help="Launch web server with given server port, default: %(default)s") @@ -73,8 +71,8 @@ def add_diag_args(p): p.add_argument('--safe', default=env_flag("SD_SAFE", False), action='store_true', help="Run in safe mode with no user extensions") p.add_argument('--test', default=env_flag("SD_TEST", False), action='store_true', help="Run test only and exit") p.add_argument('--version', default=False, action='store_true', help="Print version information") - p.add_argument("--monitor", default=os.environ.get("SD_MONITOR", 0), help="Run memory monitor, default: %(default)s") - p.add_argument("--status", default=os.environ.get("SD_STATUS", 120), help="Run server is-alive status, default: %(default)s") + p.add_argument("--monitor", default=os.environ.get("SD_MONITOR", -1), help="Run memory monitor, default: %(default)s") + p.add_argument("--status", default=os.environ.get("SD_STATUS", -1), help="Run server is-alive status, default: %(default)s") def add_log_args(p): @@ -157,6 +155,8 @@ def compatibility_args(): group_compat.add_argument("--no-metadata", default=env_flag("SD_NOMETADATA", False), action='store_true', help=argparse.SUPPRESS) group_compat.add_argument("--precision", type=str, choices=["full", "autocast"], default="autocast", help=argparse.SUPPRESS) group_compat.add_argument("--upcast-sampling", default=env_flag("SD_UPCASTSAMPLING", False), action='store_true', help=argparse.SUPPRESS) + group_compat.add_argument("--hypernetwork-dir", type=str, default=os.environ.get("SD_HYPERNETWORKDIR", None), help=argparse.SUPPRESS) + group_compat.add_argument("--share", default=env_flag("SD_SHARE", False), action="store_true", help=argparse.SUPPRESS) def settings_args(opts, args): @@ -211,3 +211,13 @@ def settings_args(opts, args): opts.onchange(d, lambda d=d: setattr(args, d, getattr(opts, d)), call=False) return args + + +def override_args(opts, args): + if opts.server_listen: + args.listen = True + if opts.server_status >= 0: + args.status = opts.server_status + if opts.server_monitor >= 0: + args.monitor = opts.server_monitor + return args diff --git a/modules/shared.py b/modules/shared.py index b07a1f34a..d6d1e043a 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -164,6 +164,7 @@ from modules.options_handler import Options config_filename = cmd_opts.config opts = Options(options_templates, restricted_opts, filename=config_filename) cmd_opts = cmd_args.settings_args(opts, cmd_opts) +cmd_opts = cmd_args.override_args(opts, cmd_opts) if cmd_opts.locale is not None: opts.data['ui_locale'] = cmd_opts.locale diff --git a/modules/ui_definitions.py b/modules/ui_definitions.py index 86f2fd74c..a4a9c8e3d 100644 --- a/modules/ui_definitions.py +++ b/modules/ui_definitions.py @@ -251,6 +251,14 @@ def create_settings(cmd_opts): "dynamic_attention_trigger_rate": OptionInfo(1, "Dynamic Attention trigger rate", gr.Slider, {"minimum": 0.01, "maximum": max(gpu_memory,4)*2, "step": 0.01}), })) + # --- Server Settings --- + options_templates.update(options_section(('server', "Server Settings"), { + "server_listen": OptionInfo(False, "Listen on all interfaces", gr.Checkbox), + "server_status": OptionInfo(120, "Automatic server status monitor rate", gr.Number, {"minimum": 0, "maximum": 1000, "step": 1}), + "server_monitor": OptionInfo(0, "Automatic server memory monitor rate", gr.Number, {"minimum": 0, "maximum": 1000, "step": 1}), + "server_rate_limit": OptionInfo(300, "API base rate limit rate", gr.Number, {"minimum": 0, "maximum": 1000, "step": 1}), + })) + # --- Backend Settings --- options_templates.update(options_section(('backends', "Backend Settings"), { "other_sep": OptionInfo("