#!/usr/bin/env python3
"""Draw the article's three data figures from results/*.tsv (needs matplotlib)."""
import csv
import os
import sys

import matplotlib

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

HERE = os.path.dirname(os.path.abspath(__file__))
RES = os.path.join(HERE, "results")
OUT = os.path.dirname(HERE)

BG, FG, HI, DIM, GRID = "#2d333b", "#adbac7", "#cdd9e5", "#768390", "#444c56"
BLUE, GREEN, ORANGE, PURPLE, RED, GRAY = "#388bfd", "#3fb950", "#f0883e", "#a371f7", "#f85149", "#636e7b"

plt.rcParams.update({
    "svg.hashsalt": "84-integer-compression", "svg.fonttype": "none",
    "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": GRID, "legend.facecolor": BG,
    "legend.edgecolor": DIM, "legend.labelcolor": FG,
})


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


def load_sizes():
    t = {}
    with open(os.path.join(RES, "sizes.tsv")) as f:
        for r in csv.DictReader(f, delimiter="\t"):
            t[(r["dataset"], r["subset"], r["codec"])] = float(r["bits_per_int"])
    return t


def density_figure(t):
    ks = list(range(1, 16))
    panels = [
        ("byte- and bit-aligned codes", [
            ("varint", "varint (LEB128)", ORANGE, "-", "o"),
            ("stream-vbyte", "Stream VByte", RED, "--", "s"),
            ("gamma", "Elias gamma", GRAY, ":", "^"),
            ("delta", "Elias delta", DIM, "-.", "v"),
            ("golomb", "Golomb (b from n/N)", GREEN, "-", "D"),
            ("interpolative", "interpolative", PURPLE, "--", "x"),
            ("elias-fano", "Elias-Fano", BLUE, "-", "+"),
        ]),
        ("word- and block-aligned codes", [
            ("fp:simple16", "Simple-16", GRAY, ":", "^"),
            ("fp:simple8b", "Simple-8b", DIM, "-.", "v"),
            ("bp128", "BP128", ORANGE, "-", "o"),
            ("lucene-10.5-docs", "Lucene 10.5 doc blocks", RED, "--", "s"),
            ("pfor128", "PFOR-128 (ours)", BLUE, "-", "+"),
            ("fp:simdfastpfor128", "SIMD-FastPFor", GREEN, "--", "D"),
            ("fp:optpfor", "OptPFD", PURPLE, ":", "x"),
        ]),
    ]
    fig, axes = plt.subplots(1, 2, figsize=(11, 5.2), sharey=True)
    for ax, (title, series) in zip(axes, panels):
        for key, label, color, ls, mk in series:
            ys = [t[("bern%d" % k, "all", key)] - t[("bern%d" % k, "all", "log2-binom")] for k in ks]
            ax.plot(ks, ys, color=color, linestyle=ls, marker=mk, markersize=4, linewidth=1.4, label=label)
        ax.axhline(0, color=HI, linewidth=0.8)
        ax.set_title(title, color=HI)
        ax.set_xlabel("log2(1/p): Bernoulli density p = 2^-k")
        ax.set_xticks(ks)
        ax.set_ylim(-0.2, 12.2)
        ax.grid(True, linewidth=0.5)
        ax.legend(loc="upper center", bbox_to_anchor=(0.5, -0.16), fontsize=8, ncol=3)
    axes[0].set_ylabel("bits per posting above log2 C(N,n) / n")
    save(fig, "bits-vs-density.svg")


def real_figure(t):
    names = [
        ("interpolative", "interpolative"), ("golomb", "Golomb"), ("rice", "Rice"),
        ("gamma", "Elias gamma"), ("fp:simple16", "Simple-16"), ("fp:optpfor", "OptPFD"),
        ("elias-fano", "Elias-Fano"), ("pfor128", "PFOR-128 (ours)"), ("fp:simple8b", "Simple-8b"),
        ("fp:simdfastpfor128", "SIMD-FastPFor"), ("lucene-10.5-docs", "Lucene 10.5 doc blocks"),
        ("bp128", "BP128"), ("varint", "varint (LEB128)"), ("stream-vbyte", "Stream VByte"),
    ]
    fig, axes = plt.subplots(1, 2, figsize=(11, 5), sharey=True)
    for ax, ds in zip(axes, ["bible", "world192"]):
        vals = [t[(ds, "n>=256", k)] for k, _ in names]
        ys = range(len(names))
        ax.barh(list(ys), vals, color=BLUE, height=0.6)
        bound = t[(ds, "n>=256", "log2-binom")]
        ax.axvline(bound, color=ORANGE, linestyle="--", linewidth=1.2)
        for y, v in zip(ys, vals):
            ax.text(v + 0.1, y, "%.2f" % v, va="center", fontsize=8, color=HI)
        ax.set_title("%s.txt, n >= 256 (dashed: log2 C(N,n) / n = %.2f)" % (ds, bound), color=HI, fontsize=10)
        ax.set_xlabel("bits per posting")
        ax.set_xlim(0, 12.5)
        ax.grid(True, axis="x", linewidth=0.5)
    axes[0].set_yticks(list(range(len(names))))
    axes[0].set_yticklabels([n for _, n in names])
    axes[0].invert_yaxis()
    save(fig, "bits-real-lists.svg")


def median(xs):
    xs = sorted(xs)
    return xs[len(xs) // 2]


def layout_figure():
    per = {}
    with open(os.path.join(RES, "layout_time.tsv")) as f:
        for r in csv.DictReader(f, delimiter="\t"):
            per.setdefault(int(r["b"]), []).append(
                (float(r["horizontal_ns_per_int"]), float(r["vertical_sse2_ns_per_int"])))
    b = sorted(per)
    h = [median([x[0] for x in per[k]]) for k in b]
    v = [median([x[1] for x in per[k]]) for k in b]
    fig, ax = plt.subplots(figsize=(7.5, 3.8))
    ax.plot(b, h, color=ORANGE, marker="o", markersize=3, label="horizontal, scalar loop")
    ax.plot(b, v, color=BLUE, marker="s", markersize=3, linestyle="--", label="vertical 4-lane, SSE2")
    ax.set_xlabel("bit width b")
    ax.set_ylabel("ns per integer (median of runs)")
    ax.set_ylim(bottom=0)
    ax.grid(True, linewidth=0.5)
    ax.legend(loc="upper left", fontsize=8)
    save(fig, "unpack-layout-time.svg")


def decode_summary():
    per = {}
    with open(os.path.join(RES, "decode_time.tsv")) as f:
        for r in csv.DictReader(f, delimiter="\t"):
            per.setdefault((r["codec"], r["dataset"]), []).append(float(r["median_ns"]))
    datasets = ["bible", "bern4", "bern12"]
    codecs = sorted({c for c, _ in per}, key=lambda c: median(per[(c, "bible")]))
    lines = ["codec\t" + "\t".join("%s median [min, max of runs]" % d for d in datasets)]
    for c in codecs:
        cells = []
        for d in datasets:
            xs = per[(c, d)]
            cells.append("%.3f [%.3f, %.3f]" % (median(xs), min(xs), max(xs)))
        lines.append(c + "\t" + "\t".join(cells))
    with open(os.path.join(RES, "decode_summary.txt"), "w") as f:
        f.write("\n".join(lines) + "\n")
    print("\n".join(lines))


def main():
    t = load_sizes()
    density_figure(t)
    real_figure(t)
    layout_figure()
    decode_summary()


if __name__ == "__main__":
    sys.exit(main())
