File size: 4,715 Bytes
d8ebb7a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0016651
d8ebb7a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
"""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())