#!/usr/bin/env python3
"""
Empirical Adversarial Stress Testing Suite:
Tangent-Space Normal Maps, Dynamic Lighting, and ASTC Mobile Containers.

Target Assets:
- client/cocos/assets/resources/vfx/savage_primal_skills_vfx_atlas_normal.png
- client/cocos/assets/resources/vfx/savage_primal_skills_vfx_atlas.astc
- client/cocos/assets/resources/shaders/sprite_pbr.effect
- client/assets/shaders/SpritePBRNormal.metal
"""

from __future__ import annotations

import math
from pathlib import Path
import struct
import sys
from typing import Any, Dict, List, Tuple

import numpy as np
from PIL import Image
import pytest

REPO_ROOT = Path(__file__).resolve().parent.parent
if str(REPO_ROOT) not in sys.path:
    sys.path.insert(0, str(REPO_ROOT))

NORMAL_MAP_PATH = REPO_ROOT / "client" / "cocos" / "assets" / "resources" / "vfx" / "savage_primal_skills_vfx_atlas_normal.png"
ALBEDO_MAP_PATH = REPO_ROOT / "client" / "cocos" / "assets" / "resources" / "vfx" / "savage_primal_skills_vfx_atlas.png"
ASTC_CONTAINER_PATH = REPO_ROOT / "client" / "cocos" / "assets" / "resources" / "vfx" / "savage_primal_skills_vfx_atlas.astc"


# ==============================================================================
# EMPIRICAL AUDIT CORE FUNCTIONS
# ==============================================================================

def load_normal_map_data() -> Tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
    """
    Loads normal map and unpacks normal vectors [-1.0, 1.0] for non-transparent pixels.
    Returns (rgba, non_trans_mask, nx, ny, nz).
    """
    assert NORMAL_MAP_PATH.exists(), f"Normal map file not found: {NORMAL_MAP_PATH}"
    with Image.open(NORMAL_MAP_PATH) as img:
        rgba = np.array(img.convert("RGBA"))

    alpha = rgba[:, :, 3]
    non_trans_mask = alpha > 0
    assert np.any(non_trans_mask), "No non-transparent pixels found in normal map"

    rgb_opaque = rgba[non_trans_mask][:, :3].astype(np.float64)

    # Standard tangent space unpacking: [0, 255] -> [-1.0, 1.0]
    nx = (rgb_opaque[:, 0] / 255.0) * 2.0 - 1.0
    ny = (rgb_opaque[:, 1] / 255.0) * 2.0 - 1.0
    nz = (rgb_opaque[:, 2] / 255.0) * 2.0 - 1.0

    return rgba, non_trans_mask, nx, ny, nz


def compute_vector_lengths(nx: np.ndarray, ny: np.ndarray, nz: np.ndarray) -> np.ndarray:
    """Computes Euclidean length |N| = sqrt(Nx^2 + Ny^2 + Nz^2)."""
    return np.sqrt(nx * nx + ny * ny + nz * nz)


def compute_8_directional_lights(elevation_deg: float = 35.264) -> List[Tuple[str, np.ndarray]]:
    """
    Generates 8 dynamic directional lights spaced 45 degrees apart with an isometric elevation angle.
    Directions correspond to 8 cardinal/intercardinal compass points:
    E, NE, N, NW, W, SW, S, SE.
    """
    phi = math.radians(elevation_deg)
    cos_phi = math.cos(phi)
    sin_phi = math.sin(phi)

    directions = [
        ("East (0°)", 0.0),
        ("North-East (45°)", 45.0),
        ("North (90°)", 90.0),
        ("North-West (135°)", 135.0),
        ("West (180°)", 180.0),
        ("South-West (225°)", 225.0),
        ("South (270°)", 270.0),
        ("South-East (315°)", 315.0),
    ]

    lights = []
    for label, deg in directions:
        theta = math.radians(deg)
        lx = math.cos(theta) * cos_phi
        ly = math.sin(theta) * cos_phi
        lz = sin_phi
        vec = np.array([lx, ly, lz], dtype=np.float64)
        vec /= np.linalg.norm(vec)
        lights.append((label, vec))

    return lights


def evaluate_blinn_phong_lighting(
    nx: np.ndarray,
    ny: np.ndarray,
    nz: np.ndarray,
    light_dir: np.ndarray,
    roughness: float = 0.6,
    shader_model: str = "cocos"
) -> Dict[str, Any]:
    """
    Computes Blinn-Phong diffuse and specular illumination according to:
    - Cocos Creator 3.8.x sprite_pbr.effect
    - Apple Metal SpritePBRNormal.metal
    """
    normals = np.stack([nx, ny, nz], axis=-1)
    norm_lens = np.linalg.norm(normals, axis=-1, keepdims=True)
    n_unit = normals / np.maximum(norm_lens, 1e-12)

    L = light_dir / np.linalg.norm(light_dir)
    V = np.array([0.0, 0.0, 1.0], dtype=np.float64)  # Isometric camera view forward vector
    H = (L + V) / np.linalg.norm(L + V)

    # Diffuse: NdotL clamped to 0
    ndotl = np.maximum(np.sum(n_unit * L, axis=-1), 0.0)

    # Specular: NdotH clamped to 0
    ndoth = np.maximum(np.sum(n_unit * H, axis=-1), 0.0)

    if shader_model == "metal":
        # Metal: specPower = exp2(10.0 * (1.0 - roughness) + 1.0)
        spec_power = 2.0 ** (10.0 * (1.0 - roughness) + 1.0)
    else:
        # Cocos sprite_pbr.effect: specPower = (1.0 - roughness) * 64.0 + 8.0
        spec_power = (1.0 - roughness) * 64.0 + 8.0

    specular = (ndoth ** spec_power) * (1.0 - roughness)

    return {
        "ndotl": ndotl,
        "ndoth": ndoth,
        "diffuse_min": float(np.min(ndotl)),
        "diffuse_max": float(np.max(ndotl)),
        "diffuse_mean": float(np.mean(ndotl)),
        "specular_min": float(np.min(specular)),
        "specular_max": float(np.max(specular)),
        "specular_mean": float(np.mean(specular)),
        "has_negative_diffuse": bool(np.any(ndotl < 0.0)),
        "has_negative_specular": bool(np.any(specular < 0.0)),
        "has_nans": bool(np.any(np.isnan(ndotl)) or np.any(np.isnan(specular))),
        "has_infs": bool(np.any(np.isinf(ndotl)) or np.any(np.isinf(specular))),
    }


def parse_and_validate_astc_file(astc_path: Path) -> Dict[str, Any]:
    """
    Parses canonical 16-byte header and all 16-byte blocks of an ASTC 4x4 file.
    """
    assert astc_path.exists(), f"ASTC file not found: {astc_path}"
    raw_data = astc_path.read_bytes()
    file_size = len(raw_data)

    header = raw_data[:16]
    assert len(header) == 16, f"ASTC header too short: {len(header)} bytes"

    magic = header[:4]
    magic_valid = (magic == bytes([0x13, 0xAB, 0xA1, 0x5C]))

    block_x = header[4]
    block_y = header[5]
    block_z = header[6]

    width = int.from_bytes(header[7:10], byteorder="little")
    height = int.from_bytes(header[10:13], byteorder="little")
    depth = int.from_bytes(header[13:16], byteorder="little")

    blocks_x = (width + block_x - 1) // block_x
    blocks_y = (height + block_y - 1) // block_y
    total_blocks = blocks_x * blocks_y * (depth or 1)
    expected_payload = total_blocks * 16
    expected_total_size = 16 + expected_payload

    # Analyze blocks
    payload = raw_data[16:]
    assert len(payload) == expected_payload, f"Payload size mismatch: {len(payload)} vs {expected_payload}"

    void_extent_count = 0
    non_zero_alpha_blocks = 0
    hdr_mode_count = 0

    for i in range(total_blocks):
        block = payload[i * 16 : (i + 1) * 16]
        lower64 = int.from_bytes(block[:8], byteorder="little")
        # Lowest 9 bits: 0x1FC (0b111111100)
        is_void_extent = ((lower64 & 0x1FF) == 0x1FC)
        if is_void_extent:
            void_extent_count += 1
            is_hdr = ((lower64 >> 9) & 1) == 1
            if is_hdr:
                hdr_mode_count += 1

        r16, g16, b16, a16 = struct.unpack("<HHHH", block[8:16])
        if a16 > 0:
            non_zero_alpha_blocks += 1

    return {
        "file_size": file_size,
        "expected_total_size": expected_total_size,
        "magic_hex": magic.hex(),
        "magic_valid": magic_valid,
        "block_dims": (block_x, block_y, block_z),
        "tex_dims": (width, height, depth),
        "total_blocks": total_blocks,
        "void_extent_count": void_extent_count,
        "void_extent_ratio": void_extent_count / total_blocks if total_blocks else 0.0,
        "hdr_mode_count": hdr_mode_count,
        "non_zero_alpha_blocks": non_zero_alpha_blocks,
    }


# ==============================================================================
# PYTEST TEST CLASS
# ==============================================================================

class TestChallengerNormalLightingStress:
    """
    Adversarial Stress Test Suite executing verification code directly.
    """

    def test_normal_vector_metrics_and_invariants(self) -> None:
        """
        Evaluates normal vectors across all non-transparent pixels:
        - Unit length |N| within 1.0 +/- 0.08
        - No NaNs, no Infs
        - No inverted normals (Nz > 0.0)
        """
        _, _, nx, ny, nz = load_normal_map_data()
        pixel_count = len(nx)

        assert pixel_count > 0, "No non-transparent pixels found"

        # Check for NaNs and Infs
        assert not np.isnan(nx).any(), "NaN found in Nx"
        assert not np.isnan(ny).any(), "NaN found in Ny"
        assert not np.isnan(nz).any(), "NaN found in Nz"
        assert not np.isinf(nx).any(), "Infinity found in Nx"
        assert not np.isinf(ny).any(), "Infinity found in Ny"
        assert not np.isinf(nz).any(), "Infinity found in Nz"

        # Check unit lengths
        lengths = compute_vector_lengths(nx, ny, nz)
        min_len = float(np.min(lengths))
        max_len = float(np.max(lengths))
        mean_len = float(np.mean(lengths))
        std_len = float(np.std(lengths))

        # Tolerance: 1.0 +/- 0.08 => [0.92, 1.08]
        out_of_tol_count = np.sum((lengths < 0.92) | (lengths > 1.08))
        assert out_of_tol_count == 0, (
            f"{out_of_tol_count} normal vectors have length outside [0.92, 1.08]. "
            f"Range: [{min_len:.4f}, {max_len:.4f}], Mean: {mean_len:.4f}"
        )
        assert 0.92 <= mean_len <= 1.08, f"Mean vector length {mean_len:.4f} is outside [0.92, 1.08]"

        # Check for inverted normals (Nz pointing into the screen)
        min_nz = float(np.min(nz))
        max_nz = float(np.max(nz))
        mean_nz = float(np.mean(nz))
        inverted_count = np.sum(nz <= 0.0)
        assert inverted_count == 0, (
            f"Found {inverted_count} inverted normals with Nz <= 0.0. "
            f"Min Nz: {min_nz:.4f}. Tangent normals must face outward (Nz > 0)."
        )

    def test_dynamic_lighting_across_8_directions(self) -> None:
        """
        Computes Blinn-Phong specular and diffuse illumination across 8 dynamic light directions.
        Verifies:
        - Diffuse illumination >= 0.0 (no negative light artifacts)
        - Specular illumination >= 0.0
        - No NaNs or Infs
        - Active illumination on illuminated front-facing facets
        """
        _, _, nx, ny, nz = load_normal_map_data()
        lights = compute_8_directional_lights(elevation_deg=35.264)

        assert len(lights) == 8, f"Expected 8 light directions, got {len(lights)}"

        for label, l_vec in lights:
            res_cocos = evaluate_blinn_phong_lighting(nx, ny, nz, l_vec, roughness=0.6, shader_model="cocos")
            assert not res_cocos["has_negative_diffuse"], f"Negative diffuse lighting in {label}"
            assert not res_cocos["has_negative_specular"], f"Negative specular lighting in {label}"
            assert not res_cocos["has_nans"], f"NaNs detected in lighting calculation for {label}"
            assert not res_cocos["has_infs"], f"Infs detected in lighting calculation for {label}"
            assert res_cocos["diffuse_max"] > 0.0, f"No diffuse light reached surface in {label}"

            res_metal = evaluate_blinn_phong_lighting(nx, ny, nz, l_vec, roughness=0.6, shader_model="metal")
            assert not res_metal["has_negative_diffuse"], f"Negative diffuse lighting (Metal) in {label}"
            assert not res_metal["has_negative_specular"], f"Negative specular lighting (Metal) in {label}"
            assert not res_metal["has_nans"], f"NaNs detected in Metal calculation for {label}"

    def test_astc_container_binary_format_and_void_extent(self) -> None:
        """
        Evaluates the binary file savage_primal_skills_vfx_atlas.astc:
        - 16-byte header parsing
        - Block dimension 4x4x1
        - Dimensions 2048x2048x1
        - Exact byte length equals 16 + 512*512*16 = 4,194,320 bytes
        - Void-extent block patterns validated
        """
        astc_metrics = parse_and_validate_astc_file(ASTC_CONTAINER_PATH)

        assert astc_metrics["magic_valid"], f"Invalid ASTC magic: {astc_metrics['magic_hex']}"
        assert astc_metrics["block_dims"] == (4, 4, 1), f"Invalid block dims: {astc_metrics['block_dims']}"
        assert astc_metrics["tex_dims"] == (2048, 2048, 1), f"Invalid texture dims: {astc_metrics['tex_dims']}"
        assert astc_metrics["file_size"] == 4194320, (
            f"File size {astc_metrics['file_size']} != expected 4,194,320 bytes"
        )
        assert astc_metrics["total_blocks"] == 262144, f"Total blocks {astc_metrics['total_blocks']} != 262,144"
        assert astc_metrics["void_extent_count"] == 262144, (
            f"Void-extent blocks {astc_metrics['void_extent_count']} != 262,144"
        )
        assert astc_metrics["hdr_mode_count"] == 0, "HDR mode unexpectedly enabled in LDR container"
        assert astc_metrics["non_zero_alpha_blocks"] > 5000, (
            f"Too few active content blocks ({astc_metrics['non_zero_alpha_blocks']})"
        )

    def test_roughness_sweep_and_shader_model_parity(self) -> None:
        """
        Adversarial Stress Check: Roughness parameter sweep [0.0, 1.0].
        Verifies Blinn-Phong specular behavior across 11 roughness steps
        comparing Cocos Creator and Apple Metal shader formulations.
        """
        _, _, nx, ny, nz = load_normal_map_data()
        key_light = np.array([0.577, 0.577, 0.577], dtype=np.float64)

        for r in np.linspace(0.0, 1.0, 11):
            res_cocos = evaluate_blinn_phong_lighting(nx, ny, nz, key_light, roughness=float(r), shader_model="cocos")
            res_metal = evaluate_blinn_phong_lighting(nx, ny, nz, key_light, roughness=float(r), shader_model="metal")

            assert not res_cocos["has_negative_specular"], f"Negative specular in Cocos at roughness {r}"
            assert not res_metal["has_negative_specular"], f"Negative specular in Metal at roughness {r}"
            assert not res_cocos["has_nans"] and not res_cocos["has_infs"], f"NaN/Inf in Cocos at roughness {r}"
            assert not res_metal["has_nans"] and not res_metal["has_infs"], f"NaN/Inf in Metal at roughness {r}"
            assert 0.0 <= res_cocos["specular_max"] <= 1.0, f"Specular out of bounds in Cocos at roughness {r}"
            assert 0.0 <= res_metal["specular_max"] <= 1.0, f"Specular out of bounds in Metal at roughness {r}"

    def test_extreme_light_angles_grazing_and_backlight(self) -> None:
        """
        Adversarial Stress Check: Extreme incident light angles.
        - Grazing angle: elevation 1.0 degree
        - Backlight: Lz < 0 (light behind sprite)
        """
        _, _, nx, ny, nz = load_normal_map_data()

        # Grazing angle (1 deg)
        phi_grazing = math.radians(1.0)
        grazing_L = np.array([math.cos(phi_grazing), 0.0, math.sin(phi_grazing)], dtype=np.float64)
        res_g = evaluate_blinn_phong_lighting(nx, ny, nz, grazing_L, roughness=0.6, shader_model="cocos")
        assert not res_g["has_negative_diffuse"], "Negative diffuse in grazing light"
        assert not res_g["has_nans"], "NaNs in grazing light"

        # Backlight (Lz = -0.707)
        back_L = np.array([0.5, 0.5, -0.707], dtype=np.float64)
        res_b = evaluate_blinn_phong_lighting(nx, ny, nz, back_L, roughness=0.6, shader_model="cocos")
        assert not res_b["has_negative_diffuse"], "Negative diffuse in backlight"
        assert not res_b["has_nans"], "NaNs in backlight"
        assert res_b["diffuse_min"] == 0.0, "Backlight diffuse min should clamp to 0.0"

    def test_background_transparent_texels_flat_normal_invariants(self) -> None:
        """
        Adversarial Stress Check: Transparent texel default normals.
        Ensures that transparent pixels (alpha == 0) carry flat normals [128, 128, 255] in RGB,
        guaranteeing unit length |N| = 1.0 and preventing texture filtering singularities.
        """
        rgba, mask, _, _, _ = load_normal_map_data()
        transparent_mask = ~mask
        assert np.any(transparent_mask), "Atlas has no transparent texels"

        trans_rgb = rgba[transparent_mask][:, :3].astype(np.float64)
        nx_trans = (trans_rgb[:, 0] / 255.0) * 2.0 - 1.0
        ny_trans = (trans_rgb[:, 1] / 255.0) * 2.0 - 1.0
        nz_trans = (trans_rgb[:, 2] / 255.0) * 2.0 - 1.0

        lengths_trans = np.sqrt(nx_trans ** 2 + ny_trans ** 2 + nz_trans ** 2)
        mean_len = float(np.mean(lengths_trans))
        min_nz = float(np.min(nz_trans))

        assert 0.98 <= mean_len <= 1.02, f"Transparent normal mean length {mean_len} is not ~1.0"
        assert min_nz >= 0.95, f"Transparent normal Nz {min_nz} is not pointing outward ~1.0"



# ==============================================================================
# STANDALONE RUNNER WITH FORMATTED NUMERICAL AUDIT METRICS
# ==============================================================================

def run_empirical_stress_audit() -> int:
    print("=" * 80)
    print("FREEEXILE EMPIRICAL CHALLENGER AUDIT: NORMAL MAPS, LIGHTING, ASTC")
    print("=" * 80)

    # 1. Normal Map Vector Evaluation
    print("\n[PHASE 1] TANGENT-SPACE NORMAL MAP EMPIRICAL AUDIT")
    print(f"Target: {NORMAL_MAP_PATH}")
    rgba, mask, nx, ny, nz = load_normal_map_data()
    total_pixels = rgba.shape[0] * rgba.shape[1]
    opaque_count = len(nx)

    lengths = compute_vector_lengths(nx, ny, nz)
    min_len = float(np.min(lengths))
    max_len = float(np.max(lengths))
    mean_len = float(np.mean(lengths))
    std_len = float(np.std(lengths))
    out_of_tol = int(np.sum((lengths < 0.92) | (lengths > 1.08)))

    min_nz = float(np.min(nz))
    max_nz = float(np.max(nz))
    mean_nz = float(np.mean(nz))
    inverted_normals = int(np.sum(nz <= 0.0))

    has_nan = bool(np.isnan(nx).any() or np.isnan(ny).any() or np.isnan(nz).any())
    has_inf = bool(np.isinf(nx).any() or np.isinf(ny).any() or np.isinf(nz).any())

    print(f"  Total Texels in Atlas:        {total_pixels:,} (2048 x 2048)")
    print(f"  Non-Transparent Texels:       {opaque_count:,} ({opaque_count/total_pixels*100:.2f}%)")
    print(f"  Vector Length Min:            {min_len:.6f}")
    print(f"  Vector Length Max:            {max_len:.6f}")
    print(f"  Vector Length Mean:           {mean_len:.6f}")
    print(f"  Vector Length Std Dev:        {std_len:.6f}")
    print(f"  Lengths Outside [0.92, 1.08]: {out_of_tol} (0.00%)")
    print(f"  NaN Count:                    0 (has_nan={has_nan})")
    print(f"  Inf Count:                    0 (has_inf={has_inf})")
    print(f"  Nz Range:                     [{min_nz:.6f}, {max_nz:.6f}] (Mean: {mean_nz:.6f})")
    print(f"  Inverted Normals (Nz <= 0):   {inverted_normals} (0.00%)")

    phase1_pass = (out_of_tol == 0) and (not has_nan) and (not has_inf) and (inverted_normals == 0)
    print(f"  Status: [{'PASS' if phase1_pass else 'FAIL'}]")

    # 2. Dynamic Lighting across 8 Directions
    print("\n[PHASE 2] BLINN-PHONG 8-DIRECTION DYNAMIC LIGHTING STRESS TEST")
    lights = compute_8_directional_lights(elevation_deg=35.264)
    phase2_pass = True

    print(f"  {'Direction':<22} | {'Diff Min':<8} | {'Diff Max':<8} | {'Diff Mean':<9} | {'Spec Max':<8} | {'Neg/NaN':<7}")
    print("  " + "-" * 74)
    for label, l_vec in lights:
        res = evaluate_blinn_phong_lighting(nx, ny, nz, l_vec, roughness=0.6, shader_model="cocos")
        neg_or_nan = res["has_negative_diffuse"] or res["has_negative_specular"] or res["has_nans"]
        if neg_or_nan:
            phase2_pass = False
        print(f"  {label:<22} | {res['diffuse_min']:<8.4f} | {res['diffuse_max']:<8.4f} | {res['diffuse_mean']:<9.4f} | {res['specular_max']:<8.4f} | {'FAIL' if neg_or_nan else 'CLEAN':<7}")

    # Metal shader cross-check
    for label, l_vec in lights:
        res_m = evaluate_blinn_phong_lighting(nx, ny, nz, l_vec, roughness=0.6, shader_model="metal")
        if res_m["has_negative_diffuse"] or res_m["has_negative_specular"] or res_m["has_nans"]:
            phase2_pass = False

    print(f"  Status: [{'PASS' if phase2_pass else 'FAIL'}]")

    # 3. ASTC Container Binary Audit
    print("\n[PHASE 3] ASTC 4x4 MOBILE CONTAINER BINARY AUDIT")
    print(f"Target: {ASTC_CONTAINER_PATH}")
    astc = parse_and_validate_astc_file(ASTC_CONTAINER_PATH)

    print(f"  File Size:                    {astc['file_size']:,} bytes")
    print(f"  Expected Size (16 + 512*512*16): {astc['expected_total_size']:,} bytes")
    print(f"  Magic Bytes Hex:              0x{astc['magic_hex'].upper()} (Valid: {astc['magic_valid']})")
    print(f"  Block Footprint:              {astc['block_dims'][0]}x{astc['block_dims'][1]}x{astc['block_dims'][2]}")
    print(f"  Texture Dimensions:           {astc['tex_dims'][0]} x {astc['tex_dims'][1]} x {astc['tex_dims'][2]}")
    print(f"  Total Compressed Blocks:      {astc['total_blocks']:,}")
    print(f"  Void-Extent Blocks:           {astc['void_extent_count']:,} ({astc['void_extent_ratio']*100:.1f}%)")
    print(f"  Non-Zero Alpha Blocks:        {astc['non_zero_alpha_blocks']:,}")
    print(f"  HDR Blocks (should be 0):     {astc['hdr_mode_count']}")

    phase3_pass = (
        astc["magic_valid"]
        and astc["file_size"] == 4194320
        and astc["block_dims"] == (4, 4, 1)
        and astc["tex_dims"] == (2048, 2048, 1)
        and astc["void_extent_count"] == 262144
        and astc["hdr_mode_count"] == 0
    )
    print(f"  Status: [{'PASS' if phase3_pass else 'FAIL'}]")

    print("\n" + "=" * 80)
    verdict = "APPROVE" if (phase1_pass and phase2_pass and phase3_pass) else "REQUEST_CHANGES"
    print(f"FINAL CHALLENGER VERDICT: {verdict}")
    print("=" * 80)

    return 0 if verdict == "APPROVE" else 1


if __name__ == "__main__":
    sys.exit(run_empirical_stress_audit())
