#!/usr/bin/env python3
"""Generate the hand-laid-out SVG diagrams of the cuckoo hashing article.

Usage: python3 draw_diagrams.py   (writes ../*.svg)
The states drawn in kickout-steps.svg and cuckoo-graph.svg are the ones
printed by `./cuckoo_sim e1`.
"""
import os

BG, FG, HI, DIM = "#2d333b", "#adbac7", "#cdd9e5", "#768390"
BLUE, GREEN, ORANGE, PURPLE, RED, GRAY = "#388bfd", "#3fb950", "#f0883e", "#a371f7", "#f85149", "#636e7b"
FONT = "ui-monospace, 'SFMono-Regular', 'DejaVu Sans Mono', Menlo, monospace"
OUT = os.path.join(os.path.dirname(os.path.abspath(__file__)), "..")


def svg(w, h, body, defs=""):
    return (f'<svg xmlns="http://www.w3.org/2000/svg" width="{w}" height="{h}" viewBox="0 0 {w} {h}" '
            f'font-family="{FONT}" font-size="13">\n<defs>{defs}</defs>\n'
            f'<rect width="{w}" height="{h}" fill="{BG}"/>\n{body}</svg>\n')


def marker(mid, color):
    return (f'<marker id="{mid}" viewBox="0 0 10 10" refX="9" refY="5" markerWidth="7" markerHeight="7" '
            f'orient="auto-start-reverse"><path d="M0,0 L10,5 L0,10 z" fill="{color}"/></marker>')


def text(x, y, s, color=FG, size=13, anchor="middle", weight="normal"):
    s = s.replace("&", "&amp;").replace("<", "&lt;").replace(">", "&gt;")
    return (f'<text x="{x}" y="{y}" fill="{color}" font-size="{size}" text-anchor="{anchor}" '
            f'font-weight="{weight}">{s}</text>\n')


def rect(x, y, w, h, stroke, fill=BG, sw=1.5, rx=4, dash=None):
    d = f' stroke-dasharray="{dash}"' if dash else ""
    return f'<rect x="{x}" y="{y}" width="{w}" height="{h}" rx="{rx}" fill="{fill}" stroke="{stroke}" stroke-width="{sw}"{d}/>\n'


# ---------------------------------------------------------------- figure 1
def kickout_steps():
    states = [
        ("before", "insert D", "C B _ _", "_ _ A _", "D (new)", None, ORANGE),
        ("step 1", "D -> T1[h1(D)=0]", "D B _ _", "_ _ A _", "C", ("T1", 0), ORANGE),
        ("step 2", "C -> T2[h2(C)=2]", "D B _ _", "_ _ C _", "A", ("T2", 2), ORANGE),
        ("step 3", "A -> T1[h1(A)=0]", "A B _ _", "_ _ C _", "D", ("T1", 0), ORANGE),
        ("step 4", "D -> T2[h2(D)=3]", "A B _ _", "_ _ C D", "none", ("T2", 3), GREEN),
    ]
    W, PW = 920, 180
    body = text(W / 2, 30, "Pagh-Rodler insertion of D: h1(D)=0, h2(D)=3", HI, 16, weight="bold")
    body += text(W / 2, 52, "hashes: A (0,2)   B (1,1)   C (0,2)   D (0,3)", DIM, 12)
    for i, (tag, what, t1, t2, nest, hl, hlc) in enumerate(states):
        x0 = 10 + i * PW
        body += rect(x0 + 4, 68, PW - 8, 318, GRAY, BG, 1, 6)
        body += text(x0 + PW / 2, 90, tag, HI, 13, weight="bold")
        body += text(x0 + PW / 2, 110, what, FG, 11)
        for col, (name, cells) in enumerate((("T1", t1.split()), ("T2", t2.split()))):
            cx = x0 + 30 + col * 70
            body += text(cx + 27, 138, name, BLUE if col == 0 else PURPLE, 12, weight="bold")
            for j, c in enumerate(cells):
                y = 148 + j * 42
                is_hl = hl == (name, j)
                stroke = hlc if is_hl else (GRAY if c == "_" else FG)
                sw = 2.5 if is_hl else 1.2
                body += rect(cx, y, 54, 34, stroke, BG, sw)
                body += text(cx + 9, y + 21, str(j), DIM, 10)
                body += text(cx + 32, y + 22, "" if c == "_" else c, HI if c != "_" else DIM, 15, weight="bold")
        ncolor = GREEN if nest == "none" else ORANGE
        body += text(x0 + PW / 2, 332, "nestless key", DIM, 11)
        body += rect(x0 + PW / 2 - 40, 340, 80, 30, ncolor, BG, 1.8)
        body += text(x0 + PW / 2, 360, nest, ncolor, 13, weight="bold")
    body += text(W / 2, 410, "orange cell = written in this step; D is evicted by A in step 3 and "
                 "ends in T2[3] after 3 evictions", DIM, 12)
    return svg(W, 425, body)


# ---------------------------------------------------------------- figure 2
def node(x, y, label, color):
    return rect(x - 38, y - 16, 76, 32, color, BG, 1.8, 6) + text(x, y + 5, label, HI, 13, weight="bold")


def cuckoo_graph():
    W, H = 920, 420
    defs = marker("ag", GREEN) + marker("ab", BLUE)
    body = ""
    for p in range(2):
        x0 = 10 + p * 455
        body += rect(x0, 10, 445, 400, GRAY, BG, 1, 6)
        title = "(a) after inserting A, B, C, D" if p == 0 else "(b) inserting E: h1(E)=0, h2(E)=2"
        body += text(x0 + 222, 36, title, HI, 14, weight="bold")
        L, R = x0 + 100, x0 + 330
        ya, yb, yc = 135, 115, 205
        comp_color = BLUE if p == 0 else RED
        # component 1: T1[0] - T2[2] (A, C, [E]) and T1[0] - T2[3] (D)
        mk = 'marker-start="url(#ab)"' if p == 0 else ""
        mk_end = 'marker-end="url(#ab)"' if p == 0 else ""
        # A: stored in T1[0] -> arrow at the T1 end
        body += f'<line x1="{L + 40}" y1="{ya}" x2="{R - 40}" y2="{yb}" stroke="{comp_color}" stroke-width="2" {mk}/>\n'
        body += text((L + R) / 2, (ya + yb) / 2 - 6, "A", HI, 13, weight="bold")
        # C: stored in T2[2] -> arrow at the T2 end, curve above
        body += (f'<path d="M{L + 40},{ya - 8} Q{(L + R) / 2},{yb - 50} {R - 40},{yb - 8}" fill="none" '
                 f'stroke="{comp_color}" stroke-width="2" {mk_end}/>\n')
        body += text((L + R) / 2, yb - 26, "C", HI, 13, weight="bold")
        # D: stored in T2[3]
        body += f'<line x1="{L + 40}" y1="{ya + 8}" x2="{R - 40}" y2="{yc}" stroke="{comp_color}" stroke-width="2" {mk_end}/>\n'
        body += text((L + R) / 2 + 10, (ya + yc) / 2 + 22, "D", HI, 13, weight="bold")
        if p == 1:
            body += (f'<path d="M{L + 40},{ya + 2} Q{(L + R) / 2},{ya + 20} {R - 40},{yb + 10}" fill="none" '
                     f'stroke="{RED}" stroke-width="2.5" stroke-dasharray="6,4"/>\n')
            body += text((L + R) / 2 - 10, ya + 27, "E", RED, 13, weight="bold")
        body += node(L, ya, "T1[0]", comp_color)
        body += node(R, yb, "T2[2]", comp_color)
        body += node(R, yc, "T2[3]", comp_color)
        if p == 0:
            body += text(x0 + 222, 250, "unicyclic: 3 cells, 3 keys -> placeable", BLUE, 12)
        else:
            body += text(x0 + 222, 250, "3 cells, 4 keys -> no placement, rehash", RED, 12, weight="bold")
        # component 2: T1[1] - T2[1] (B)
        yt = 298
        mkb = 'marker-start="url(#ag)"'
        body += f'<line x1="{L + 40}" y1="{yt}" x2="{R - 40}" y2="{yt}" stroke="{GREEN}" stroke-width="2" {mkb}/>\n'
        body += text((L + R) / 2, yt - 8, "B", HI, 13, weight="bold")
        body += node(L, yt, "T1[1]", GREEN)
        body += node(R, yt, "T2[1]", GREEN)
        body += text(x0 + 222, 336, "tree: 2 cells, 1 key", GREEN, 12)
        body += text(x0 + 222, 362, "untouched cells: T1[2] T1[3] T2[0]", DIM, 11)
        if p == 0:
            body += text(x0 + 222, 390, "arrowhead = cell currently holding the key", DIM, 11)
        else:
            body += text(x0 + 222, 390, "keys (edges) > cells (vertices) in one component", DIM, 11)
    return svg(W, H, body, defs)


# ---------------------------------------------------------------- figure 3
def bfs_backward():
    W, H = 920, 500
    defs = marker("ao", ORANGE) + marker("abl", BLUE)
    body = text(W / 2, 28, "(2,4) cuckoo insert of x: find the path first, then move the hole backwards",
                HI, 15, weight="bold")
    # left panel: BFS discovery
    body += rect(10, 44, 330, 446, GRAY, BG, 1, 6)
    body += text(175, 68, "1. path discovery (reads only)", BLUE, 13, weight="bold")
    levels = [(110, [("b1", True), ("b2", False)]),
              (200, [("b4", True), ("b6", False), ("b3", False), ("b5", False)]),
              (290, [("b7", True), ("b8", False), ("b9", False)])]
    pos = {}
    for y, nodes in levels:
        for i, (n, on) in enumerate(nodes):
            x = 80 + i * 66
            pos[n] = (x, y)
            col = BLUE if on else GRAY
            body += rect(x - 24, y - 15, 48, 30, col, BG, 1.8 if on else 1, 5)
            body += text(x, y + 5, n, HI if on else DIM, 12, weight="bold")
    for a, b, lab in [("b1", "b4", "a"), ("b1", "b6", ""), ("b2", "b3", ""), ("b2", "b5", ""),
                      ("b4", "b7", "c"), ("b4", "b8", ""), ("b6", "b9", "")]:
        (x1, y1), (x2, y2) = pos[a], pos[b]
        col = BLUE if lab else GRAY
        body += (f'<line x1="{x1}" y1="{y1 + 15}" x2="{x2}" y2="{y2 - 16}" stroke="{col}" '
                 f'stroke-width="{2 if lab else 1}" marker-end="url(#abl)"/>\n')
        if lab:
            body += text((x1 + x2) / 2 - 12, (y1 + y2) / 2 + 4, lab, BLUE, 13, weight="bold")
    body += text(18, 342, "b1, b2: candidate buckets of x, full", FG, 11, "start")
    body += text(18, 362, "edge label: key that could move there", FG, 11, "start")
    body += text(18, 382, "b7 has a free slot", FG, 11, "start")
    body += text(18, 402, "path: x->b1, a->b4, c->b7", FG, 11, "start")
    body += text(18, 436, "BFS returns a shortest path, so no", DIM, 11, "start")
    body += text(18, 454, "bucket appears twice on it", DIM, 11, "start")
    # right panel: execution
    body += rect(350, 44, 560, 446, GRAY, BG, 1, 6)
    body += text(630, 68, "2. execution, starting at the free end", ORANGE, 13, weight="bold")
    b1 = ["a", "e", "f", "g"]
    b4 = ["c", "h", "i", "j"]
    b7 = ["k", "l", "m", ""]
    rows = [("start", list(b1), list(b4), list(b7), None)]
    b7s = list(b7); b7s[3] = "c"
    rows.append(("move 1: c  b4 -> b7", list(b1), list(b4), b7s, (2, 3)))
    b4s = list(b4); b4s[0] = "a"
    rows.append(("move 2: a  b1 -> b4", list(b1), b4s, b7s, (1, 0)))
    b1s = list(b1); b1s[0] = "x"
    rows.append(("move 3: x  -> b1", b1s, b4s, b7s, (0, 0)))
    y = 88
    for label, r1, r4, r7, hl in rows:
        body += text(366, y + 12, label, HI, 12, "start", "bold")
        for bi, (name, keys) in enumerate((("b1", r1), ("b4", r4), ("b7", r7))):
            bx = 400 + bi * 172
            body += text(bx - 8, y + 44, name, BLUE, 12, "end", "bold")
            for si, k in enumerate(keys):
                is_hl = hl == (bi, si)
                st = ORANGE if is_hl else (GRAY if not k else FG)
                body += rect(bx + si * 35, y + 22, 31, 32, st, BG, 2.6 if is_hl else 1.2)
                body += text(bx + si * 35 + 15.5, y + 43, k if k else "-", HI if k else DIM, 13, weight="bold")
        y += 84
    body += text(630, 432, "after move 1, c is in b4 and b7 at once; after move 2, a is in", DIM, 11)
    body += text(630, 450, "b1 and b4 at once: a key can be seen twice, never zero times", DIM, 11)
    body += text(630, 472, "orange = slot written in this step", DIM, 11)
    return svg(W, H, body, defs)


if __name__ == "__main__":
    for name, fn in (("kickout-steps.svg", kickout_steps), ("cuckoo-graph.svg", cuckoo_graph),
                     ("bfs-backward-move.svg", bfs_backward)):
        with open(os.path.join(OUT, name), "w", encoding="utf-8") as f:
            f.write(fn())
        print("wrote", name)
