mirror of
https://github.com/vladmandic/automatic
synced 2026-09-11 07:18:44 +02:00
new server settings section
Signed-off-by: vladmandic <mandic00@live.com>
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
+38
-19
@@ -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)
|
||||
|
||||
+19
-9
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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("<h2>Torch Options</h2>", "", gr.HTML),
|
||||
|
||||
@@ -15,6 +15,7 @@ pyyaml
|
||||
toml
|
||||
voluptuous
|
||||
fasteners
|
||||
limits
|
||||
orjson
|
||||
websockets
|
||||
|
||||
|
||||
Reference in New Issue
Block a user