"""Generate simhash-fingerprint.svg and lsh-banding.svg.

The toy hash values are chosen by hand; every derived number in the
figures (column sums, fingerprints, band matches) is computed here.
"""
import os

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


def text(x, y, s, size=13, fill=FG, anchor="start"):
    return f'<text x="{x}" y="{y}" font-size="{size}" fill="{fill}" text-anchor="{anchor}">{s}</text>'


def rect(x, y, w, h, stroke, fill=None, op=0.25, sw=1.2):
    f = f'fill="{fill}" fill-opacity="{op}"' if fill else 'fill="none"'
    return f'<rect x="{x}" y="{y}" width="{w}" height="{h}" rx="4" {f} stroke="{stroke}" stroke-width="{sw}"/>'


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


def write(name, content):
    path = os.path.join(OUT, name)
    with open(path, "w", encoding="utf-8") as f:
        f.write(content)
    print("wrote", path)


# ---------------- SimHash, f = 8 ----------------
F = 8
doc_a = [("near", 3, "10110010"), ("duplicate", 2, "01101011"),
         ("web", 1, "11001001"), ("page", 1, "10111100")]
replace = ("site", 1, "01010110")          # document B: "web" replaced by "site"


def colsum(feats):
    return [sum(w if bits[j] == "1" else -w for _, w, bits in feats) for j in range(F)]


va = colsum(doc_a)
doc_b = [f for f in doc_a if f[0] != "web"] + [replace]
vb = colsum(doc_b)
fpa = [1 if v > 0 else 0 for v in va]
fpb = [1 if v > 0 else 0 for v in vb]
assert all(v != 0 for v in va + vb)
diff = [j for j in range(F) if fpa[j] != fpb[j]]

X0, CW, GAP, CH = 330, 64, 8, 36
cx = lambda j: X0 + j * (CW + GAP)
body = [text(470, 30, "SimHash fingerprint with f = 8 bits (toy hash values)", 16, HI, "middle")]
body.append(text(30, 66, "feature", 12, SUB))
body.append(text(200, 66, "weight", 12, SUB))
for j in range(F):
    body.append(text(cx(j) + CW / 2, 66, f"bit {j}", 12, SUB, "middle"))


def feature_row(y, name, w, bits):
    out = [text(30, y + 24, name, 14, HI), text(200, y + 24, f"w = {w}", 13)]
    for j in range(F):
        one = bits[j] == "1"
        c = GREEN if one else RED
        out.append(rect(cx(j), y, CW, CH, c, c, 0.22))
        out.append(text(cx(j) + 8, y + 14, bits[j], 10, SUB))
        out.append(text(cx(j) + CW / 2, y + 27, ("+" if one else "-") + str(w), 14, HI, "middle"))
    return out


def value_row(y, label, vals):
    out = [text(30, y + 24, label, 13)]
    for j, v in enumerate(vals):
        out.append(rect(cx(j), y, CW, CH, GREY, GREY, 0.25))
        out.append(text(cx(j) + CW / 2, y + 24, f"{v:+d}", 14, HI, "middle"))
    return out


def fp_row(y, label, bits, mark=()):
    out = [text(30, y + 24, label, 13, HI)]
    for j, b in enumerate(bits):
        c = RED if j in mark else (BLUE if b else GREY)
        out.append(rect(cx(j), y, CW, CH, c, c, 0.35 if b or j in mark else 0.15,
                        2.5 if j in mark else 1.2))
        out.append(text(cx(j) + CW / 2, y + 24, str(b), 15, HI, "middle"))
    return out


y = 80
for name, w, bits in doc_a:
    body += feature_row(y, name, w, bits)
    y += CH + 8
body += value_row(y, "V = column sum", va)
y += CH + 8
body += fp_row(y, "fingerprint A (V &gt; 0)", fpa)
y += CH + 22
body.append(f'<line x1="30" y1="{y}" x2="910" y2="{y}" stroke="{GREY}" stroke-dasharray="4 4"/>')
y += 26
body.append(text(30, y, 'document B: replace "web" (w = 1) by "site" (w = 1); other features unchanged', 13))
y += 14
body += feature_row(y, replace[0], replace[1], replace[2])
y += CH + 8
body += value_row(y, "V for document B", vb)
y += CH + 8
body += fp_row(y, "fingerprint B", fpb, mark=diff)
y += CH + 26
body.append(text(470, y, f"Hamming(A, B) = {len(diff)}: only bit {diff[0]} flipped, where |V| was 1 in A",
                 13, HI, "middle"))
write("simhash-fingerprint.svg", svg(940, y + 22, body))
print("SimHash A:", va, fpa, " B:", vb, fpb, " diff:", diff)

# ---------------- banding, k = 12, b = 4, r = 3 ----------------
B, R = 4, 3
sig = {"D1": [5, 2, 7, 1, 4, 4, 9, 0, 3, 6, 6, 2],
       "D2": [5, 2, 8, 1, 4, 4, 9, 1, 3, 2, 6, 2],
       "D3": [3, 2, 7, 0, 4, 1, 8, 0, 3, 6, 5, 2]}
docs = list(sig)
band_color = [BLUE, GREEN, ORANGE, PURPLE]
buckets = []
for bi in range(B):
    groups = {}
    for d in docs:
        groups.setdefault(tuple(sig[d][bi * R:(bi + 1) * R]), []).append(d)
    buckets.append(groups)
cands = sorted({(g[i], g[j]) for grp in buckets for g in grp.values()
                for i in range(len(g)) for j in range(i + 1, len(g))})
agree = {(a, b): sum(x == y for x, y in zip(sig[a], sig[b]))
         for i, a in enumerate(docs) for b in docs[i + 1:]}

body = [text(470, 30, "Banding a MinHash signature: k = 12 rows, b = 4 bands of r = 3 rows", 16, HI, "middle")]
MX, MW, MH = 200, 60, 24
for i, d in enumerate(docs):
    body.append(text(MX + i * (MW + 10) + MW / 2, 66, d, 13, HI, "middle"))
body.append(text(560, 66, "one table per band, key = the band's r values", 12, SUB))
y = 76
for bi in range(B):
    col = band_color[bi]
    top = y
    for ri in range(R):
        row = bi * R + ri
        body.append(text(MX - 12, y + 17, f"row {row}", 11, SUB, "end"))
        for i, d in enumerate(docs):
            body.append(rect(MX + i * (MW + 10), y, MW, MH, col, col, 0.18))
            body.append(text(MX + i * (MW + 10) + MW / 2, y + 17, str(sig[d][row]), 13, HI, "middle"))
        y += MH + 3
    body.append(text(30, top + 46, f"band {bi + 1}", 13, col))
    # bucket table for this band
    by = top + 4
    for key, grp in buckets[bi].items():
        hit = len(grp) > 1
        bcol = GREEN if hit else GREY
        body.append(rect(560, by, 340, 20, bcol, bcol, 0.3 if hit else 0.12, 2 if hit else 1))
        members = ", ".join(grp)
        body.append(text(570, by + 15, f"key ({', '.join(map(str, key))})  -&gt;  {{{members}}}", 12,
                         HI if hit else FG))
        by += 24
    mid = top + (R * (MH + 3)) / 2
    body.append(f'<line x1="{MX + 3 * (MW + 10) + 4}" y1="{mid}" x2="552" y2="{mid}" stroke="{col}" '
                f'stroke-width="1.5" marker-end="url(#arr{bi})"/>')
    y += 16
defs = "<defs>" + "".join(
    f'<marker id="arr{i}" markerWidth="8" markerHeight="8" refX="7" refY="4" orient="auto">'
    f'<path d="M0,0 L8,4 L0,8 z" fill="{c}"/></marker>' for i, c in enumerate(band_color)) + "</defs>"
body.insert(0, defs)
y += 8
pairs = "; ".join(f"{a}-{b} {agree[(a, b)]}/12" for a, b in agree)
body.append(text(30, y, "candidate pairs (share a bucket in at least one band): "
                 + ", ".join(f"({a}, {b})" for a, b in cands), 13, GREEN))
body.append(text(30, y + 24, f"rows that agree: {pairs}", 13))
body.append(text(30, y + 48, "D1-D3 agree on 7 of 12 rows but on no whole band, so they never become a candidate.",
                 13, SUB))
write("lsh-banding.svg", svg(940, y + 66, body))
print("banding candidates:", cands, "agree:", agree)
