#!/usr/bin/env python3
from pathlib import Path
import csv

ROOT = Path(__file__).resolve().parents[1]
BG = '#2d333b'; FG = '#adbac7'; HI = '#cdd9e5'; MUTED = '#768390'
BLUE = '#388bfd'; GREEN = '#3fb950'; ORANGE = '#f0883e'; PURPLE = '#a371f7'; RED = '#f85149'; GRAY = '#636e7b'

def svg(w, h, body):
    return f'''<svg xmlns="http://www.w3.org/2000/svg" width="{w}" height="{h}" viewBox="0 0 {w} {h}">
  <rect width="{w}" height="{h}" fill="{BG}"/>
  <style>text{{font-family:ui-monospace,SFMono-Regular,Menlo,Consolas,monospace;font-size:14px}}.title{{font-size:18px;font-weight:700}}.small{{font-size:12px}}</style>
  <g style="fill:{FG}">
{body}
  </g>
</svg>
'''

def t(x, y, text, cls='', anchor='', fill=None):
    c = f' class="{cls}"' if cls else ''
    a = f' text-anchor="{anchor}"' if anchor else ''
    f = f' style="fill:{fill}"' if fill else ''
    return f'<text x="{x}" y="{y}"{a}{c}{f}>{text}</text>'

def rect(x, y, w, h, label, color, sub=''):
    out = f'<rect x="{x}" y="{y}" width="{w}" height="{h}" rx="10" fill="none" stroke="{color}" stroke-width="2"/>'
    out += t(x + w / 2, y + 24, label, 'title', 'middle', HI)
    if sub:
        out += t(x + w / 2, y + 47, sub, 'small', 'middle', MUTED)
    return out

def arrow(x1, y1, x2, y2, color=FG):
    return f'<path d="M{x1} {y1} L{x2} {y2}" stroke="{color}" stroke-width="2" marker-end="url(#arrow)"/>'

def defs():
    return f'<defs><marker id="arrow" viewBox="0 0 10 10" refX="8" refY="5" markerWidth="6" markerHeight="6" orient="auto"><path d="M0 0 L10 5 L0 10 Z" fill="{FG}"/></marker></defs>'

def draw_layers():
    body = [defs(), t(30, 34, 'Three allocator paths share the same pressure points', 'title', fill=HI)]
    cols = [('jemalloc', BLUE), ('gperftools', GREEN), ('mimalloc', PURPLE)]
    labels = [
        [('tcache', 'thread-local bins'), ('arena bins', 'slabs by size class'), ('pa shard', 'extents and decay')],
        [('thread cache', 'per-thread free lists'), ('central freelist', 'spans by size class'), ('page heap', 'page map and spans')],
        [('heap pages', 'current page per size'), ('sharded free', 'free/local/xthread'), ('segments', 'aligned chunks')],
    ]
    for c, (name, color) in enumerate(cols):
        x = 40 + c * 270
        body.append(t(x + 95, 74, name, 'title', 'middle', color))
        for r, (lab, sub) in enumerate(labels[c]):
            y = 95 + r * 105
            body.append(rect(x, y, 190, 70, lab, color, sub))
            if r < 2:
                body.append(arrow(x + 95, y + 72, x + 95, y + 101))
    body.append(t(40, 430, 'Fast path tries to stop at the top box; fragmentation and RSS are decided by lower boxes.', 'small', fill=MUTED))
    (ROOT / 'allocator-layers.svg').write_text(svg(850, 460, '\n'.join(body)), encoding='utf-8')

def draw_sharding():
    body = [defs(), t(30, 34, 'mimalloc free list sharding', 'title', fill=HI)]
    body.append(rect(45, 80, 150, 72, 'owner thread', 'none', 'malloc result'))
    body.append(rect(690, 80, 150, 72, 'remote thread', 'none', 'cross-thread free'))
    body.append(rect(270, 65, 230, 290, 'mi_page_t', ORANGE, 'one size class'))
    lists = [('free', 135, BLUE, 'malloc pops'), ('local_free', 205, GREEN, 'owner-local staging'), ('xthread_free', 275, RED, 'atomic remote list')]
    for lab, y, color, sub in lists:
        body.append(f'<rect x="300" y="{y}" width="150" height="34" rx="6" fill="none" stroke="{color}" stroke-width="2"/>')
        body.append(t(375, y + 22, lab, anchor='middle'))
        body.append(t(512, y + 22, sub, 'small', fill=MUTED))
    body += [
        arrow(300, 152, 198, 130, BLUE),
        f'<path d="M765 152 L765 372 L385 372 L385 312" fill="none" stroke="{RED}" stroke-width="2" marker-end="url(#arrow)"/>',
        arrow(385, 275, 385, 242, RED),
        arrow(385, 205, 385, 172, GREEN),
        t(395, 260, 'collect', 'small', fill=MUTED),
        t(395, 195, 'merge', 'small', fill=MUTED),
        t(222, 160, 'pop', 'small', fill=MUTED),
        t(55, 395, 'Remote frees go to xthread_free; collection moves them upward.', 'small', fill=MUTED),
        t(55, 419, 'Allocation pops only from free on the fast path.', 'small', fill=MUTED),
    ]
    (ROOT / 'remote-free-sharding.svg').write_text(svg(870, 455, '\n'.join(body)), encoding='utf-8')

def draw_rss():
    path = ROOT / 'reproduce' / 'results' / 'summary.tsv'
    rows = []
    if path.exists():
        with path.open(newline='') as f:
            rows = list(csv.DictReader(f, delimiter='\t'))
    rows = [r for r in rows if r.get('mode') == 'cross']
    order = ['glibc', 'jemalloc', 'tcmalloc', 'mimalloc']
    rows.sort(key=lambda r: order.index(r['allocator']) if r['allocator'] in order else 99)
    body = [defs(), t(30, 34, 'Cross-thread workload: median ratios, 3 seeds', 'title', fill=HI)]
    x0, y0, w, h = 80, 70, 560, 250
    body.append(f'<line x1="{x0}" y1="{y0+h}" x2="{x0+w}" y2="{y0+h}" stroke="{MUTED}"/>')
    body.append(f'<line x1="{x0}" y1="{y0}" x2="{x0}" y2="{y0+h}" stroke="{MUTED}"/>')
    maxv = max([float(r['rss_live_ratio']) for r in rows] + [1.6])
    maxv = max(1.6, round(maxv + 0.2, 1))
    for tick in range(0, 5):
        v = 1.0 + tick * (maxv - 1.0) / 4
        y = y0 + h - (v - 1.0) / (maxv - 1.0) * h
        body.append(f'<line x1="{x0-5}" y1="{y:.1f}" x2="{x0+w}" y2="{y:.1f}" stroke="{GRAY}" stroke-opacity="0.35"/>')
        body.append(t(x0 - 12, f'{y+4:.1f}', f'{v:.2f}', 'small', 'end', MUTED))
    for i, r in enumerate(rows):
        x = x0 + 35 + i * 125
        barw = 48
        for j, m in enumerate(['usable_live_ratio', 'rss_live_ratio']):
            v = float(r[m])
            bh = (v - 1.0) / (maxv - 1.0) * h
            color = BLUE if j else ORANGE
            body.append(f'<rect x="{x+j*(barw+8)}" y="{y0+h-bh:.1f}" width="{barw}" height="{bh:.1f}" fill="{color}" opacity="0.9"/>')
            body.append(t(x + j * (barw + 8) + barw / 2, f'{y0+h-bh-6:.1f}', f'{v:.2f}', 'small', 'middle', MUTED))
        body.append(t(x + barw, y0 + h + 24, r['allocator'], anchor='middle'))
    body.append(f'<rect x="90" y="355" width="16" height="16" fill="{ORANGE}"/>')
    body.append(t(113, 368, 'usable/live', 'small', fill=MUTED))
    body.append(f'<rect x="230" y="355" width="16" height="16" fill="{BLUE}"/>')
    body.append(t(253, 368, 'RSS/live', 'small', fill=MUTED))
    body.append(t(80, 405, 'Lower is better. RSS includes allocator caches and unmapped-delay effects.', 'small', fill=MUTED))
    (ROOT / 'rss-ratio.svg').write_text(svg(700, 455, '\n'.join(body)), encoding='utf-8')

if __name__ == '__main__':
    draw_layers()
    draw_sharding()
    draw_rss()
