"""Draw the article's matplotlib figures from results/ (run run.sh first).

Writes ../red-drop-curve.svg, ../codel-trace.svg, ../aqm-mix.svg, ../aqm-sweep.svg.
"""
from pathlib import Path

import matplotlib

matplotlib.use("Agg")
import matplotlib.pyplot as plt  # noqa: E402

from summarize import load, medians  # noqa: E402

HERE = Path(__file__).resolve().parent
OUT = HERE.parent
RES = HERE / "results"

BG, FG, HI, DIM = "#2d333b", "#adbac7", "#cdd9e5", "#768390"
COL = {"taildrop": "#f85149", "red": "#f0883e", "codel": "#388bfd", "fqcodel": "#3fb950"}
NAME = {"taildrop": "tail drop", "red": "RED", "codel": "CoDel", "fqcodel": "FQ-CoDel"}
MARK = {"taildrop": "s", "red": "^", "codel": "o", "fqcodel": "D"}
LINE = {"taildrop": ":", "red": "--", "codel": "-", "fqcodel": "-."}

plt.rcParams.update({
    "svg.hashsalt": "aqm-71", "font.family": "monospace", "font.size": 10,
    "figure.facecolor": BG, "axes.facecolor": BG, "savefig.facecolor": BG,
    "axes.edgecolor": DIM, "axes.labelcolor": FG, "text.color": FG,
    "xtick.color": FG, "ytick.color": FG, "grid.color": "#444c56",
    "legend.facecolor": BG, "legend.edgecolor": DIM,
})


def save(fig, name):
    fig.savefig(OUT / name, format="svg", metadata={"Date": None}, bbox_inches="tight")
    plt.close(fig)


def red_curve():
    mn, mx, maxp = 20, 60, 0.02
    fig, (a, b) = plt.subplots(1, 2, figsize=(10, 3.8))
    xs = [0, mn, mx, mx, 80]
    ys = [0, 0, maxp, 1.0, 1.0]
    a.plot(xs[:3], ys[:3], color=COL["red"], lw=2, label="p_b (linear)")
    a.plot([mx, mx, 80], [maxp, 1.0, 1.0], color=COL["taildrop"], lw=2, ls="--",
           label="avg >= max_th: drop all")
    a.set_yscale("symlog", linthresh=0.01)
    a.set_xlabel("average queue avg (packets)")
    a.set_ylabel("drop probability")
    a.set_title("RED: p_b vs avg (min_th=20, max_th=60, max_p=0.02)", color=HI, fontsize=10)
    a.grid(True, ls=":")
    a.legend(loc="upper left", fontsize=8)
    pb = 0.02
    n = list(range(1, 121))
    geo = [pb * (1 - pb) ** (k - 1) for k in n]
    uni = [1 / 50 if k <= 50 else 0 for k in n]
    b.step(n, geo, where="mid", color=DIM, lw=1.5, label="fixed p_b: geometric")
    b.step(n, uni, where="mid", color=COL["red"], lw=2, label="p_a: uniform")
    b.set_xlabel("packets between drops")
    b.set_ylabel("probability")
    b.set_title("gap distribution at p_b = 1/50", color=HI, fontsize=10)
    b.grid(True, ls=":")
    b.legend(loc="upper right", fontsize=8)
    save(fig, "red-drop-curve.svg")


def read_trace(q):
    t, soj, drops = [], [], []
    for line in (RES / f"trace_{q}.txt").read_text().splitlines():
        f = line.split()
        if f[0] == "S":
            t.append(float(f[1]))
            soj.append(float(f[2]))
        else:
            drops.append(float(f[1]))
    return t, soj, drops


def codel_trace():
    fig, axes = plt.subplots(2, 1, figsize=(10, 5.6), sharex=True)
    for ax, q in zip(axes, ["taildrop", "codel"]):
        t, soj, drops = read_trace(q)
        ax.plot(t, soj, color=COL[q], lw=1, label=f"{NAME[q]}: sojourn time")
        top = max(soj) * 1.08
        ax.plot(drops, [top] * len(drops), "|", color=HI, ms=8, label="drop")
        ax.set_ylabel("sojourn (ms)")
        ax.grid(True, ls=":")
        ax.set_xlim(0, 20)
        if q == "codel":
            ax.axhline(5, color=DIM, ls="--", lw=1, label="target 5 ms")
        ax.legend(loc="upper right", fontsize=8)
    axes[0].set_title("one Reno flow, 20 Mbit/s, RTT 50 ms, limit 1000 packets", color=HI,
                      fontsize=10)
    axes[1].set_xlabel("time (s)")
    save(fig, "codel-trace.svg")


def mix():
    tab = medians(load("e1_mix.txt"))
    qs = ["taildrop", "red", "codel", "fqcodel"]
    rows = {q: tab[(q, "20", "4")] for q in qs}
    fig, (a, b) = plt.subplots(1, 2, figsize=(10, 3.8))
    w = 0.38
    for i, q in enumerate(qs):
        r = rows[q]
        a.bar(i - w / 2, r["probe_p50"], w, color=COL[q])
        a.bar(i + w / 2, r["probe_p99"], w, color=COL[q], alpha=0.45, hatch="//", edgecolor=HI)
        b.bar(i - w / 2, r["fct_p50"], w, color=COL[q])
        b.bar(i + w / 2, r["fct_p95"], w, color=COL[q], alpha=0.45, hatch="//", edgecolor=HI)
        a.text(i + w / 2, r["probe_p99"] * 1.15, f"{r['probe_p99']:.3g}", ha="center",
               fontsize=8, color=HI)
        b.text(i + w / 2, r["fct_p95"] * 1.04, f"{r['fct_p95']:.0f}", ha="center",
               fontsize=8, color=HI)
    for ax in (a, b):
        ax.set_xticks(range(4), [NAME[q] for q in qs])
        ax.grid(True, axis="y", ls=":")
    a.set_yscale("log")
    a.set_ylabel("probe queueing delay (ms)")
    a.set_title("sparse probe: p50 (solid) / p99 (hatched)", color=HI, fontsize=10)
    b.set_ylabel("completion time (ms)")
    b.set_title("40-packet transfers: p50 (solid) / p95 (hatched)", color=HI, fontsize=10)
    save(fig, "aqm-mix.svg")


def sweep():
    rate = medians(load("e2_rate.txt"))
    flows = medians(load("e2_flows.txt"))
    fig, ax = plt.subplots(2, 2, figsize=(10, 6.4))
    rates = [5, 10, 20, 50, 100]
    ns = [1, 2, 4, 8, 16, 32]
    for q in ["red", "codel", "fqcodel"]:
        kw = dict(color=COL[q], marker=MARK[q], ls=LINE[q], label=NAME[q])
        ax[0][0].plot(rates, [rate[(q, str(r), "4")]["q_p50"] for r in rates], **kw)
        ax[1][0].plot(rates, [rate[(q, str(r), "4")]["util"] for r in rates], **kw)
        ax[0][1].plot(ns, [flows[(q, "20", str(n))]["q_p50"] for n in ns], **kw)
        ax[1][1].plot(ns, [flows[(q, "20", str(n))]["util"] for n in ns], **kw)
    for a in ax[0][0], ax[1][0]:
        a.set_xscale("log")
        a.set_xticks(rates, [str(r) for r in rates])
    for a in ax[0][1], ax[1][1]:
        a.set_xscale("log", base=2)
        a.set_xticks(ns, [str(n) for n in ns])
    ax[0][0].set_yscale("log")
    ax[0][0].set_title("4 flows, varying link rate", color=HI, fontsize=10)
    ax[0][1].set_title("20 Mbit/s, varying number of flows", color=HI, fontsize=10)
    ax[0][0].set_ylabel("median sojourn (ms)")
    ax[0][1].set_ylabel("median sojourn (ms)")
    ax[1][0].set_ylabel("link utilization")
    ax[1][1].set_ylabel("link utilization")
    ax[1][0].set_xlabel("link rate (Mbit/s)")
    ax[1][1].set_xlabel("long-lived flows")
    for row in ax:
        for a in row:
            a.grid(True, ls=":")
            a.axhline(5, color=DIM, ls=":", lw=1) if a in (ax[0][0], ax[0][1]) else None
    ax[1][0].set_ylim(0.8, 1.01)
    ax[1][1].set_ylim(0.8, 1.01)
    ax[0][0].text(5.2, 5.4, "5 ms", color=DIM, fontsize=8)
    ax[0][1].text(1.05, 5.8, "5 ms", color=DIM, fontsize=8)
    ax[0][1].legend(loc="upper left", fontsize=8)
    save(fig, "aqm-sweep.svg")


if __name__ == "__main__":
    red_curve()
    codel_trace()
    mix()
    sweep()
