#!/usr/bin/env python3
"""Interactive live plot of prefix-sums execution time vs. n.

Usage:
  ./show-times [options]

Example:
  ./show-times
  ./show-times --procs 8 --warmup 1

Opens a browser window with a plot and a panel listing every impl in main.sml.
Each impl has buttons to run it (sweeping n, adding each point to the plot as
soon as it has been measured), clear its measurements, and show/hide it.
Runs happen one at a time; clicking run while another impl is running queues
it. Once an impl takes longer than --max-time at some n, its larger n are
skipped.
"""

import argparse
import http.server
import json
import os
import re
import subprocess
import sys
import threading
import urllib.parse
import webbrowser

# impls in main.sml left out of the panel
EXCLUDE = set()

# each of these gets its own entry in the panel, run with -chunk-size K
CHUNK_SIZES = {"chunked-contraction": [2, 10, 1000], "scan": [1000]}

NS = [100, 500, 1000, 5000, 10000, 50000, 100000, 500000, 1000000, 5000000, 10000000]
HERE = os.path.dirname(os.path.abspath(__file__))


def find_impls():
    """The impl names in main.sml's `case impl of ...`, in order."""
    with open(os.path.join(HERE, "main.sml")) as f:
        return re.findall(r'^\s*\|?\s*"([^"]+)"\s*=>', f.read(), re.MULTILINE)


# c5 (cyan) is left out: too close to the k0..k5 teal shades for chunk sizes
BASE_COLORS = ["c0", "c1", "c2", "c3", "c4", "c6", "c7"]
NUM_SHADES = 6


def make_entries(impls):
    """Panel entries: name -> (impl, chunk size or None), in order, plus each
    entry's color. An impl with several chunk sizes gets shades of one color."""
    entries, colors, c = {}, {}, 0
    for impl in impls:
        if impl in EXCLUDE:
            continue
        sizes = CHUNK_SIZES.get(impl, [None])
        for i, k in enumerate(sizes):
            name = impl if k is None else f"{impl} k={k}"
            entries[name] = (impl, k)
            if len(sizes) > 1:
                # spread the shades k0..k5 across the sizes
                colors[name] = f"k{round(i * (NUM_SHADES - 1) / (len(sizes) - 1))}"
            else:
                colors[name] = BASE_COLORS[c % len(BASE_COLORS)]
                c += 1
    return entries, colors


# ---------------------------------------------------------------------------
# event log shared with the browser (replayed in full on every (re)connect)

events = []
events_cv = threading.Condition()


def emit(**msg):
    with events_cv:
        events.append(msg)
        events_cv.notify_all()


# ---------------------------------------------------------------------------
# run queue: one sweep at a time, so runs don't compete for cores

class Runner:
    def __init__(self, args, entries):
        self.args = args
        self.entries = entries
        self.lock = threading.Condition()
        self.queue = []          # impls waiting to run
        self.running = None      # impl currently running
        self.cancelled = False   # set when the running impl is cleared
        self.proc = None         # its current subprocess

    def run(self, impl):
        with self.lock:
            if impl == self.running or impl in self.queue:
                return
            self.queue.append(impl)
            emit(type="queued", impl=impl)
            self.lock.notify_all()

    def clear(self, impl):
        with self.lock:
            if impl in self.queue:
                self.queue.remove(impl)
            if impl == self.running:
                self.cancelled = True
                if self.proc is not None:
                    self.proc.kill()
            emit(type="clear", impl=impl)

    def worker(self):
        while True:
            with self.lock:
                while not self.queue:
                    self.lock.wait()
                impl = self.running = self.queue.pop(0)
                self.cancelled = False
                emit(type="clear", impl=impl)   # a new run replaces old data
            self.sweep(impl)
            with self.lock:
                if not self.cancelled:
                    emit(type="finished", impl=impl)
                self.running = None

    def sweep(self, impl):
        for n in NS:
            with self.lock:
                if self.cancelled:
                    return
                emit(type="status", impl=impl, n=n)
            try:
                t = self.run_one(impl, n)
            except RuntimeError as e:
                with self.lock:
                    if self.cancelled:
                        return
                    print(e, file=sys.stderr)
                    emit(type="error", impl=impl, n=n, msg=str(e).splitlines()[0])
                return
            with self.lock:
                if self.cancelled:
                    return
                print(f"{impl:>28}  n={n:<8}  {t:.6f}s", flush=True)
                emit(type="point", impl=impl, n=n, time=t)
            if t > self.args.max_time and n != NS[-1]:
                print(f"{impl:>28}  skipping larger n (> {self.args.max_time}s)", flush=True)
                emit(type="skip", impl=impl, n=n)
                return

    def run_one(self, name, n):
        a = self.args
        impl, chunk_size = self.entries[name]
        cmd = [os.path.join(HERE, "main"), "-impl", impl, "-n", str(n),
               "-repeat", str(a.repeat), "-warmup", str(a.warmup)]
        if chunk_size is not None:
            cmd += ["-chunk-size", str(chunk_size)]
        if a.procs is not None:
            cmd += ["@mpl", "procs", str(a.procs), "--"]
        with self.lock:
            if self.cancelled:
                raise RuntimeError("cancelled")
            self.proc = subprocess.Popen(cmd, stdout=subprocess.PIPE,
                                         stderr=subprocess.PIPE, text=True)
        out, err = self.proc.communicate()
        code = self.proc.returncode
        with self.lock:
            self.proc = None
        if code != 0:
            raise RuntimeError(f"{' '.join(cmd)} failed:\n{out}{err}")
        m = re.search(r"^average\s+([0-9.]+)s", out, re.MULTILINE)
        if not m:
            raise RuntimeError(f"could not parse output of {' '.join(cmd)}:\n{out}")
        return float(m.group(1))


runner = None


# ---------------------------------------------------------------------------
# HTTP: the page, the event stream, and run/clear requests

class Handler(http.server.BaseHTTPRequestHandler):
    def log_message(self, *args):
        pass

    def do_GET(self):
        if self.path == "/":
            body = PAGE.encode()
            self.send_response(200)
            self.send_header("Content-Type", "text/html; charset=utf-8")
            self.send_header("Content-Length", str(len(body)))
            self.end_headers()
            self.wfile.write(body)
        elif self.path == "/events":
            self.send_response(200)
            self.send_header("Content-Type", "text/event-stream")
            self.send_header("Cache-Control", "no-cache")
            self.end_headers()
            i = 0
            try:
                while True:
                    with events_cv:
                        while i >= len(events):
                            events_cv.wait()
                        batch = events[i:]
                        i = len(events)
                    for msg in batch:
                        self.wfile.write(b"data: " + json.dumps(msg).encode() + b"\n\n")
                    self.wfile.flush()
            except (BrokenPipeError, ConnectionResetError):
                pass
        else:
            self.send_error(404)

    def do_POST(self):
        url = urllib.parse.urlparse(self.path)
        impl = urllib.parse.parse_qs(url.query).get("impl", [None])[0]
        if impl not in runner.entries:
            self.send_error(400, "unknown impl")
            return
        if url.path == "/run":
            runner.run(impl)
        elif url.path == "/clear":
            runner.clear(impl)
        else:
            self.send_error(404)
            return
        self.send_response(204)
        self.end_headers()


def main():
    global runner
    p = argparse.ArgumentParser(
        description="Interactive live plot of execution time vs. n.")
    p.add_argument("--procs", type=int, default=None,
                   help="number of processors (@mpl procs P --)")
    p.add_argument("--repeat", type=int, default=5,
                   help="repetitions per point; the average is plotted (default 5)")
    p.add_argument("--warmup", type=float, default=0.0,
                   help="seconds of warmup runs before each point is measured (default 0)")
    p.add_argument("--max-time", type=float, default=2.0,
                   help="skip larger n for an impl once it exceeds this many seconds (default 2)")
    p.add_argument("--linear", action="store_true",
                   help="start with linear axes instead of log-log (toggle in the page)")
    p.add_argument("--port", type=int, default=0,
                   help="HTTP port (default: any free port)")
    p.add_argument("--no-browser", action="store_true",
                   help="don't open a browser window, just print the URL")
    args = p.parse_args()

    subprocess.run(["make", "-C", HERE, "main"], check=True)

    entries, colors = make_entries(find_impls())
    runner = Runner(args, entries)
    threading.Thread(target=runner.worker, daemon=True).start()

    opts = [f"procs={args.procs if args.procs is not None else 'default'}",
            f"repeat={args.repeat}", f"warmup={args.warmup:g}s"]
    emit(type="config", impls=list(entries), colors=colors, ns=NS, opts=", ".join(opts),
         linear=args.linear)

    server = http.server.ThreadingHTTPServer(("127.0.0.1", args.port), Handler)
    server.daemon_threads = True
    url = f"http://127.0.0.1:{server.server_address[1]}/"
    print(f"plot at {url}  (Ctrl-C to quit)", flush=True)
    threading.Thread(target=server.serve_forever, daemon=True).start()
    if not args.no_browser:
        webbrowser.open(url)

    try:
        threading.Event().wait()
    except KeyboardInterrupt:
        with runner.lock:
            if runner.proc is not None:
                runner.proc.kill()


# ---------------------------------------------------------------------------
# the page: a hand-drawn SVG plot (log-log or linear), no external dependencies

PAGE = r"""<!doctype html>
<html>
<head>
<meta charset="utf-8">
<meta name="viewport" content="width=device-width, initial-scale=1">
<title>Prefix Sums Times</title>
<style>
  :root {
    --bg: #ffffff; --fg: #1f2328; --muted: #6e7781; --grid: #e5e7eb; --axis: #9ca3af;
    --btn: #f3f4f6; --btn-on: #1f2328; --btn-on-fg: #ffffff; --sel: rgba(37, 99, 235, 0.12);
    --panel: #f9fafb; --err: #b91c1c;
    --c0: #2563eb; --c1: #dc2626; --c2: #16a34a; --c3: #9333ea;
    --c4: #ea580c; --c5: #0891b2; --c6: #ca8a04; --c7: #db2777;
    --k0: #2dd4bf; --k1: #14b8a6; --k2: #0d9488; --k3: #0f766e; --k4: #115e59; --k5: #134e4a;
  }
  @media (prefers-color-scheme: dark) {
    :root {
      --bg: #0d1117; --fg: #e6edf3; --muted: #8b949e; --grid: #21262d; --axis: #484f58;
      --btn: #21262d; --btn-on: #e6edf3; --btn-on-fg: #0d1117; --sel: rgba(96, 165, 250, 0.18);
      --panel: #161b22; --err: #f87171;
      --c0: #60a5fa; --c1: #f87171; --c2: #4ade80; --c3: #c084fc;
      --c4: #fb923c; --c5: #22d3ee; --c6: #facc15; --c7: #f472b6;
      --k0: #0d9488; --k1: #14b8a6; --k2: #2dd4bf; --k3: #5eead4; --k4: #99f6e4; --k5: #ccfbf1;
    }
  }
  body { margin: 0; background: var(--bg); color: var(--fg);
         font: 16px/1.4 system-ui, -apple-system, sans-serif; }
  main { max-width: 1200px; margin: 0 auto; padding: 20px 16px; }
  h1 { font-size: 22px; margin: 0 0 2px; }
  #opts, #status, .hint { color: var(--muted); font-size: 14px; }
  #status { margin-top: 6px; min-height: 1.4em; }
  .controls { display: flex; flex-wrap: wrap; align-items: center; gap: 8px 14px; margin-top: 10px; }
  .seg { display: inline-flex; border: 1px solid var(--axis); border-radius: 6px; overflow: hidden; }
  button { font: inherit; font-size: 14px; color: var(--fg); background: var(--btn);
           border: 0; padding: 4px 12px; cursor: pointer; }
  button:disabled { opacity: 0.4; cursor: default; }
  .seg button + button { border-left: 1px solid var(--axis); }
  .seg button[aria-pressed="true"] { background: var(--btn-on); color: var(--btn-on-fg); }
  #reset, .impl button { border: 1px solid var(--axis); border-radius: 6px; }
  #reset[hidden] { display: none; }
  .layout { display: flex; gap: 20px; align-items: flex-start; margin-top: 8px; }
  .plotwrap { flex: 1; min-width: 0; }
  svg { width: 100%; height: auto; display: block;
        touch-action: none; user-select: none; -webkit-user-select: none; }
  svg text { fill: var(--fg); font-size: 13px; }
  .grid { stroke: var(--grid); }
  .axis { stroke: var(--axis); }
  .lbl { fill: var(--muted); }
  .plotarea { cursor: crosshair; }
  .selrect { fill: var(--sel); stroke: var(--c0); stroke-dasharray: 4 3; }
  #impls { width: 290px; flex: none; background: var(--panel); border: 1px solid var(--grid);
           border-radius: 8px; padding: 6px 12px; }
  .impl { padding: 8px 0; }
  .impl + .impl { border-top: 1px solid var(--grid); }
  .impl .head { display: flex; flex-wrap: wrap; align-items: center; gap: 2px 8px; }
  .impl .name { font-weight: 600; font-size: 15px; }
  .impl .state { color: var(--muted); font-size: 13px; margin-left: auto; text-align: right; }
  .impl .state.err { color: var(--err); }
  .impl .btns { display: flex; gap: 6px; margin-top: 6px; }
  .impl .btns button { padding: 2px 10px; font-size: 13px; }
  .impl.hidden .name, .impl.hidden .swatch { opacity: 0.35; }
  #more { display: block; width: 100%; margin: 8px 0 4px; padding: 2px 0; font-size: 18px;
          line-height: 1.3; border: 1px dashed var(--axis); border-radius: 6px; }
  #more[hidden] { display: none; }
  .swatch { width: 12px; height: 12px; border-radius: 50%; flex: none; }
  @media (max-width: 760px) {
    .layout { flex-direction: column; }
    #impls { width: auto; align-self: stretch; }
  }
</style>
</head>
<body>
<main>
  <h1>Prefix sums: time vs. n</h1>
  <div id="opts"></div>
  <div class="controls">
    <div class="seg" role="group" aria-label="axis scale">
      <button data-scale="log" aria-pressed="true">log-log</button>
      <button data-scale="linear" aria-pressed="false">linear</button>
    </div>
    <button id="reset" hidden>reset zoom</button>
    <span class="hint">drag on the plot to zoom &middot; double-click or Esc to reset</span>
  </div>
  <div id="status">connecting&hellip;</div>
  <div class="layout">
    <div class="plotwrap"><svg id="plot" viewBox="0 0 800 540"></svg></div>
    <div id="impls"></div>
  </div>
</main>
<script>
const W = 800, H = 540, M = {l: 70, r: 20, t: 16, b: 50};
const PW = W - M.l - M.r, PH = H - M.t - M.b;
let impls = [], ns = [], data = {}, st = {};
let linear = false, configured = false;
let hidden = {};   // per-viewer: impls not drawn
let revealed = 1;  // per-viewer: how many impls the panel shows; "+" reveals the next
let zoom = null;   // null = auto range, else {x: [lo, hi], y: [lo, hi]} in data units
let drag = null;   // {x0, y0, x1, y1} in svg coords while dragging
let sx = null, sy = null;   // current scales, kept for mapping drags back to data

let colors = {};
const color = impl => `var(--${colors[impl]})`;
const clean = v => Number(v.toPrecision(6));
function fmtN(n) {
  n = clean(n);
  if (n >= 1e6) return clean(n / 1e6) + "M";
  if (n >= 1e3) return clean(n / 1e3) + "K";
  return "" + n;
}
function fmtT(t) {
  t = clean(t);
  if (t === 0) return "0";
  if (t >= 1) return t + "s";
  if (t >= 1e-3) return clean(t * 1e3) + "ms";
  if (t >= 1e-6) return clean(t * 1e6) + "µs";
  return clean(t * 1e9) + "ns";
}

// ---- scales and ticks ------------------------------------------------------

function scale(log, lo, hi, p0, p1) {
  const tf = log ? Math.log10 : (v => v), itf = log ? (v => Math.pow(10, v)) : (v => v);
  const a = tf(lo), b = tf(hi);
  return {log, lo, hi,
          f: v => p0 + (tf(v) - a) / (b - a) * (p1 - p0),
          inv: p => itf(a + (p - p0) / (p1 - p0) * (b - a))};
}

function niceStep(span, count) {
  const raw = Math.max(span, 1e-12) / count, p = Math.pow(10, Math.floor(Math.log10(raw)));
  return [1, 2, 5, 10].map(k => k * p).find(s => s >= raw);
}

// multiples of a 1/2/5 step covering [lo, hi], roughly `count` of them
function linTicks(lo, hi, count = 5) {
  const step = niceStep(hi - lo, count);
  const ticks = [];
  for (let v = Math.ceil(lo / step) * step; v <= hi * (1 + 1e-9); v += step)
    ticks.push(Number(v.toPrecision(12)));
  return ticks;
}

// labeled ticks at decades (or 1/2/5 per decade when zoomed in), minor ticks at 2..9
function logTicks(lo, hi) {
  const inr = v => v >= lo * (1 - 1e-9) && v <= hi * (1 + 1e-9);
  const e0 = Math.floor(Math.log10(lo)), e1 = Math.ceil(Math.log10(hi));
  const mults = ks => { const r = [];
    for (let e = e0; e <= e1; e++) for (const k of ks) r.push(Number((k * Math.pow(10, e)).toPrecision(12)));
    return r.filter(inr); };
  let major = mults([1]);
  if (major.length < 3) major = mults([1, 2, 5]);
  if (major.length < 3) major = linTicks(lo, hi);
  const minor = mults([1, 2, 3, 4, 5, 6, 7, 8, 9]).filter(v => !major.includes(v));
  return {major, minor};
}

function autoDomain(pts) {
  const ts = pts.map(p => p.time).filter(t => linear || t > 0);
  if (linear) {
    // round the top up to a tick so the largest point stays in view
    const up = v => { const s = niceStep(v, 5); return Math.ceil(v / s - 1e-9) * s; };
    return {x: [0, up(ns[ns.length - 1])], y: [0, up(ts.length ? Math.max(...ts) : 1e-3)]};
  }
  let ylo = -5, yhi = -2;
  if (ts.length) {
    ylo = Math.floor(Math.log10(Math.min(...ts)));
    yhi = Math.ceil(Math.log10(Math.max(...ts)));
    if (yhi - ylo < 2) yhi = ylo + 2;
  }
  const pad = Math.pow(10, 0.1);
  return {x: [ns[0] / pad, ns[ns.length - 1] * pad], y: [Math.pow(10, ylo), Math.pow(10, yhi)]};
}

// ---- drawing ---------------------------------------------------------------

function el(tag, attrs, text) {
  const e = document.createElementNS("http://www.w3.org/2000/svg", tag);
  for (const k in attrs) e.setAttribute(k, attrs[k]);
  if (text !== undefined) e.textContent = text;
  return e;
}

function draw() {
  const svg = document.getElementById("plot");
  svg.replaceChildren();
  if (!ns.length) return;

  const shown = impls.filter(impl => !hidden[impl]);
  const dom = zoom || autoDomain(shown.flatMap(impl => data[impl] || []));
  sx = scale(!linear, dom.x[0], dom.x[1], M.l, M.l + PW);
  sy = scale(!linear, dom.y[0], dom.y[1], M.t + PH, M.t);
  const X = sx.f, Y = sy.f;

  let xticks, yticks, yminor = [];
  if (linear) {
    xticks = linTicks(dom.x[0], dom.x[1]);
    yticks = linTicks(dom.y[0], dom.y[1]);
  } else {
    xticks = ns.filter(n => n >= dom.x[0] && n <= dom.x[1]);
    if (xticks.length < 2) xticks = logTicks(dom.x[0], dom.x[1]).major;
    ({major: yticks, minor: yminor} = logTicks(dom.y[0], dom.y[1]));
  }

  const defs = el("defs", {});
  const clip = el("clipPath", {id: "clip"});
  clip.append(el("rect", {x: M.l - 6, y: M.t - 6, width: PW + 12, height: PH + 12}));
  defs.append(clip);
  svg.append(defs);
  svg.append(el("rect", {x: M.l, y: M.t, width: PW, height: PH, fill: "transparent", class: "plotarea"}));

  for (const t of yticks) {
    const y = Y(t);
    svg.append(el("line", {x1: M.l, x2: M.l + PW, y1: y, y2: y, class: "grid"}));
    svg.append(el("text", {x: M.l - 8, y: y + 4, "text-anchor": "end", class: "lbl"}, fmtT(t)));
  }
  for (const t of yminor) {
    const y = Y(t);
    svg.append(el("line", {x1: M.l - 3, x2: M.l, y1: y, y2: y, class: "axis"}));
  }
  for (const n of xticks) {
    const x = X(n);
    svg.append(el("line", {x1: x, x2: x, y1: M.t, y2: M.t + PH, class: "grid"}));
    svg.append(el("text", {x: x, y: M.t + PH + 18, "text-anchor": "middle", class: "lbl"}, fmtN(n)));
  }
  svg.append(el("line", {x1: M.l, x2: M.l, y1: M.t, y2: M.t + PH, class: "axis"}));
  svg.append(el("line", {x1: M.l, x2: M.l + PW, y1: M.t + PH, y2: M.t + PH, class: "axis"}));
  svg.append(el("text", {x: M.l + PW / 2, y: H - 8, "text-anchor": "middle"}, "n"));
  svg.append(el("text", {x: 16, y: M.t + PH / 2, "text-anchor": "middle",
                         transform: `rotate(-90 16 ${M.t + PH / 2})`}, "time (s)"));

  const series = el("g", {"clip-path": "url(#clip)"});
  for (const impl of shown) {
    const pts = (data[impl] || []).filter(p => linear || p.time > 0);
    if (pts.length > 1)
      series.append(el("polyline", {
        points: pts.map(p => `${X(p.n)},${Y(p.time)}`).join(" "),
        fill: "none", stroke: color(impl), "stroke-width": 2.5, "stroke-linejoin": "round"}));
    for (const p of pts) {
      const c = el("circle", {cx: X(p.n), cy: Y(p.time), r: 4.5, fill: color(impl)});
      c.append(el("title", {}, `${impl}  n=${p.n}  ${p.time.toFixed(6)}s`));
      series.append(c);
    }
  }
  svg.append(series);

  if (drag) {
    const r = selection(drag);
    svg.append(el("rect", {x: r.x0, y: r.y0, width: r.x1 - r.x0, height: r.y1 - r.y0, class: "selrect"}));
  }
}

// ---- impl panel ------------------------------------------------------------

function post(action, impl) {
  fetch(`/${action}?impl=${encodeURIComponent(impl)}`, {method: "POST"})
    .catch(() => status("could not reach the script (stopped?)"));
}

function buildPanel() {
  const panel = document.getElementById("impls");
  panel.replaceChildren();
  for (const impl of impls) {
    const row = document.createElement("div");
    row.className = "impl";
    row.dataset.impl = impl;
    row.innerHTML = `
      <div class="head">
        <span class="swatch"></span><span class="name"></span><span class="state"></span>
      </div>
      <div class="btns">
        <button data-act="run">run</button>
        <button data-act="clear">clear</button>
        <button data-act="toggle"></button>
      </div>`;
    row.querySelector(".swatch").style.background = color(impl);
    row.querySelector(".name").textContent = impl;
    row.querySelector('[data-act="run"]').onclick = () => post("run", impl);
    row.querySelector('[data-act="clear"]').onclick = () => post("clear", impl);
    row.querySelector('[data-act="toggle"]').onclick = () => {
      hidden[impl] = !hidden[impl];
      updatePanel(); draw();
    };
    panel.append(row);
  }
  const more = document.createElement("button");
  more.id = "more";
  more.textContent = "+";
  more.setAttribute("aria-label", "show the next implementation");
  more.onclick = () => { revealed++; updatePanel(); };
  panel.append(more);
  updatePanel();
}

function updatePanel() {
  // anything with data or a run in progress stays revealed (e.g. after a reload)
  impls.forEach((impl, i) => {
    if ((data[impl] || []).length || st[impl].state !== "idle" || st[impl].error)
      revealed = Math.max(revealed, i + 1);
  });
  document.querySelectorAll(".impl").forEach((row, i) => row.hidden = i >= revealed);
  const more = document.getElementById("more");
  if (more) more.hidden = revealed >= impls.length;
  for (const row of document.querySelectorAll(".impl")) {
    const impl = row.dataset.impl, s = st[impl], npts = (data[impl] || []).length;
    let text = "", err = false;
    if (s.state === "queued") text = "queued";
    else if (s.state === "running") text = `running n=${fmtN(s.n)}…`;
    else if (s.error) { text = `error at n=${fmtN(s.error)}`; err = true; }
    else if (s.skipped) text = `cut off after n=${fmtN(s.skipped)}`;
    else if (npts) text = `${npts} points`;
    const state = row.querySelector(".state");
    state.textContent = text;
    state.classList.toggle("err", err);
    row.classList.toggle("hidden", !!hidden[impl]);
    row.querySelector('[data-act="run"]').disabled = s.state === "queued" || s.state === "running";
    row.querySelector('[data-act="clear"]').disabled = !npts && s.state === "idle" && !s.error;
    row.querySelector('[data-act="toggle"]').textContent = hidden[impl] ? "show" : "hide";
  }
  const running = impls.find(impl => st[impl].state === "running");
  const queued = impls.filter(impl => st[impl].state === "queued");
  if (running)
    status(`running ${running}, n = ${st[running].n.toLocaleString()} …` +
           (queued.length ? `  (queued: ${queued.join(", ")})` : ""));
  else if (configured)
    status("idle — click run on an impl");
}

// ---- interaction -----------------------------------------------------------

const MIN_DRAG = 6;   // px; a drag thinner than this in one direction zooms only the other axis

// the selected rectangle; spans the full plot in a direction the drag barely moved in
function selection(d) {
  const w = Math.abs(d.x1 - d.x0), h = Math.abs(d.y1 - d.y0);
  return {
    x0: w < MIN_DRAG ? M.l : Math.min(d.x0, d.x1), x1: w < MIN_DRAG ? M.l + PW : Math.max(d.x0, d.x1),
    y0: h < MIN_DRAG ? M.t : Math.min(d.y0, d.y1), y1: h < MIN_DRAG ? M.t + PH : Math.max(d.y0, d.y1),
    tiny: w < MIN_DRAG && h < MIN_DRAG,
  };
}

const svgEl = document.getElementById("plot");
function svgPoint(ev) {
  const p = svgEl.createSVGPoint();
  p.x = ev.clientX; p.y = ev.clientY;
  const q = p.matrixTransform(svgEl.getScreenCTM().inverse());
  return {x: Math.min(Math.max(q.x, M.l), M.l + PW), y: Math.min(Math.max(q.y, M.t), M.t + PH), raw: q};
}

svgEl.addEventListener("pointerdown", ev => {
  if (ev.button !== 0) return;
  const p = svgPoint(ev);
  if (p.raw.x < M.l || p.raw.x > M.l + PW || p.raw.y < M.t || p.raw.y > M.t + PH) return;
  drag = {x0: p.x, y0: p.y, x1: p.x, y1: p.y};
  svgEl.setPointerCapture(ev.pointerId);
});
svgEl.addEventListener("pointermove", ev => {
  if (!drag) return;
  const p = svgPoint(ev);
  drag.x1 = p.x; drag.y1 = p.y;
  draw();
});
svgEl.addEventListener("pointerup", ev => {
  if (!drag) return;
  const r = selection(drag);
  drag = null;
  if (!r.tiny) setZoom({x: [sx.inv(r.x0), sx.inv(r.x1)], y: [sy.inv(r.y1), sy.inv(r.y0)]});
  else draw();
});
svgEl.addEventListener("pointercancel", () => { drag = null; draw(); });
svgEl.addEventListener("dblclick", () => setZoom(null));
document.addEventListener("keydown", ev => {
  if (ev.key !== "Escape") return;
  if (drag) { drag = null; draw(); } else setZoom(null);
});

function setZoom(z) {
  zoom = z;
  document.getElementById("reset").hidden = !zoom;
  draw();
}
document.getElementById("reset").addEventListener("click", () => setZoom(null));

function setLinear(on) {
  linear = on;
  for (const b of document.querySelectorAll(".seg button"))
    b.setAttribute("aria-pressed", String((b.dataset.scale === "linear") === linear));
  // a zoomed range reaching 0 or below can't be shown on log axes
  if (zoom && !linear && (zoom.x[0] <= 0 || zoom.y[0] <= 0)) setZoom(null);
  else draw();
}
for (const b of document.querySelectorAll(".seg button"))
  b.addEventListener("click", () => setLinear(b.dataset.scale === "linear"));

// ---- live data -------------------------------------------------------------

const fresh = () => ({state: "idle", n: null, skipped: null, error: null});
const status = s => document.getElementById("status").textContent = s;
const es = new EventSource("/events");
es.onerror = () => status("disconnected (script stopped?)");
es.onmessage = ev => {
  const m = JSON.parse(ev.data);
  const s = st[m.impl];
  if (m.type === "config") {
    impls = m.impls; colors = m.colors; ns = m.ns; data = {}; st = {};
    for (const impl of impls) st[impl] = fresh();
    document.getElementById("opts").textContent = m.opts;
    buildPanel();
    if (!configured) { configured = true; setLinear(m.linear); }
  } else if (m.type === "queued") {
    s.state = "queued";
  } else if (m.type === "clear") {
    data[m.impl] = [];
    // clearing a queued/running impl cancels it; a run starting also sends clear
    st[m.impl] = fresh();
  } else if (m.type === "status") {
    s.state = "running"; s.n = m.n;
  } else if (m.type === "point") {
    (data[m.impl] = data[m.impl] || []).push(m);
  } else if (m.type === "skip") {
    s.skipped = m.n;
  } else if (m.type === "error") {
    s.error = m.n;
    s.state = "idle";
  } else if (m.type === "finished") {
    s.state = "idle";
  }
  updatePanel();
  draw();
};
</script>
</body>
</html>
"""

if __name__ == "__main__":
    main()
