"""Draw the four data figures from li_bench output.

Usage: python3 plot.py results      (reads results/{segs,bench,dump}.txt, writes ../*.svg)
"""
import collections
import os
import sys

import matplotlib

matplotlib.use("Agg")
import matplotlib.pyplot as plt

BG, FG, HI, SUB, GRID = "#2d333b", "#adbac7", "#cdd9e5", "#768390", "#444c56"
COL = {"uniform": "#388bfd", "normal": "#3fb950", "lognormal": "#f0883e", "clustered": "#a371f7"}
DS = ["uniform", "normal", "lognormal", "clustered"]
plt.rcParams.update({
    "svg.hashsalt": "learned-index", "font.family": "monospace", "font.size": 10,
    "figure.facecolor": BG, "axes.facecolor": BG, "axes.edgecolor": SUB, "axes.labelcolor": FG,
    "xtick.color": FG, "ytick.color": FG, "text.color": FG, "grid.color": GRID,
    "legend.facecolor": BG, "legend.edgecolor": SUB, "legend.labelcolor": FG,
})
res = sys.argv[1] if len(sys.argv) > 1 else "results"
out = os.path.join(os.path.dirname(os.path.abspath(__file__)), "..")


def rows(name, tag):
    for line in open(os.path.join(res, name)):
        if line.startswith(tag + ","):
            yield line.strip().split(",")[1:]


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


# ---- figure 1: CDF of every dataset
fig, axes = plt.subplots(1, 4, figsize=(13, 3.4))
cdf = collections.defaultdict(lambda: ([], []))
for ds, k, i in rows("dump.txt", "C"):
    cdf[ds][0].append(int(k))
    cdf[ds][1].append(int(i))
for ax, ds in zip(axes, DS):
    n = cdf[ds][1][-1] + 1
    ax.plot(cdf[ds][0], [i / n for i in cdf[ds][1]], color=COL[ds], lw=1.6)
    if ds == "lognormal":
        ax.set_xscale("log")
        ax.set_xlabel("key (log scale)")
    else:
        ax.set_xlabel("key")
    ax.set_title(ds, color=HI)
    ax.grid(True, lw=0.5)
axes[0].set_ylabel("rank / n")
fig.tight_layout()
save(fig, "dataset-cdfs.svg")

# ---- figure 2: a zoomed slice with its eps-segments
fig, a2 = plt.subplots(figsize=(8, 4.6))
pts = [(int(k), int(i)) for k, i in rows("dump.txt", "Z")]
k0 = pts[0][0]
xs = [k - k0 for k, _ in pts]
y0 = pts[0][1]
ys = [i - y0 for _, i in pts]
a2.step(xs, ys, where="post", color=FG, lw=1.0, label="rank(key)")
first = True
for key, end, slope, icpt, eps in rows("dump.txt", "G"):
    key, end, slope, icpt, eps = int(key), int(end), float(slope), int(icpt), int(eps)
    lo, hi = max(key, pts[0][0]), min(end, pts[-1][0])
    if lo >= hi:
        continue
    gx = [lo - k0, hi - k0]
    gy = [icpt - y0 + slope * (lo - key), icpt - y0 + slope * (hi - key)]
    a2.fill_between(gx, [g - eps for g in gy], [g + eps for g in gy], color="#388bfd", alpha=0.18,
                    lw=0, label="+/- eps band" if first else None)
    a2.plot(gx, gy, color="#f0883e", lw=1.8, label="segment" if first else None)
    a2.axvline(gx[0], color=SUB, lw=0.6, ls=":")
    first = False
a2.set_xlabel("key - %d" % k0)
a2.set_ylabel("position - %d" % y0)
a2.set_title("clustered (n = 200k), positions %d-%d, eps = %d, dotted = segment start" % (y0, y0 + ys[-1], eps), color=HI, fontsize=10)
a2.grid(True, lw=0.5)
a2.legend(loc="upper left")
save(fig, "pla-error-band.svg")

# ---- figure 3: number of segments vs eps, optimal PLA vs ShrinkingCone
fig, ax = plt.subplots(figsize=(7.5, 4.6))
seg = collections.defaultdict(list)
n_of = {}
for ds, n, eps, opt, cone, *_ in rows("segs.txt", "S"):
    seg[ds].append((int(eps), int(opt), int(cone)))
    n_of[ds] = int(n)
for ds in DS:
    e = [r[0] for r in seg[ds]]
    ax.plot(e, [r[1] for r in seg[ds]], color=COL[ds], marker="o", ms=4, lw=1.6, label=ds + " optimal")
    ax.plot(e, [r[2] for r in seg[ds]], color=COL[ds], marker="x", ms=5, lw=1.2, ls="--", label=ds + " cone")
e = [r[0] for r in seg["uniform"]]
ax.plot(e, [n_of["uniform"] / (2 * x) for x in e], color=SUB, lw=1.0, ls=":", label="n / (2 eps) bound")
ax.set_xscale("log", base=2)
ax.set_yscale("log")
ax.set_xlabel("eps (max position error)")
ax.set_ylabel("number of segments (n = 10M)")
ax.set_title("Optimal PLA (solid) vs ShrinkingCone (dashed)", color=HI)
ax.grid(True, which="major", lw=0.5)
ax.legend(fontsize=8, ncol=2, loc="lower left")
save(fig, "segments-vs-eps.svg")

# ---- figure 4: comparisons and cache lines per lookup vs index size
bench = collections.defaultdict(list)
for ds, st, param, b, cmps, lines, ns, extra in rows("bench.txt", "R"):
    bench[(ds, st)].append((int(param), int(b), float(cmps), float(lines)))
STY = {"btree": ("#3fb950", "s", "static B+tree (B = 16..1024)"),
       "pgm": ("#388bfd", "o", "PGM (eps = 8..2048)"),
       "rmi": ("#f0883e", "^", "2-stage RMI (2^10..2^20 leaves)")}
fig, axes = plt.subplots(2, 4, figsize=(13, 6.4), sharex=True)
for c, ds in enumerate(DS):
    for r, (idx, lab) in enumerate([(2, "key comparisons / lookup"), (3, "distinct cache lines / lookup")]):
        ax = axes[r][c]
        bs = bench[(ds, "bs")][0][idx]
        ax.axhline(bs, color=SUB, lw=1.0, ls="--", label="binary search (no index)")
        for st, (col, mk, name) in STY.items():
            d = bench[(ds, st)]
            ax.plot([max(x[1], 16) for x in d], [x[idx] for x in d], color=col, marker=mk, ms=5, lw=1.4, label=name)
        ax.set_xscale("log")
        ax.set_ylim(0, 26)
        ax.grid(True, lw=0.5)
        if r == 0:
            ax.set_title(ds, color=HI)
        else:
            ax.set_xlabel("index size (bytes)")
        if c == 0:
            ax.set_ylabel(lab)
h, l = axes[0][0].get_legend_handles_labels()
fig.legend(h, l, loc="lower center", ncol=4, bbox_to_anchor=(0.5, -0.04))
fig.tight_layout(rect=(0, 0.04, 1, 1))
save(fig, "lookup-cost-vs-size.svg")
print("wrote dataset-cdfs.svg pla-error-band.svg segments-vs-eps.svg lookup-cost-vs-size.svg")
