"""Test harness: server lifecycle, HTTP helpers, result recording. Standard library only. The bundle deliberately ships no pytest, so the suites must run against nothing but the interpreter and the installed gradio. Written to run on Python 3.9 as well as 3.12. """ import json import os import socket import subprocess import sys import threading import time import urllib.error import urllib.request HERE = os.path.dirname(os.path.abspath(__file__)) APPS = os.path.join(HERE, "apps") # --- results ----------------------------------------------------------------- class Results: """Collects PASS/FAIL/SKIP rows and measured metrics.""" def __init__(self): self.rows = [] self.metrics = [] def add(self, suite, name, status, detail=""): assert status in ("PASS", "FAIL", "SKIP", "WARN"), status self.rows.append({"suite": suite, "test": name, "status": status, "detail": detail}) line = " %-4s %-46s %s" % (status, name, detail) print(line.rstrip(), flush=True) def metric(self, name, value, unit, note=""): self.metrics.append({"metric": name, "value": value, "unit": unit, "note": note}) shown = ("%.2f" % value) if isinstance(value, float) else str(value) print(" %-52s %10s %s %s" % (name, shown, unit, note), flush=True) @property def failed(self): return sum(1 for r in self.rows if r["status"] == "FAIL") @property def passed(self): return sum(1 for r in self.rows if r["status"] == "PASS") @property def skipped(self): return sum(1 for r in self.rows if r["status"] == "SKIP") @property def warned(self): return sum(1 for r in self.rows if r["status"] == "WARN") def check(results, suite, name): """Context manager: a test passes unless it raises. Raise Skip to record SKIP. Any other exception is a FAIL carrying the exception text, so a suite never dies on one bad case. """ return _Check(results, suite, name) class Skip(Exception): """Raise inside a check to record SKIP: the case does not apply here.""" class Warn(Exception): """Raise inside a check to record WARN. For a known limitation of the installed gradio that the operator should be aware of but which is not a defect in this bundle. Does not fail the run. """ class _Check: def __init__(self, results, suite, name): self.results, self.suite, self.name = results, suite, name self.detail = "" def __enter__(self): return self def __exit__(self, exc_type, exc, tb): if exc_type is None: self.results.add(self.suite, self.name, "PASS", self.detail) elif exc_type is Skip: self.results.add(self.suite, self.name, "SKIP", str(exc)) elif exc_type is Warn: self.results.add(self.suite, self.name, "WARN", str(exc)) else: text = "%s: %s" % (exc_type.__name__, exc) self.results.add(self.suite, self.name, "FAIL", text.replace("\n", " ")[:200]) return True # swallow; the row already records the outcome # --- HTTP -------------------------------------------------------------------- def http(url, data=None, timeout=30, method=None, headers=None): """Return (status, body_text). Never raises: an HTTP error status comes back as that status, and a connection-level failure comes back as status 0 with the error text. """ hdrs = {"Content-Type": "application/json"} if data is not None else {} hdrs.update(headers or {}) body = json.dumps(data).encode() if data is not None else None req = urllib.request.Request(url, data=body, headers=hdrs, method=method) try: with urllib.request.urlopen(req, timeout=timeout) as r: return r.getcode(), r.read().decode("utf-8", "replace") except urllib.error.HTTPError as e: return e.code, e.read().decode("utf-8", "replace") except (urllib.error.URLError, socket.timeout, OSError) as e: # Status 0 means "did not get an HTTP response at all". Callers poll # this while a server is still starting, so a refused connection must # be an ordinary return value rather than an exception. return 0, str(e) def http_json(url, data=None, timeout=30): code, text = http(url, data=data, timeout=timeout) try: return code, json.loads(text) except ValueError: return code, None def free_port(): s = socket.socket() s.bind(("127.0.0.1", 0)) port = s.getsockname()[1] s.close() return port # --- server ------------------------------------------------------------------ class Server: """Launch a gradio app as a subprocess and wait until it answers. Used as a context manager so a failing test cannot leak a live server into the next one. """ def __init__(self, app, python=None, port=None, env=None, timeout=90): self.app = app if os.path.isabs(app) else os.path.join(APPS, app) self.python = python or sys.executable self.port = port or free_port() self.timeout = timeout self.proc = None self.log_path = None self.startup_seconds = None self.env = dict(os.environ) # Offline defaults: without these each launch stalls on outbound calls # that cannot fail fast on an isolated host. self.env.update({ "GRADIO_ANALYTICS_ENABLED": "False", "DO_NOT_TRACK": "1", "HF_HUB_OFFLINE": "1", "HF_HUB_DISABLE_TELEMETRY": "1", "MPLBACKEND": "Agg", "no_proxy": "127.0.0.1,localhost,::1", "NO_PROXY": "127.0.0.1,localhost,::1", }) self.env.update(env or {}) @property def url(self): return "http://127.0.0.1:%d" % self.port def start(self): import tempfile fd, self.log_path = tempfile.mkstemp(prefix="gradio-test-", suffix=".log") self.logfile = os.fdopen(fd, "w+") began = time.time() self.proc = subprocess.Popen( [self.python, self.app, "--port", str(self.port)], stdout=self.logfile, stderr=subprocess.STDOUT, env=self.env, cwd=os.path.dirname(self.app)) deadline = began + self.timeout while time.time() < deadline: if self.proc.poll() is not None: raise RuntimeError("server exited early:\n%s" % self.log()[-1500:]) code, _ = http(self.url + "/", timeout=3) if code == 200: self.startup_seconds = time.time() - began return self time.sleep(0.25) raise RuntimeError("server did not answer within %ss:\n%s" % (self.timeout, self.log()[-1500:])) def log(self): try: self.logfile.flush() with open(self.log_path) as fh: return fh.read() except Exception: return "" def stop(self): if self.proc and self.proc.poll() is None: self.proc.terminate() try: self.proc.wait(timeout=15) except subprocess.TimeoutExpired: self.proc.kill() self.proc.wait(timeout=10) try: self.logfile.close() except Exception: pass def __enter__(self): return self.start() def __exit__(self, *exc): self.stop() if self.log_path and os.path.exists(self.log_path): os.unlink(self.log_path) return False # --- gradio version differences ---------------------------------------------- def api_root(major): """Route prefix for the HTTP API. gradio 5 moved it under /gradio_api.""" return "" if major < 5 else "/gradio_api" def call_endpoint(server, major, name, payload, timeout=60): """Invoke a named endpoint over HTTP and return the output list. gradio 4 exposes a blocking POST /api/. gradio 5+ splits it into a POST that returns an event id and a GET that streams the result, so the two paths differ enough to be worth hiding here. """ if major < 5: code, body = http_json("%s/api/%s" % (server.url, name), data={"data": payload}, timeout=timeout) if code != 200 or not body: raise AssertionError("POST /api/%s -> %s" % (name, code)) return body["data"] code, body = http_json("%s/gradio_api/call/%s" % (server.url, name), data={"data": payload}, timeout=timeout) if code != 200 or not body or "event_id" not in body: raise AssertionError("POST call/%s -> %s %s" % (name, code, body)) event = body["event_id"] code, text = http("%s/gradio_api/call/%s/%s" % (server.url, name, event), timeout=timeout) if code != 200: raise AssertionError("GET call/%s/%s -> %s" % (name, event, code)) # Server-sent events: take the payload of the final complete/data frame. out = None for line in text.splitlines(): if line.startswith("data:"): try: out = json.loads(line[5:].strip()) except ValueError: pass if out is None: raise AssertionError("no data frame in stream for %s" % name) return out def threaded(fn, count, workers): """Run fn(i) count times across workers threads; return list of results. A tiny thread pool rather than concurrent.futures so behaviour is identical on 3.9 and 3.12 and failures surface as values, not swallowed exceptions. """ results = [None] * count lock = threading.Lock() nxt = [0] def worker(): while True: with lock: i = nxt[0] nxt[0] += 1 if i >= count: return try: results[i] = ("ok", fn(i)) except Exception as exc: # recorded, not raised results[i] = ("err", "%s: %s" % (type(exc).__name__, exc)) threads = [threading.Thread(target=worker) for _ in range(workers)] for t in threads: t.start() for t in threads: t.join() return results