#!/usr/bin/env python3
"""
FreeExile 2.5D Isometric Animation Pipeline & Motion Matching Generator.
Generates multi-frame animation sprite atlases, PBR normal maps, and kinematics manifests.

Features:
- Character & Monster Kinematic Frame Synthesizer (Idle, Run, Attack, Skill, Dodge, Hurt, Death)
- Foot-Speed Synchronizer (Stride-length to Velocity motion matching to eliminate foot sliding)
- Action-Speed Scaler (Windup, Impact Hit Frame, Recovery phase alignment)
- 8-Directional Isometric Heading Resolver
- Tangent-Space Normal Map Baker for PBR dynamic lighting
"""

from __future__ import annotations

import argparse
import json
import math
import os
from pathlib import Path
import sys
from typing import Any, Dict, List, Tuple

import numpy as np
from PIL import Image, ImageChops, ImageDraw, ImageEnhance, ImageFilter

WORKSPACE_ROOT = Path(__file__).resolve().parent.parent.parent
ANIMATIONS_OUT_DIR = WORKSPACE_ROOT / "client" / "webapp" / "assets" / "animations"


def generate_sobel_normal_map(albedo_img: Image.Image, strength: float = 2.5) -> Image.Image:
    """Computes tangent-space normal map using vectorized NumPy operations."""
    gray = albedo_img.convert("L")
    arr = np.array(gray, dtype=np.float32) / 255.0
    padded = np.pad(arr, 1, mode="edge")

    # Vectorized Sobel convolution
    dx = (
        -1.0 * padded[:-2, :-2] + 1.0 * padded[:-2, 2:]
        - 2.0 * padded[1:-1, :-2] + 2.0 * padded[1:-1, 2:]
        - 1.0 * padded[2:, :-2] + 1.0 * padded[2:, 2:]
    ) * strength

    dy = (
        -1.0 * padded[:-2, :-2] - 2.0 * padded[:-2, 1:-1] - 1.0 * padded[:-2, 2:]
        + 1.0 * padded[2:, :-2] + 2.0 * padded[2:, 1:-1] + 1.0 * padded[2:, 2:]
    ) * strength

    dz = np.ones_like(arr)
    len_vec = np.sqrt(dx**2 + dy**2 + dz**2)
    len_vec[len_vec == 0] = 1.0
    nx = dx / len_vec
    ny = dy / len_vec
    nz = dz / len_vec

    r = ((nx * 0.5 + 0.5) * 255).astype(np.uint8)
    g = ((ny * 0.5 + 0.5) * 255).astype(np.uint8)
    b = ((nz * 0.5 + 0.5) * 255).astype(np.uint8)

    alpha = albedo_img.split()[-1] if albedo_img.mode == "RGBA" else Image.new("L", albedo_img.size, 255)
    return Image.merge("RGBA", (Image.fromarray(r), Image.fromarray(g), Image.fromarray(b), alpha))


def synthesize_hero_frame(
    base_img: Image.Image,
    anim_name: str,
    frame_idx: int,
    total_frames: int,
    frame_size: Tuple[int, int] = (160, 192),
) -> Image.Image:
    """Synthesizes an animation frame for the Savage Primal Hero with kinematically sound deformation."""
    fw, fh = frame_size
    frame = Image.new("RGBA", (fw, fh), (0, 0, 0, 0))

    # Resize base hero proportionally to fit the frame canvas
    aspect = base_img.width / base_img.height
    target_h = int(fh * 0.82)
    target_w = int(target_h * aspect)
    scaled_hero = base_img.resize((target_w, target_h), Image.Resampling.LANCZOS)

    t = frame_idx / max(1, total_frames)
    phase = t * math.pi * 2.0

    ox = (fw - target_w) // 2
    oy = fh - target_h - 14

    squash_x = 1.0
    squash_y = 1.0
    rot_deg = 0.0
    shift_x = 0
    shift_y = 0

    if anim_name == "idle":
        # Natural breathing, subtle abdominal/chest cycle
        breathe = math.sin(phase)
        squash_y = 1.0 + breathe * 0.025
        squash_x = 1.0 - breathe * 0.015
        shift_y = int(breathe * 2.5)

    elif anim_name == "run":
        # 8-phase bipedal running gait:
        # Pelvis drops on foot contact, rises on passing, peaks on flight phase
        # Footstep events at frame 0 (left) and frame 4 (right)
        pelvis_bob = math.cos(phase * 2.0) * 5.0
        shift_y = int(pelvis_bob)
        rot_deg = math.sin(phase) * 3.5  # torso tilt forward & back
        squash_y = 1.0 - abs(math.sin(phase)) * 0.04
        squash_x = 1.0 + abs(math.sin(phase)) * 0.03
        shift_x = int(math.sin(phase) * 3.0)

    elif anim_name == "attack_slash":
        # 3-phase ARPG strike:
        # Frames 0-1: Windup (draw weapon back, coil body)
        # Frame 2: Hit frame (explosive forward lunge, maximum elongation)
        # Frames 3-5: Recovery & follow-through
        if frame_idx == 0:
            rot_deg = -8.0
            shift_x = -4
            shift_y = -2
        elif frame_idx == 1:
            rot_deg = -14.0
            shift_x = -8
            shift_y = -4
            squash_y = 0.95
        elif frame_idx == 2:  # IMPACT HIT FRAME
            rot_deg = 18.0
            shift_x = 14
            shift_y = 4
            squash_x = 1.08
            squash_y = 0.94
        elif frame_idx == 3:
            rot_deg = 10.0
            shift_x = 10
            shift_y = 2
        elif frame_idx == 4:
            rot_deg = 4.0
            shift_x = 5
        else:
            rot_deg = 0.0
            shift_x = 1

    elif anim_name == "skill_whirlwind":
        # 360 spin attack: rapid horizontal squash & scale flip
        angle = t * math.pi * 2.0
        scale_x = math.cos(angle)
        squash_x = abs(scale_x) if abs(scale_x) > 0.1 else 0.1
        shift_y = int(math.sin(angle * 2.0) * 3.0)
        rot_deg = math.sin(angle) * 6.0

    elif anim_name == "dodge":
        # Agility phantom roll / slide
        roll_angle = t * 360.0
        rot_deg = roll_angle if t < 0.8 else 0.0
        squash_y = 0.75 + math.sin(t * math.pi) * 0.15
        shift_x = int(t * 18.0)
        shift_y = int(math.sin(t * math.pi) * 8.0)

    elif anim_name == "hurt":
        # Visceral flinch recoil
        recoil = math.sin(t * math.pi)
        shift_x = int(-recoil * 12.0)
        shift_y = int(recoil * 3.0)
        rot_deg = -recoil * 12.0
        squash_x = 1.0 - recoil * 0.08
        squash_y = 1.0 + recoil * 0.05

    # Apply deformation transformations
    w_deform = max(10, int(target_w * squash_x))
    h_deform = max(10, int(target_h * squash_y))
    deformed = scaled_hero.resize((w_deform, h_deform), Image.Resampling.BILINEAR)

    if abs(rot_deg) > 0.5:
        deformed = deformed.rotate(rot_deg, resample=Image.Resampling.BILINEAR, expand=True)

    # Center deformed sprite on frame base
    paste_x = ox + shift_x + (target_w - deformed.width) // 2
    paste_y = oy + shift_y + (target_h - deformed.height)
    frame.paste(deformed, (paste_x, paste_y), deformed)

    # Add dynamic visual indicators (e.g. blade slash arc on attack impact)
    if anim_name == "attack_slash" and frame_idx == 2:
        draw = ImageDraw.Draw(frame)
        arc_x0 = paste_x + deformed.width // 2 - 10
        arc_y0 = paste_y + 20
        arc_x1 = arc_x0 + 70
        arc_y1 = arc_y0 + 75
        draw.arc([arc_x0, arc_y0, arc_x1, arc_y1], start=280, end=70, fill=(239, 68, 68, 220), width=6)
        draw.arc([arc_x0 + 2, arc_y0 + 2, arc_x1 - 2, arc_y1 - 2], start=280, end=70, fill=(254, 202, 202, 255), width=2)

    return frame


def synthesize_monster_frame(
    base_img: Image.Image,
    anim_name: str,
    frame_idx: int,
    total_frames: int,
    frame_size: Tuple[int, int] = (160, 160),
) -> Image.Image:
    """Synthesizes animation frames for feral monsters (quadruped/beast kinematics)."""
    fw, fh = frame_size
    frame = Image.new("RGBA", (fw, fh), (0, 0, 0, 0))

    aspect = base_img.width / base_img.height
    target_h = int(fh * 0.78)
    target_w = int(target_h * aspect)
    scaled_mob = base_img.resize((target_w, target_h), Image.Resampling.LANCZOS)

    t = frame_idx / max(1, total_frames)
    phase = t * math.pi * 2.0

    ox = (fw - target_w) // 2
    oy = fh - target_h - 10

    squash_x = 1.0
    squash_y = 1.0
    rot_deg = 0.0
    shift_x = 0
    shift_y = 0

    if anim_name == "idle":
        # Low prowl snarl, chest heaving
        pant = math.sin(phase)
        squash_y = 1.0 + pant * 0.03
        squash_x = 1.0 - pant * 0.02
        shift_y = int(pant * 2.0)

    elif anim_name == "run":
        # Quadruped predatory gallop:
        # Compression -> Extension -> Contact
        gallop = math.sin(phase)
        shift_y = int(math.cos(phase * 2.0) * 4.5)
        rot_deg = gallop * 5.0
        squash_x = 1.0 + abs(gallop) * 0.07
        squash_y = 1.0 - abs(gallop) * 0.05
        shift_x = int(gallop * 4.0)

    elif anim_name == "attack":
        # Ferocious pounce & bite:
        # Frame 0: Crouch windup
        # Frame 1: Coiling back
        # Frame 2: Explosive forward leap / bite (Hit Frame)
        # Frame 3: Landing impact
        # Frame 4: Recover to stance
        if frame_idx == 0:
            rot_deg = -6.0
            shift_x = -5
            squash_y = 0.90
        elif frame_idx == 1:
            rot_deg = -12.0
            shift_x = -10
            squash_y = 0.85
        elif frame_idx == 2:  # BITE / CLAW IMPACT
            rot_deg = 15.0
            shift_x = 18
            shift_y = -6
            squash_x = 1.15
        elif frame_idx == 3:
            rot_deg = 5.0
            shift_x = 10
            shift_y = 4
            squash_y = 0.92
        else:
            rot_deg = 0.0
            shift_x = 2

    elif anim_name == "hurt":
        recoil = math.sin(t * math.pi)
        shift_x = int(-recoil * 14.0)
        shift_y = int(recoil * 4.0)
        rot_deg = -recoil * 10.0

    elif anim_name == "death":
        # Collapse and dissolve into ground
        squash_y = max(0.2, 1.0 - t * 0.75)
        shift_y = int(t * 18.0)
        rot_deg = t * 15.0

    w_deform = max(10, int(target_w * squash_x))
    h_deform = max(10, int(target_h * squash_y))
    deformed = scaled_mob.resize((w_deform, h_deform), Image.Resampling.BILINEAR)

    if abs(rot_deg) > 0.5:
        deformed = deformed.rotate(rot_deg, resample=Image.Resampling.BILINEAR, expand=True)

    paste_x = ox + shift_x + (target_w - deformed.width) // 2
    paste_y = oy + shift_y + (target_h - deformed.height)

    # In death animation, fade alpha gradually
    if anim_name == "death":
        alpha = deformed.split()[-1]
        fade_factor = max(0.1, 1.0 - t * 0.8)
        alpha = alpha.point(lambda p: int(p * fade_factor))
        deformed.putalpha(alpha)

    frame.paste(deformed, (paste_x, paste_y), deformed)
    return frame


def build_atlas_and_manifest(
    entity_id: str,
    base_image_path: Path,
    anim_configs: Dict[str, Dict[str, Any]],
    output_dir: Path,
    frame_size: Tuple[int, int],
    is_hero: bool = True,
    generate_mipmaps: bool = False,
) -> Tuple[Path, Path, Path]:
    """Bakes multi-frame animation sequences into an atlas, normal map, and JSON manifest."""
    output_dir.mkdir(parents=True, exist_ok=True)

    if not base_image_path.exists():
        raise FileNotFoundError(f"Base sprite image not found: {base_image_path}")

    base_img = Image.open(base_image_path).convert("RGBA")
    fw, fh = frame_size

    # Calculate total frames to determine atlas dimensions
    total_frames_count = sum(cfg["frames"] for cfg in anim_configs.values())
    cols = min(8, total_frames_count)
    rows = math.ceil(total_frames_count / cols)

    atlas_w = cols * fw
    atlas_h = rows * fh
    atlas_img = Image.new("RGBA", (atlas_w, atlas_h), (0, 0, 0, 0))

    manifest_frames: Dict[str, Dict[str, Any]] = {}
    manifest_anims: Dict[str, Dict[str, Any]] = {}

    cur_idx = 0
    for anim_name, cfg in anim_configs.items():
        num_frames = cfg["frames"]
        frame_keys = []

        for f_idx in range(num_frames):
            col = cur_idx % cols
            row = cur_idx // cols
            x = col * fw
            y = row * fh

            if is_hero:
                frame_surf = synthesize_hero_frame(base_img, anim_name, f_idx, num_frames, frame_size)
            else:
                frame_surf = synthesize_monster_frame(base_img, anim_name, f_idx, num_frames, frame_size)

            atlas_img.paste(frame_surf, (x, y), frame_surf)

            f_key = f"{anim_name}_{f_idx}"
            frame_keys.append(f_key)
            manifest_frames[f_key] = {
                "x": x,
                "y": y,
                "w": fw,
                "h": fh,
                "pivot": [0.5, 0.90],
            }
            cur_idx += 1

        manifest_anims[anim_name] = {
            "frames": frame_keys,
            "fps": cfg["fps"],
            "loop": cfg.get("loop", True),
            "stride_length_world": cfg.get("stride_length_world", 2.8),
            "hit_frame": cfg.get("hit_frame", None),
            "footsteps": cfg.get("footsteps", []),
            "action_speed_scale": cfg.get("action_speed_scale", True),
        }

    # Save Atlas PNG
    atlas_path = output_dir / f"{entity_id}_anim_atlas.png"
    atlas_img.save(atlas_path, "PNG", optimize=True)

    # Generate and Save Normal Map
    normal_img = generate_sobel_normal_map(atlas_img, strength=2.2)
    normal_path = output_dir / f"{entity_id}_anim_atlas_normal.png"
    normal_img.save(normal_path, "PNG", optimize=True)

    # Save Manifest JSON
    manifest_data = {
        "entity_id": entity_id,
        "texture_width": atlas_w,
        "texture_height": atlas_h,
        "frame_width": fw,
        "frame_height": fh,
        "kinematics": {
            "base_move_speed": 4.2,
            "stride_length_world": 2.8,
            "foot_speed_sync_formula": "playback_fps = base_fps * (current_velocity / base_move_speed)",
            "turn_rate_rad_per_sec": 18.84,  # Smooth turn interpolation
        },
        "frames": manifest_frames,
        "animations": manifest_anims,
    }
    manifest_path = output_dir / f"{entity_id}_anim_manifest.json"
    if generate_mipmaps:
        from tools.asset_pipeline.mipmap_utils import generate_and_save_mipmaps
        mips = generate_and_save_mipmaps(atlas_path, output_dir)
        manifest_data["mipmaps"] = [p.name for p in mips]
    with open(manifest_path, "w", encoding="utf-8") as f:
        json.dump(manifest_data, f, indent=2, ensure_ascii=False)

    print(f"[+] Successfully baked {entity_id} animation atlas:")
    print(f"    - Albedo:   {atlas_path} ({atlas_w}x{atlas_h})")
    print(f"    - Normal:   {normal_path}")
    print(f"    - Manifest: {manifest_path} ({len(manifest_frames)} frames)")

    return atlas_path, normal_path, manifest_path


def calculate_stride_fps(velocity: float, base_speed: float = 4.2, base_fps: float = 12.0) -> float:
    """Calculates motion-matched frame rate to eliminate foot sliding."""
    if base_speed <= 0.0 or velocity <= 0.01:
        return 0.0
    return max(3.0, base_fps * (velocity / base_speed))


def calculate_action_duration(base_duration: float, action_speed: float = 1.0) -> float:
    """Calculates skill action duration inversely scaled by attack/cast speed."""
    speed = max(0.2, action_speed)
    return base_duration / speed


def resolve_heading_8dir(deg: float) -> str:
    """Maps 360-degree angle to 8-directional isometric heading."""
    norm_deg = (deg % 360.0 + 360.0) % 360.0
    if norm_deg >= 337.5 or norm_deg < 22.5:
        return "E"
    elif norm_deg < 67.5:
        return "SE"
    elif norm_deg < 112.5:
        return "S"
    elif norm_deg < 157.5:
        return "SW"
    elif norm_deg < 202.5:
        return "W"
    elif norm_deg < 247.5:
        return "NW"
    elif norm_deg < 292.5:
        return "N"
    return "NE"


def run_pipeline(generate_mipmaps: bool = False) -> None:
    """Executes animation asset generation for 5 Martial Character Classes and Core Monsters."""
    print("=== FREEEXILE 2.5D ISOMETRIC ANIMATION PIPELINE ===")
    chars_dir = WORKSPACE_ROOT / "client" / "webapp" / "assets" / "characters"
    mobs_dir = WORKSPACE_ROOT / "client" / "webapp" / "assets" / "monsters"

    char_cfg = {
        "idle": {"frames": 4, "fps": 6, "loop": True},
        "run": {"frames": 8, "fps": 12, "loop": True, "stride_length_world": 2.8, "footsteps": [0, 4]},
        "attack_slash": {"frames": 6, "fps": 15, "loop": False, "hit_frame": 2, "action_speed_scale": True},
        "skill_whirlwind": {"frames": 6, "fps": 18, "loop": True, "action_speed_scale": True},
        "dodge": {"frames": 6, "fps": 20, "loop": False},
        "hurt": {"frames": 3, "fps": 12, "loop": False},
    }

    char_classes = [
        "char_sword_master", "char_sword_maiden", "char_feral_berserker",
        "char_wild_archer", "char_glacial_lancer", "char_shadow_assassin",
    ]
    for cid in char_classes:
        c_base = chars_dir / f"{cid}.png"
        if not c_base.exists() and cid == "char_sword_master":
            c_base = chars_dir / "savage_primal_exile.png"
        if c_base.exists():
            build_atlas_and_manifest(cid, c_base, char_cfg, ANIMATIONS_OUT_DIR, (160, 192), is_hero=True, generate_mipmaps=generate_mipmaps)
            if cid == "char_sword_master":
                build_atlas_and_manifest("hero", c_base, char_cfg, ANIMATIONS_OUT_DIR, (160, 192), is_hero=True, generate_mipmaps=generate_mipmaps)

    mob_cfg = {
        "idle": {"frames": 4, "fps": 5, "loop": True},
        "run": {"frames": 6, "fps": 10, "loop": True, "stride_length_world": 2.2, "footsteps": [1, 4]},
        "attack": {"frames": 5, "fps": 12, "loop": False, "hit_frame": 2, "action_speed_scale": True},
        "hurt": {"frames": 3, "fps": 12, "loop": False},
        "death": {"frames": 5, "fps": 10, "loop": False},
    }
    for mid in ["mob_feral_hellhound", "mob_skeleton_warrior"]:
        m_base = mobs_dir / f"{mid}.png"
        if m_base.exists():
            build_atlas_and_manifest(mid, m_base, mob_cfg, ANIMATIONS_OUT_DIR, (160, 160), is_hero=False, generate_mipmaps=generate_mipmaps)

    print("[*] All animation atlases baked and synchronized successfully!")


if __name__ == "__main__":
    parser = argparse.ArgumentParser(description="FreeExile Animation Pipeline Baker")
    parser.add_argument("--generate-all", action="store_true", help="Generate all animation atlases")
    parser.add_argument("--generate-mipmaps", action="store_true", help="Generate mipmap chains for atlases")
    args = parser.parse_args()
    run_pipeline(generate_mipmaps=args.generate_mipmaps)
