#!/usr/bin/env python3
"""RC17 pose sprite cleanup + runtime sequence linking.

This is intentionally deterministic and conservative:
- backs up active pose sprites before overwrite;
- removes app-arrow remnants / detached neighbor fragments from active transparent sprites;
- re-extracts the two light-background keyed sets from their source illustrations;
- updates pose cleanup metadata from live pixels;
- links every exercise that has an available pose family and appears in DATA.
"""
from __future__ import annotations
import json, re, shutil
from collections import Counter, defaultdict
from pathlib import Path

import numpy as np
from PIL import Image
from scipy import ndimage

ROOT = Path(__file__).resolve().parents[1]
LIB_PATH = ROOT / 'assets/data/image-admin/pose-sprite-library.json'
SEQ_PATH = ROOT / 'assets/data/image-admin/exercise-pose-sequences.json'
DATA_PATH = ROOT / 'assets/js/data.js'
SPRITE_BACKUP = ROOT / 'incoming/pose-cleanup/rc17-original-active-sprites'
CANDIDATE_DIR = ROOT / 'incoming/pose-cleanup/rc17-cleaned-candidates'
LIGHT_SOURCE_REEXTRACT = {'burpee', 'lateral-line-hops'}

SEQUENCE_DEFS = {
    'linear-pogos': dict(label='Linear Pogos', loopMs=1100, y=[0, -12, -58, 0]),
    'lateral-pogo-jumps': dict(label='Lateral Pogo Jumps', loopMs=1100, y=[0, -24, -18, 0]),
    'lateral-line-hops': dict(label='Lateral Line Hops', loopMs=1050, y=[0, -22, -46, 0]),
    'broad-jump': dict(label='Broad Jump', loopMs=1450, poseKeys=['broad-jump__pose01','broad-jump__pose02','broad-jump__pose04','broad-jump__pose05'], y=[0,-44,0,0]),
    'seated-box-jump': dict(label='Seated Box Jump', loopMs=1450, poseKeys=['seated-box-jump__pose01','seated-box-jump__pose02','seated-box-jump__pose04'], y=[0,-34,0], floorY=86),
    'carioca-shuffle': dict(label='Carioca Shuffle', loopMs=1150),
    'lateral-power-shuffle': dict(label='Lateral Power Shuffle', loopMs=1150),
    'db-farmer-carry': dict(label='DB Farmer Carry', loopMs=1400),
    'goblet-squat': dict(label='Goblet Squat', loopMs=1300),
    'single-arm-row': dict(label='Single Arm Row', loopMs=1150),
    'med-ball-slam': dict(label='Med Ball Slam', loopMs=1350, poseKeys=['med-ball-slam__pose01','med-ball-slam__pose03','med-ball-slam__pose04']),
    'kneeling-soccer-toss': dict(label='Kneeling Soccer Toss', loopMs=1300, poseKeys=['kneeling-soccer-toss__pose01','kneeling-soccer-toss__pose02','kneeling-soccer-toss__pose03']),
    'side-plank-hip-lift': dict(label='Side Plank Hip Lift', loopMs=1250),
    'sit-up': dict(label='Sit Up', loopMs=1500),
    'plank': dict(label='Plank', loopMs=1200),
    'plank-walkouts': dict(label='Plank Walkouts', loopMs=1450),
    'bear-crawl': dict(label='Bear Crawl', loopMs=1100),
    'push-up': dict(label='Push Up', loopMs=1250, poseKeys=['push-up__pose01','push-up__pose02','push-up__pose04']),
    'pull-up': dict(label='Pull Up', loopMs=1300, poseKeys=['pull-up__pose01','pull-up__pose03','pull-up__pose05']),
    'chin-up': dict(label='Chin Up', loopMs=1300),
    'buddy-hamstring-curls': dict(label='Buddy Hamstring Curls', loopMs=1450, floorY=86),
    'burpee': dict(label='Burpee', loopMs=1800, poseKeys=['burpee__pose01','burpee__pose02','burpee__pose03','burpee__pose05'], y=[0,0,0,-48]),
}

DEFAULT_LAYOUT = {
    'slotCount': 5,
    'floorY': 84,
    'baseScale': 1,
    'referenceStandingPx': 820,
    'baseStandingHeightPct': 78,
}


def normalize_key(value: str) -> str:
    return re.sub(r'^-+|-+$', '', re.sub(r'[^a-z0-9]+', '-', value.strip().lower()))


def load_data_exercise_keys() -> set[str]:
    text = DATA_PATH.read_text(encoding='utf-8')
    m = re.search(r'const DATA = (\[.*?\]);\n', text, re.S)
    if not m:
        return set()
    rows = json.loads(m.group(1))
    return {normalize_key(r.get('exercise', '')) for r in rows if r.get('exercise')}


def bright_arrow_mask(arr: np.ndarray) -> np.ndarray:
    rgb = arr[:, :, :3].astype(np.int16)
    r, g, b = rgb[:, :, 0], rgb[:, :, 1], rgb[:, :, 2]
    return (((b > 135) & (g > 105) & (r < 95) & ((b - r) > 45) & ((g - r) > 30)) |
            ((g > 140) & (b > 100) & (r < 85) & ((g - r) > 55)))


def keep_meaningful_components(mask: np.ndarray, min_ratio: float = 0.08, min_area: int = 1200) -> tuple[np.ndarray, dict]:
    lab, n = ndimage.label(mask, structure=np.ones((3, 3), int))
    metrics = {'componentCount': int(n), 'largestArea': 0, 'keptComponents': 0}
    if not n:
        return np.zeros_like(mask, dtype=bool), metrics
    areas = np.bincount(lab.ravel())[1:]
    largest = int(areas.max())
    keep = np.zeros_like(mask, dtype=bool)
    kept = 0
    for idx, area in enumerate(areas, start=1):
        if area == largest or area > max(min_area, largest * min_ratio):
            keep |= lab == idx
            kept += 1
    keep = ndimage.binary_closing(keep, structure=np.ones((3, 3), bool), iterations=1)
    metrics.update({'largestArea': largest, 'keptComponents': kept})
    return keep, metrics


def clean_existing_sprite(path: Path) -> Image.Image:
    im = Image.open(path).convert('RGBA')
    arr = np.array(im)
    arr[:, :, 3] = np.where(bright_arrow_mask(arr), 0, arr[:, :, 3]).astype(np.uint8)
    keep, _ = keep_meaningful_components(arr[:, :, 3] > 20)
    arr[:, :, 3] = np.where(keep, arr[:, :, 3], 0).astype(np.uint8)
    return Image.fromarray(arr, 'RGBA')


def reextract_light_source(pose: dict) -> Image.Image:
    src = Image.open(ROOT / 'assets/img/exercises/figures' / f"{pose['sourceExercise']}.webp").convert('RGBA')
    x1, y1, x2, y2 = pose['bbox']
    arr = np.array(src.crop((x1, y1, x2, y2)).convert('RGBA'))
    rgb = arr[:, :, :3].astype(np.int16)
    r, g, b = rgb[:, :, 0], rgb[:, :, 1], rgb[:, :, 2]
    bg = ((r > 232) & (g > 232) & (b > 232) &
          ((np.maximum.reduce([r, g, b]) - np.minimum.reduce([r, g, b])) < 32))
    mask = (~bg) & (~bright_arrow_mask(arr)) & (arr[:, :, 3] > 5)
    keep, _ = keep_meaningful_components(mask, min_ratio=0.035, min_area=400)
    keep = ndimage.binary_fill_holes(keep)
    arr[:, :, 3] = np.where(keep, 255, 0).astype(np.uint8)
    return Image.fromarray(arr, 'RGBA')


def sprite_metrics(path: Path) -> dict:
    arr = np.array(Image.open(path).convert('RGBA'))
    alpha = arr[:, :, 3]
    visible = alpha > 20
    lab, n = ndimage.label(visible, structure=np.ones((3, 3), int))
    areas = np.bincount(lab.ravel())[1:] if n else np.array([], dtype=int)
    largest = int(areas.max()) if len(areas) else 0
    largest_label = int(areas.argmax() + 1) if len(areas) else 0
    separated_mask = visible & (lab != largest_label) if largest_label else np.zeros_like(visible, dtype=bool)
    separated = int(separated_mask.sum())
    edge = int(visible[0, :].sum() + visible[-1, :].sum() + visible[:, 0].sum() + visible[:, -1].sum())
    # Bright cyan inside the athlete shirt/shorts is not an artifact. Only count it
    # as a cleanup flag when it remains in detached/non-primary components.
    bright_separated = bright_arrow_mask(arr) & separated_mask
    return {
        'visiblePixels': int(visible.sum()),
        'separatedPixels': separated,
        'separatedComponentCount': int(max(0, n - 1)),
        'brightBlueSeparatedPixels': int(bright_separated.sum()),
        'brightMagentaSeparatedPixels': 0,
        'largestComponentRatio': round((largest / max(1, int(visible.sum()))), 3),
        'edgeTouchPixels': edge,
    }


def cleanup_flag(metrics: dict) -> list[str]:
    flags = []
    if metrics['brightBlueSeparatedPixels'] > 30:
        flags.append('embedded-blue-arrow-remnant')
    if metrics['separatedPixels'] > max(1100, metrics['visiblePixels'] * 0.025):
        flags.append('extra-separated-parts')
    if metrics['edgeTouchPixels'] > 120:
        flags.append('touches-canvas-edge')
    return flags


def build_sequences(library: dict) -> dict:
    data_keys = load_data_exercise_keys()
    by_exercise = defaultdict(list)
    for pose in library.get('poses', []):
        by_exercise[pose.get('sourceExercise', '')].append(pose['poseKey'])
    sequences = {}
    linked = sorted((data_keys | {'burpee'}) & set(SEQUENCE_DEFS.keys()))
    for key in linked:
        definition = SEQUENCE_DEFS[key]
        pose_keys = definition.get('poseKeys') or sorted(by_exercise.get(key, []))
        if not pose_keys:
            continue
        y_values = definition.get('y') or [0] * len(pose_keys)
        x_values = definition.get('x') or [0] * len(pose_keys)
        scale_values = definition.get('scale') or [1] * len(pose_keys)
        poses = []
        for idx, pose_key in enumerate(pose_keys):
            poses.append({
                'poseKey': pose_key,
                'x': x_values[idx] if idx < len(x_values) else 0,
                'y': y_values[idx] if idx < len(y_values) else 0,
                'scale': scale_values[idx] if idx < len(scale_values) else 1,
            })
        layout = DEFAULT_LAYOUT.copy()
        if 'floorY' in definition:
            layout['floorY'] = definition['floorY']
        sequences[key] = {
            'exerciseKey': key,
            'label': definition['label'],
            'status': 'runtime-linked-prototype',
            'loopMs': definition['loopMs'],
            'poses': poses,
            'notes': 'RC17: linked from cleaned pose sprites for pass-player motion fallback. App-owned arrows remain optional and default to none.',
            'layout': layout,
            'arrows': [],
        }
    return sequences


def main() -> None:
    library = json.loads(LIB_PATH.read_text(encoding='utf-8'))
    SPRITE_BACKUP.mkdir(parents=True, exist_ok=True)
    CANDIDATE_DIR.mkdir(parents=True, exist_ok=True)

    for pose in library.get('poses', []):
        src = ROOT / pose['src']
        backup = SPRITE_BACKUP / src.name
        if not backup.exists():
            shutil.copy2(src, backup)
        if pose.get('sourceExercise') in LIGHT_SOURCE_REEXTRACT:
            cleaned = reextract_light_source(pose)
            pose['sourceType'] = 'source-reextract-light-bg'
        else:
            cleaned = clean_existing_sprite(src)
        cleaned.save(CANDIDATE_DIR / src.name, 'WEBP', quality=92, method=4)
        cleaned.save(src, 'WEBP', quality=92, method=4)
        metrics = sprite_metrics(src)
        flags = cleanup_flag(metrics)
        cleanup_status = 'ok' if not flags else 'review'
        pose['cleanup'] = {
            'status': cleanup_status,
            'flags': flags,
            'metrics': metrics,
            'defaultUse': 'allowed-for-runtime-prototype' if cleanup_status == 'ok' else 'review-before-runtime',
            'rc17': 'cleaned active sprite from deterministic cleanup/re-extract pipeline',
        }
        tags = [t for t in pose.get('tags', []) if t not in {'cleanup_review', 'clean_pose'}]
        tags.append('clean_pose' if cleanup_status == 'ok' else 'cleanup_review')
        pose['tags'] = sorted(set(tags))

    status = Counter((p.get('cleanup') or {}).get('status', 'unknown') for p in library.get('poses', []))
    flag_counts = Counter(flag for p in library.get('poses', []) for flag in (p.get('cleanup') or {}).get('flags', []))
    library['version'] = 'v97-rc17-pose-runtime-polish'
    library['rc17Notes'] = {
        'spriteBackup': str(SPRITE_BACKUP.relative_to(ROOT)),
        'cleanedCandidates': str(CANDIDATE_DIR.relative_to(ROOT)),
        'method': 'detached component cleanup + bright app-arrow removal; burpee/lateral-line-hops re-extracted from light source figures',
    }
    library['cleanupSummary'] = {
        'ok': int(status.get('ok', 0)),
        'review': int(status.get('review', 0)),
        'needsCleanup': int(status.get('needs-cleanup', 0)),
        'flags': dict(flag_counts),
        'method': library['rc17Notes']['method'],
    }
    LIB_PATH.write_text(json.dumps(library, indent=2, ensure_ascii=False) + '\n', encoding='utf-8')

    seq_data = json.loads(SEQ_PATH.read_text(encoding='utf-8'))
    seq_data['version'] = 'v97-rc17-pose-runtime-polish'
    seq_data['intent'] = 'Cleaned pose sprites and linked every available pose family that appears in Stage 2 pass data to the optional runtime motion fallback.'
    seq_data['defaultLoopMs'] = 1300
    seq_data['sequences'] = build_sequences(library)
    seq_data['defaultLayout'] = DEFAULT_LAYOUT
    seq_data['defaultArrows'] = []
    seq_data['uiDefaults'] = {
        'arrowVisibility': 'app-owned overlays remain optional and default off',
        'runtimeMotion': 'CSS-only stepped highlight driven by per-sequence loopMs',
    }
    SEQ_PATH.write_text(json.dumps(seq_data, indent=2, ensure_ascii=False) + '\n', encoding='utf-8')

    print('RC17 pose cleanup + sequence link complete')
    print('poses:', len(library.get('poses', [])), 'cleanup:', dict(status), 'flags:', dict(flag_counts))
    print('sequences:', len(seq_data['sequences']))


if __name__ == '__main__':
    main()
