"""Deterministic, constrained meshes of generated RGBA parts, never video pixels."""
import cv2
import numpy as np
from shapely.geometry import LineString, Point, Polygon, box
from shapely.ops import unary_union
from skimage.morphology import skeletonize
import triangle


def _polygons(shape):
    if shape.is_empty:
        return []
    if shape.geom_type == 'Polygon':
        return [shape]
    return [p for item in shape.geoms for p in _polygons(item)]


def _lines(shape):
    if shape.is_empty:
        return []
    if shape.geom_type in ('LineString', 'LinearRing'):
        return [shape]
    if hasattr(shape, 'geoms'):
        return [line for item in shape.geoms for line in _lines(item)]
    return []


def _edge_paths(binary):
    """Trace a skeleton once per undirected edge, including closed loops."""
    pixels = set(map(tuple, np.argwhere(skeletonize(binary > 0))))
    adjacency = {}
    for y, x in sorted(pixels):
        neighbors = []
        for dy, dx in [(-1, -1), (-1, 0), (-1, 1), (0, -1), (0, 1),
                       (1, -1), (1, 0), (1, 1)]:
            other = (y+dy, x+dx)
            if other not in pixels:
                continue
            if dx and dy and ((y, x+dx) in pixels or (y+dy, x) in pixels):
                continue
            neighbors.append(other)
        adjacency[y, x] = neighbors
    used, paths = set(), []
    starts = sorted(pixels, key=lambda p: (len(adjacency[p]) == 2, p))
    for start in starts:
        for next_point in adjacency[start]:
            edge = tuple(sorted((start, next_point)))
            if edge in used:
                continue
            used.add(edge)
            path = [start, next_point]
            while len(adjacency[path[-1]]) == 2 and path[-1] != start:
                candidates = [p for p in adjacency[path[-1]] if p != path[-2]]
                end = candidates[0]
                edge = tuple(sorted((path[-1], end)))
                if edge in used:
                    break
                used.add(edge)
                path.append(end)
            if len(path) >= 3:
                paths.append(LineString([(x+.5, y+.5) for y, x in path]))
    return paths


def edge_mesh(images, spacing=22, anchors=(), max_size=512):
    """Return UV mesh with enforced boundary/feature edges and measured coverage.

    Spacing is in analysis pixels (long side <= max_size). Alpha geometry is
    the conservative union of all supplied expression states. Internal image
    lines come from Canny, closing and skeleton tracing, with short noise paths
    removed. Anchors are normalized UV points, never inferred from a loss.
    """
    if not images or spacing <= 0:
        raise ValueError('Images and positive spacing are required')
    width, height = images[0].size
    if any(image.size != (width, height) for image in images):
        raise ValueError('Texture states must have matching dimensions')
    scale = min(1., max_size/max(width, height))
    w, h = max(2, round(width*scale)), max(2, round(height*scale))
    rasters = [np.asarray(image.convert('RGBA')) for image in images]
    alpha = np.maximum.reduce([item[..., 3] for item in rasters])
    small_alpha = cv2.resize(alpha, (w, h), interpolation=cv2.INTER_AREA)
    mask = np.uint8(small_alpha > 3)*255
    contours, hierarchy = cv2.findContours(mask, cv2.RETR_CCOMP, cv2.CHAIN_APPROX_SIMPLE)
    polygons = []
    if hierarchy is not None:
        for i, contour in enumerate(contours):
            if hierarchy[0, i, 3] != -1 or len(contour) < 3 or cv2.contourArea(contour) <= 0:
                continue
            holes = []
            child = hierarchy[0, i, 2]
            while child >= 0:
                if len(contours[child]) >= 3:
                    holes.append(contours[child][:, 0].astype(float)+.5)
                child = hierarchy[0, child, 0]
            polygons.extend(_polygons(Polygon(contour[:, 0].astype(float)+.5, holes).buffer(0)))
    if not polygons:
        raise ValueError('Part has no usable alpha domain')
    # Bound each margin by neighboring components, so padding cannot join islands.
    # Shrink holes conservatively, keeping an interior disk even for small holes.
    expanded = []
    for i, polygon in enumerate(polygons):
        separation = min((polygon.distance(other) for j, other in enumerate(polygons) if i != j), default=float('inf'))
        margin = min(1.5, max(0, separation/3))
        shell = Polygon(polygon.exterior).buffer(margin, quad_segs=2)
        for ring in polygon.interiors:
            hole = Polygon(ring)
            clearance = hole.representative_point().distance(hole.boundary)
            shell = shell.difference(hole.buffer(-min(margin, .45*clearance), quad_segs=2))
        expanded.extend(_polygons(shell))
    domain = unary_union(expanded).simplify(.6, preserve_topology=True)
    domain = domain.intersection(box(0, 0, w, h))
    polygons = _polygons(domain)
    interior = domain.buffer(-3)
    rgb = cv2.resize(rasters[0][..., :3], (w, h), interpolation=cv2.INTER_AREA)
    grey = cv2.cvtColor(rgb, cv2.COLOR_RGB2GRAY)
    grey[small_alpha < 16] = 220
    edges = cv2.Canny(cv2.GaussianBlur(grey, (3, 3), .7), 45, 110)
    edges[cv2.erode(mask, np.ones((7, 7), np.uint8)) == 0] = 0
    edges = cv2.morphologyEx(edges, cv2.MORPH_CLOSE, np.ones((3, 3), np.uint8))
    paths = sorted(_edge_paths(edges), key=lambda line: (-line.length, line.bounds))
    features = []
    for path in paths[:100]:
        if path.length < max(12, spacing*.65):
            continue
        simple = path.simplify(1.3, preserve_topology=False)
        if simple.length > 0:
            features.extend(line for line in _lines(simple.intersection(interior)) if line.length >= 8)
    features = _lines(unary_union(features)) if features else []
    points, segments, markers, lookup = [], [], [], {}

    def index(point):
        key = tuple(round(float(v), 8) for v in point)
        if key not in lookup:
            lookup[key] = len(points)
            points.append(key)
        return lookup[key]

    def add_path(coords, marker):
        ids = [index(point) for point in coords]
        for a, b in zip(ids, ids[1:]):
            if a != b:
                segments.append([a, b])
                markers.append([marker])

    holes = []
    for polygon in polygons:
        add_path(polygon.exterior.coords, 1)
        for ring in polygon.interiors:
            add_path(ring.coords, 1)
            holes.append(list(Polygon(ring).representative_point().coords)[0])
    for line in features:
        add_path(line.coords, 2)
    accepted_anchors = 0
    for uv in anchors:
        point = np.asarray(uv)*[w, h]
        if domain.contains(Point(point)):
            index(point)
            accepted_anchors += 1
    request = {'vertices': points, 'segments': segments, 'segment_markers': markers}
    if holes:
        request['holes'] = holes
    result = triangle.triangulate(request, f'pq20a{spacing*spacing/2:.6f}Q')
    vertices = np.asarray(result['vertices'])
    triangles = np.asarray(result['triangles'])
    pos = vertices[triangles]
    signed_area = np.cross(pos[:, 1]-pos[:, 0], pos[:, 2]-pos[:, 0])
    triangles[signed_area < 0] = triangles[signed_area < 0][:, [0, 2, 1]]
    all_edges = {tuple(sorted((int(t[i]), int(t[(i+1) % 3]))))
                 for t in triangles for i in range(3)}
    constraints = {'boundary_edges': [], 'feature_edges': []}
    for edge, marker in zip(result['segments'], result['segment_markers']):
        if tuple(sorted(edge)) not in all_edges:
            raise RuntimeError('Triangulator lost a constraint')
        constraints['feature_edges' if marker[0] == 2 else 'boundary_edges'].append(edge.tolist())
    coverage = np.zeros((height, width), np.uint8)
    source_xy = vertices/[w, h]*[width, height]-.5
    for tri in triangles:
        cv2.fillConvexPoly(coverage, np.rint(source_xy[tri]).astype(np.int32), 1)
    visible = alpha > 16
    coverage_score = float(coverage[visible].mean())
    lengths = np.linalg.norm(pos-np.roll(pos, -1, axis=1), axis=2)
    angles = []
    for i in range(3):
        a, b, c = lengths[:, i], lengths[:, (i+1) % 3], lengths[:, (i+2) % 3]
        angles.append(np.degrees(np.arccos(np.clip((a*a+c*c-b*b)/(2*a*c), -1, 1))))
    minimum_angles = np.min(angles, axis=0)
    return {'uv': (vertices/[w, h]).tolist(), 'triangles': triangles.tolist(), **constraints,
            'stats': {'vertices': len(vertices), 'triangles': len(triangles),
                      'components': len(polygons), 'holes': len(holes),
                      'feature_edges': len(constraints['feature_edges']),
                      'boundary_edges': len(constraints['boundary_edges']),
                      'alpha_coverage': coverage_score, 'uncovered_alpha_pixels': int(((coverage == 0) & visible).sum()),
                      'minimum_angle': float(minimum_angles.min()),
                      'p05_minimum_angle': float(np.percentile(minimum_angles, 5)),
                      'accepted_anchors': accepted_anchors, 'analysis_size': [w, h],
                      'spacing': spacing}}
