"""Summarise results/zipf.tsv (median of 3 seeds) and draw the two figures."""
import csv
import statistics
import sys
from collections import defaultdict
from pathlib import Path

import matplotlib

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

here = Path(__file__).resolve().parent
tsv = here / "results" / "zipf.tsv"
out_dir = here.parent

plt.rcParams.update({
    "svg.hashsalt": "frequency-estimation",
    "font.family": "monospace",
    "font.size": 10,
    "figure.facecolor": "#2d333b",
    "axes.facecolor": "#2d333b",
    "axes.edgecolor": "#768390",
    "axes.labelcolor": "#adbac7",
    "text.color": "#adbac7",
    "xtick.color": "#adbac7",
    "ytick.color": "#adbac7",
    "grid.color": "#444c56",
    "legend.facecolor": "#2d333b",
    "legend.edgecolor": "#768390",
})
COLORS = {"MG": "#f0883e", "SS": "#388bfd", "LC": "#3fb950", "DS": "#a371f7"}
NAMES = {"MG": "Misra-Gries", "SS": "Space-Saving", "LC": "Lossy Counting", "DS": "DataSketches 5.2.0"}

rows = list(csv.DictReader(open(tsv), delimiter="\t"))
n = int(rows[0]["n"])
agg = defaultdict(list)
for r in rows:
    agg[(r["alpha"], r["algo"], int(r["m"]))].append(r)

def med(key, field, cast=float):
    return statistics.median(cast(r[field]) for r in agg[key])

alphas = sorted({r["alpha"] for r in rows}, key=float)
ms = sorted({int(r["m"]) for r in rows})
algos = ["MG", "SS", "LC", "DS"]

lines = [f"n = {n}, universe 2^20, median of 3 seeds", ""]
lines.append("alpha  algo     m  recall100  err_top100  guaranteed  entries")
for a in alphas:
    for algo in algos:
        for m in ms:
            k = (a, algo, m)
            lines.append(f"{a:>5}  {algo:>4}  {m:>4}  {med(k,'recall100'):9.2f}  {med(k,'err_top100',int):10d}"
                         f"  {med(k,'guaranteed',int):10d}  {med(k,'entries',int):7d}")
    lines.append("")
same = sum(1 for r in rows if r["algo"] == "DS")
eq = sum(1 for (a, algo, m), v in agg.items() if algo == "DS"
         for r, q in zip(v, agg[(a, "MG", m)]) if r["guaranteed"] == q["guaranteed"])
lines.append(f"DS offset == MG decrement count D in {eq} of {same} runs")
(here / "results" / "summary.txt").write_text("\n".join(lines) + "\n")
print("\n".join(lines))

meta = {"Date": None, "Creator": None}

fig, axes = plt.subplots(1, len(alphas), figsize=(12, 3.6), sharey=True)
for ax, a in zip(axes, alphas):
    for algo in algos:
        ax.plot(ms, [med((a, algo, m), "recall100") for m in ms], marker="o", ms=4,
                color=COLORS[algo], label=NAMES[algo], alpha=0.9)
    ax.set_xscale("log", base=2)
    ax.set_xticks(ms, [str(m) for m in ms], fontsize=8)
    ax.set_title(f"Zipf alpha = {a}")
    ax.set_xlabel("counters m")
    ax.grid(True, alpha=0.4)
axes[0].set_ylabel("top-100 recall")
axes[0].set_ylim(0, 1.05)
axes[-1].legend(loc="lower right", fontsize=8)
fig.tight_layout()
fig.savefig(out_dir / "topk-recall.svg", metadata=meta)

fig, axes = plt.subplots(1, len(alphas), figsize=(12, 3.8), sharey=True)
for ax, a in zip(axes, alphas):
    ax.plot(ms, [n / (m + 1) for m in ms], color="#768390", ls="--", label="n/(m+1)")
    for algo, label in [("MG", "Misra-Gries (= its D)"), ("SS", "Space-Saving"), ("LC", "Lossy Counting")]:
        ys = [med((a, algo, m), "err_top100", int) + 1 for m in ms]
        ax.plot(ms, ys, marker="o", ms=4, color=COLORS[algo], label=label)
    ax.set_xscale("log", base=2)
    ax.set_yscale("log")
    ax.set_xticks(ms, [str(m) for m in ms], fontsize=8)
    ax.set_title(f"Zipf alpha = {a}")
    ax.set_xlabel("counters m")
    ax.grid(True, alpha=0.4)
axes[0].set_ylabel("max error on true top-100, + 1")
axes[0].legend(loc="lower left", fontsize=7)
fig.tight_layout()
fig.savefig(out_dir / "topk-error.svg", metadata=meta)
sys.exit(0)
