image-lab / scripts /load_test.py
Magnus Persson
Add the ideas popup with cached example images; drop the format picker
0016651
Raw History Blame Contribute Delete
4.72 kB
"""Load test for the Create flow against a running app (fakes only, no keys).
uv run python scripts/load_test.py --url https://magnusp-image-lab.hf.space
uv run python scripts/load_test.py --url http://127.0.0.1:7860 --duration 30
The target needs `access.expose_api: true` so the Gradio client may call events; turn it off
again afterwards. When the app asks for a login, set LOAD_TEST_PASSWORD (the workshop password;
never pass it on the command line).
Each worker is one simulated visitor with its own Gradio session and device id. Exit code 1 if any
Create fails or takes longer than --stall-seconds.
"""
from __future__ import annotations
import argparse
import os
import random
import statistics
import sys
import threading
import time
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass, field
from gradio_client import Client
API_NAME = "/on_create"
IDEAS = (
"a dragon that loves pancakes",
"en robot som spelar fotboll",
"a treehouse on the moon",
"en katt i en rymddräkt",
"a castle made of ice cream",
)
@dataclass
class Stats:
latencies: list[float] = field(default_factory=list)
failures: list[str] = field(default_factory=list)
stalls: int = 0
lock: threading.Lock = field(default_factory=threading.Lock)
def record(self, seconds: float, error: str | None, stalled: bool) -> None:
with self.lock:
if error:
self.failures.append(error)
else:
self.latencies.append(seconds)
self.stalls += stalled
def one_create(client: Client, device: object) -> tuple[str | None, object]:
"""Returns (error description or None, device id to send next time)."""
idea = random.choice(IDEAS)
new_device, status, image = client.predict(idea, device, api_name=API_NAME)
# Outputs are Gradio updates: a refused Create still returns a truthy dict, without a value.
if not (isinstance(image, dict) and image.get("value")):
message = status.get("value") if isinstance(status, dict) else status
return f"no image: {str(message)[:80]}", new_device
return None, new_device
def worker(args: argparse.Namespace, deadline: float, stats: Stats) -> None:
password = os.environ.get("LOAD_TEST_PASSWORD")
auth = (args.username, password) if password else None
client = Client(args.url, auth=auth, verbose=False)
device: object = None
while time.monotonic() < deadline:
start = time.monotonic()
try:
error, device = one_create(client, device)
except Exception as exc: # any failure counts against the run
error = f"{type(exc).__qualname__}: {str(exc).split(' for url')[0][:120]}"
elapsed = time.monotonic() - start
stats.record(elapsed, error, elapsed > args.stall_seconds)
time.sleep(random.uniform(args.min_think_seconds, args.think_seconds))
def percentile(values: list[float], pct: float) -> float:
ordered = sorted(values)
return ordered[min(len(ordered) - 1, int(len(ordered) * pct))]
def main() -> int:
parser = argparse.ArgumentParser(description=__doc__.splitlines()[0])
parser.add_argument("--url", required=True)
parser.add_argument("--workers", type=int, default=10)
parser.add_argument("--duration", type=float, default=600, help="seconds")
parser.add_argument("--think-seconds", type=float, default=30, help="max pause between Creates")
parser.add_argument(
"--min-think-seconds", type=float, default=10, help="min pause; keep it above the cooldown"
)
parser.add_argument("--username", default="workshop")
parser.add_argument("--stall-seconds", type=float, default=30)
args = parser.parse_args()
stats = Stats()
deadline = time.monotonic() + args.duration
print(f"{args.workers} visitors for {args.duration:.0f}s against {args.url}")
with ThreadPoolExecutor(args.workers) as pool:
for _ in range(args.workers):
pool.submit(worker, args, deadline, stats)
total = len(stats.latencies) + len(stats.failures)
print(f"creates: {total} ok: {len(stats.latencies)} failed: {len(stats.failures)}")
print(f"stalls (> {args.stall_seconds:.0f}s): {stats.stalls}")
if stats.latencies:
print(
f"latency s: median {statistics.median(stats.latencies):.1f} "
f"p95 {percentile(stats.latencies, 0.95):.1f} max {max(stats.latencies):.1f}"
)
for message in sorted(set(stats.failures))[:5]:
print(f" {stats.failures.count(message)}x {message[:200]}")
return 1 if stats.failures or stats.stalls or not total else 0
if __name__ == "__main__":
sys.exit(main())