#!/usr/bin/env python3
"""
make_embosser.py — Generate embosser plates using manifold3d.

Builds plate + ridge geometry directly in Python and uses manifold3d for robust
boolean operations (no OpenSCAD/CGAL/CadQuery required).

Outputs:
    embossed_plate.stl  — triangular ridge raised above plate (male die)
    debossed_plate.stl  — matching groove recessed into plate (female die)

Usage:
    python3 make_embosser.py <input.svg> [options]

Options:
    --plate-width    W   mm  (default 100)
    --plate-length   L   mm  (default 80)
    --plate-thickness T  mm  (default 3)
    --stroke-width   S   mm  (default 1.0)   ridge base width
    --emboss-depth   D   mm  (default 1.5)   ridge height / groove depth
    --margin         M   mm  (default 5)
    --tolerance      T   mm  (default 0.35)  extra clearance on debossed groove
    --samples        N        (default 200)  points sampled per subpath
    --smooth         N        (default 2)    Laplacian smoothing passes
    --max-miter      M        (default 3.0)  miter scale clamp
    --out-embossed   PATH     (default embossed_plate.stl)
    --out-debossed   PATH     (default debossed_plate.stl)

Requirements:
    pip install manifold3d svgpathtools numpy
"""

import sys, re, argparse, struct
import numpy as np
from xml.etree import ElementTree as ET
import svgpathtools
import manifold3d as m3d


# ---------------------------------------------------------------------------
# SVG parsing
# ---------------------------------------------------------------------------

def parse_matrix(s):
    if not s:
        return np.eye(3)
    m = re.search(r'matrix\(\s*([\d\s,.eE+-]+)\)', s)
    if not m:
        return np.eye(3)
    a, b, c, d, e, f = map(float, re.split(r'[\s,]+', m.group(1).strip()))
    return np.array([[a, c, e], [b, d, f], [0, 0, 1]], dtype=float)


def get_paths_with_transforms(root):
    results = []
    def walk(elem, mat):
        tag = elem.tag.split('}')[-1] if '}' in elem.tag else elem.tag
        local = mat
        if tag == 'g':
            t = elem.get('transform', '')
            if t:
                local = mat @ parse_matrix(t)
        if tag == 'path':
            d = elem.get('d', '')
            t = elem.get('transform', '')
            if d:
                results.append((d, local @ parse_matrix(t) if t else local))
        for child in elem:
            walk(child, local)
    walk(root, np.eye(3))
    return results


def apply_mat(mat, pts):
    pts = np.asarray(pts, dtype=float)
    h = np.hstack([pts, np.ones((len(pts), 1))])
    return (mat @ h.T).T[:, :2]


def split_into_subpaths(path):
    subpaths, current = [], []
    for seg in path:
        if current and abs(seg.start - current[-1].end) > 1e-3:
            subpaths.append(svgpathtools.Path(*current))
            current = []
        current.append(seg)
    if current:
        subpaths.append(svgpathtools.Path(*current))
    return [sp for sp in subpaths if len(sp) >= 2]


def sample_arclen(path, n):
    try:
        length = path.length(error=1e-4)
    except Exception:
        return []
    if length < 1e-6:
        return []
    pts = []
    for i in range(n):
        try:
            t = path.ilength(length * i / n, error=1e-4)
            p = path.point(t)
            pts.append([p.real, p.imag])
        except Exception:
            pass
    return pts


def deduplicate(pts, min_dist=0.1):
    if len(pts) < 3:
        return pts
    keep = [pts[0]]
    for p in pts[1:]:
        if np.linalg.norm(np.array(p) - np.array(keep[-1])) >= min_dist:
            keep.append(p)
    while len(keep) > 3 and np.linalg.norm(np.array(keep[-1]) - np.array(keep[0])) < min_dist:
        keep.pop()
    return keep


def smooth_path(pts, passes=2):
    pts = np.array(pts, dtype=float)
    for _ in range(passes):
        pts = 0.25*np.roll(pts, 1, axis=0) + 0.5*pts + 0.25*np.roll(pts, -1, axis=0)
    return pts.tolist()


def load_subpaths(svg_file, samples, smooth_passes, scale_x, scale_y, vbw, vbh):
    tree = ET.parse(svg_file)
    root = tree.getroot()

    flip   = np.array([[1,0,0],[0,-1,vbh],[0,0,1]], dtype=float)
    centre = np.array([[1,0,-vbw/2],[0,1,-vbh/2],[0,0,1]], dtype=float)
    scale  = np.array([[scale_x,0,0],[0,scale_y,0],[0,0,1]], dtype=float)
    post   = scale @ centre @ flip

    subpaths = []
    for d, svg_mat in get_paths_with_transforms(root):
        full_mat = post @ svg_mat
        try:
            path = svgpathtools.parse_path(d)
        except Exception:
            continue
        for sp in split_into_subpaths(path):
            raw = sample_arclen(sp, samples)
            if len(raw) < 3:
                continue
            pts = apply_mat(full_mat, np.array(raw)).tolist()
            pts = deduplicate(pts)
            if len(pts) < 3:
                continue
            if smooth_passes > 0:
                pts = smooth_path(pts, smooth_passes)
            subpaths.append(pts)
    return subpaths


# ---------------------------------------------------------------------------
# Polyhedron generation  (same algorithm as svg_to_ridge.py)
# ---------------------------------------------------------------------------

def path_signed_area(pts):
    pts = np.asarray(pts, dtype=float)
    x, y = pts[:, 0], pts[:, 1]
    return 0.5 * float(np.sum(x * np.roll(y, -1) - np.roll(x, -1) * y))


def miter_frame(pts, max_miter):
    N = len(pts)
    nxt = np.roll(pts, -1, axis=0)
    seg_t = nxt - pts
    seg_len = np.linalg.norm(seg_t, axis=1, keepdims=True)
    seg_len = np.where(seg_len < 1e-10, 1e-10, seg_len)
    seg_t /= seg_len
    seg_n = np.column_stack([-seg_t[:, 1], seg_t[:, 0]])
    prev_n = np.roll(seg_n, 1, axis=0)
    bisect = seg_n + prev_n
    bl = np.linalg.norm(bisect, axis=1, keepdims=True)
    degenerate = (bl < 1e-6).flatten()
    bl = np.where(bl < 1e-6, 1.0, bl)
    bisect /= bl
    cos_half = np.clip(np.sum(bisect * seg_n, axis=1), 1e-6, None)
    scale = np.clip(1.0 / cos_half, 0.0, max_miter)
    result = bisect * scale[:, None]
    result[degenerate] = seg_n[degenerate]
    return result


def ridge_polyhedron(pts, half_width, ridge_height, embed, max_miter):
    """Upward-pointing triangular ridge. L/R base at z=-embed, tip at z=ridge_height."""
    pts = np.asarray(pts, dtype=float)
    N = len(pts)
    if path_signed_area(pts) < 0:
        pts = pts[::-1]
    mperp = miter_frame(pts, max_miter)

    Lv = np.column_stack([pts + mperp * half_width, np.full(N, -embed)])
    Rv = np.column_stack([pts - mperp * half_width, np.full(N, -embed)])
    Tv = np.column_stack([pts,                       np.full(N,  ridge_height)])

    verts = np.empty((3 * N, 3))
    verts[0::3] = Lv
    verts[1::3] = Rv
    verts[2::3] = Tv

    faces = []
    for i in range(N):
        j = (i + 1) % N
        L0, R0, T0 = 3*i,   3*i+1, 3*i+2
        L1, R1, T1 = 3*j,   3*j+1, 3*j+2
        faces += [[R0, L0, L1], [R0, L1, R1]]   # bottom (−Z outward)
        faces += [[L0, T0, T1], [L0, T1, L1]]   # left slope
        faces += [[R0, R1, T1], [R0, T1, T0]]   # right slope

    return verts.tolist(), faces


def to_manifold(verts, faces):
    mesh = m3d.Mesh(
        vert_properties=np.array(verts, dtype=np.float32),
        tri_verts=np.array(faces, dtype=np.uint32),
    )
    result = m3d.Manifold(mesh=mesh)
    if result.is_empty():
        raise ValueError("manifold3d rejected mesh (likely self-intersecting geometry)")
    return result


# ---------------------------------------------------------------------------
# Plate builders
# ---------------------------------------------------------------------------

def box_manifold(pw, pl, pt):
    """Box from (−pw/2, −pl/2, 0) to (pw/2, pl/2, pt)."""
    return m3d.Manifold.cube([pw, pl, pt], center=True).translate([0.0, 0.0, pt / 2.0])


def build_embossed_plate(subpaths, pw, pl, pt, stroke_width, emboss_depth, max_miter):
    print(f"  Base plate {pw}×{pl}×{pt} mm")
    plate = box_manifold(pw, pl, pt)
    embed = emboss_depth * 0.1

    for i, pts in enumerate(subpaths):
        print(f"  Ridge {i+1}/{len(subpaths)} ({len(pts)} pts)...", end=' ', flush=True)
        try:
            verts, faces = ridge_polyhedron(pts, stroke_width / 2, emboss_depth, embed, max_miter)
            ridge = to_manifold(verts, faces).translate([0.0, 0.0, pt])
            plate = plate + ridge
            print("ok")
        except Exception as ex:
            print(f"SKIPPED ({ex})")

    return plate


def build_debossed_plate(subpaths, pw, pl, pt, stroke_width, emboss_depth, tolerance, max_miter):
    print(f"  Base plate {pw}×{pl}×{pt} mm")
    plate = box_manifold(pw, pl, pt)
    groove_half  = stroke_width / 2 + tolerance
    groove_depth = emboss_depth + 0.2

    for i, pts in enumerate(subpaths):
        print(f"  Groove {i+1}/{len(subpaths)} ({len(pts)} pts)...", end=' ', flush=True)
        try:
            # Mirror path in X so groove aligns when plate is flipped face-to-face
            mirrored = [[-x, y] for x, y in pts]
            verts, faces = ridge_polyhedron(mirrored, groove_half, groove_depth, 0.0, max_miter)
            # Flip the ridge downward then raise to plate surface
            groove = to_manifold(verts, faces).mirror([0, 0, 1]).translate([0.0, 0.0, pt])
            plate = plate - groove
            print("ok")
        except Exception as ex:
            print(f"SKIPPED ({ex})")

    return plate


# ---------------------------------------------------------------------------
# STL export (binary)
# ---------------------------------------------------------------------------

def export_stl(manifold_obj, path):
    mesh = manifold_obj.to_mesh()
    verts = np.array(mesh.vert_properties, dtype=np.float32)
    faces = np.array(mesh.tri_verts, dtype=np.int64)

    with open(path, 'wb') as f:
        f.write(b'\x00' * 80)
        f.write(struct.pack('<I', len(faces)))
        for tri in faces:
            v0, v1, v2 = verts[tri[0]], verts[tri[1]], verts[tri[2]]
            n = np.cross(v1 - v0, v2 - v0)
            nl = np.linalg.norm(n)
            if nl > 1e-10:
                n /= nl
            f.write(struct.pack('<3f', *n))
            f.write(struct.pack('<3f', *v0))
            f.write(struct.pack('<3f', *v1))
            f.write(struct.pack('<3f', *v2))
            f.write(b'\x00\x00')

    print(f"  → {path}  ({len(faces)} triangles)")


# ---------------------------------------------------------------------------
# Main
# ---------------------------------------------------------------------------

def main():
    ap = argparse.ArgumentParser(formatter_class=argparse.RawDescriptionHelpFormatter,
                                 description=__doc__)
    ap.add_argument('svg_file')
    ap.add_argument('--plate-width',     type=float, default=100)
    ap.add_argument('--plate-length',    type=float, default=80)
    ap.add_argument('--plate-thickness', type=float, default=3)
    ap.add_argument('--stroke-width',    type=float, default=1.0)
    ap.add_argument('--emboss-depth',    type=float, default=1.5)
    ap.add_argument('--margin',          type=float, default=5)
    ap.add_argument('--tolerance',       type=float, default=0.35)
    ap.add_argument('--samples',         type=int,   default=200)
    ap.add_argument('--smooth',          type=int,   default=2)
    ap.add_argument('--max-miter',       type=float, default=3.0)
    ap.add_argument('--out-embossed',    default='embossed_plate.stl')
    ap.add_argument('--out-debossed',    default='debossed_plate.stl')
    args = ap.parse_args()

    root = ET.parse(args.svg_file).getroot()
    vb_str = root.get('viewBox', '0 0 799 585')
    _, _, vbw, vbh = map(float, vb_str.split())

    s = min((args.plate_width  - 2*args.margin) / vbw,
            (args.plate_length - 2*args.margin) / vbh)

    print(f"SVG: {args.svg_file}  scale={s:.5f}")
    subpaths = load_subpaths(
        args.svg_file, args.samples, args.smooth,
        scale_x=s, scale_y=s, vbw=vbw, vbh=vbh,
    )
    print(f"Found {len(subpaths)} subpath(s).")
    if not subpaths:
        sys.exit("ERROR: no usable subpaths found.")

    print("\n--- Embossed plate ---")
    ep = build_embossed_plate(
        subpaths, args.plate_width, args.plate_length, args.plate_thickness,
        args.stroke_width, args.emboss_depth, args.max_miter,
    )
    export_stl(ep, args.out_embossed)

    print("\n--- Debossed plate ---")
    dp = build_debossed_plate(
        subpaths, args.plate_width, args.plate_length, args.plate_thickness,
        args.stroke_width, args.emboss_depth, args.tolerance, args.max_miter,
    )
    export_stl(dp, args.out_debossed)

    print("\nDone.")


if __name__ == '__main__':
    main()
