"""
Empirical Adversarial Stress Test: UV Coordinate Geometry & Texture Clamping
Audits all generated Cocos 3.8.x manifests across VFX, Monsters, and Characters.
"""

from __future__ import annotations

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

REPO_ROOT = Path(__file__).resolve().parent.parent

RESOURCES_DIR = REPO_ROOT / "client" / "cocos" / "assets" / "resources"
VFX_MANIFEST = RESOURCES_DIR / "vfx" / "savage_primal_skills_vfx_atlas.json"
MONSTERS_DIR = RESOURCES_DIR / "monsters" / "archetypes"
CHARACTERS_DIR = RESOURCES_DIR / "characters"

EPSILON = 1e-6


def to_float32(val: float) -> float:
    """Converts a Python float to single-precision IEEE 754 float."""
    return struct.unpack("f", struct.pack("f", float(val)))[0]


def is_power_of_two(n: int) -> bool:
    """Returns True if n is a positive power of two."""
    return n > 0 and (n & (n - 1)) == 0


def collect_target_manifests() -> List[Path]:
    """Collects all canonical generated manifests (VFX, 10 Monsters, 6 Characters)."""
    manifests: List[Path] = []
    if VFX_MANIFEST.exists():
        manifests.append(VFX_MANIFEST)

    monster_manifests = sorted(MONSTERS_DIR.glob("*/*_anim_manifest.json"))
    manifests.extend(monster_manifests)

    character_manifests = sorted(CHARACTERS_DIR.glob("*/*_anim_manifest.json"))
    manifests.extend(character_manifests)

    return manifests


class FrameAuditMetric:
    """Holds empirical audit findings for a single manifest."""

    def __init__(self, manifest_path: Path):
        self.path = manifest_path
        self.total_frames = 0
        self.min_u = 1.0
        self.max_u = 0.0
        self.min_v = 1.0
        self.max_v = 0.0
        self.uv_out_of_bounds_violations = 0
        self.uv_degenerate_violations = 0
        self.area_violations = 0
        self.pivot_violations = 0
        self.sampler_clamp_violations = 0
        self.float32_precision_violations = 0
        self.pixel_sync_violations = 0
        self.pot_violations = 0
        self.texture_boundary_violations = 0


def audit_manifest(manifest_path: Path) -> FrameAuditMetric:
    """Executes exhaustive adversarial verification against a single manifest."""
    metric = FrameAuditMetric(manifest_path)
    with open(manifest_path, "r", encoding="utf-8") as f:
        data: Dict[str, Any] = json.load(f)

    tex_w = int(data.get("textureWidth", data.get("texture_width", 0)))
    tex_h = int(data.get("textureHeight", data.get("texture_height", 0)))

    if not (is_power_of_two(tex_w) and is_power_of_two(tex_h)):
        metric.pot_violations += 1

    def_pivot = data.get("defaultPivot") or data.get("pivot")
    if not def_pivot or len(def_pivot) != 2:
        metric.pivot_violations += 1
    elif abs(def_pivot[0] - 0.5) > EPSILON or abs(def_pivot[1] - 0.90) > EPSILON:
        metric.pivot_violations += 1

    # Extract all frame descriptors from clips and top-level frames
    frame_descriptors: List[Tuple[str, Dict[str, Any]]] = []
    if "clips" in data and isinstance(data["clips"], dict):
        for clip_name, clip_data in data["clips"].items():
            for idx, frame in enumerate(clip_data.get("frames", [])):
                frame_descriptors.append((f"{clip_name}[{idx}]", frame))

    if "frames" in data and isinstance(data["frames"], dict):
        for frame_id, frame in data["frames"].items():
            frame_descriptors.append((f"top_level[{frame_id}]", frame))

    for frame_label, frame in frame_descriptors:
        metric.total_frames += 1

        # 1. Area assertion: w * h > 0
        w = frame.get("w", 0)
        h = frame.get("h", 0)
        if w <= 0 or h <= 0 or (w * h) <= 0:
            metric.area_violations += 1

        x = frame.get("x", 0)
        y = frame.get("y", 0)
        if x < 0 or y < 0 or (x + w) > tex_w or (y + h) > tex_h:
            metric.texture_boundary_violations += 1

        # 2. Pivot assertion: exactly [0.5, 0.90]
        p = frame.get("pivot", def_pivot)
        if not p or len(p) != 2:
            metric.pivot_violations += 1
        elif abs(p[0] - 0.5) > EPSILON or abs(p[1] - 0.90) > EPSILON:
            metric.pivot_violations += 1

        # 3. UV coordinate geometry: 0.0 <= u0 < u1 <= 1.0 and 0.0 <= v0 < v1 <= 1.0
        uv = frame.get("uv")
        if not uv or len(uv) != 4:
            metric.uv_out_of_bounds_violations += 1
            continue

        u0, v0, u1, v1 = float(uv[0]), float(uv[1]), float(uv[2]), float(uv[3])
        metric.min_u = min(metric.min_u, u0)
        metric.max_u = max(metric.max_u, u1)
        metric.min_v = min(metric.min_v, v0)
        metric.max_v = max(metric.max_v, v1)

        if not (0.0 <= u0 <= 1.0 and 0.0 <= u1 <= 1.0 and 0.0 <= v0 <= 1.0 and 0.0 <= v1 <= 1.0):
            metric.uv_out_of_bounds_violations += 1

        if u1 <= u0 or v1 <= v0:
            metric.uv_degenerate_violations += 1

        # 4. Pixel-to-UV synchronization
        if tex_w > 0 and tex_h > 0:
            exp_u0 = x / float(tex_w)
            exp_v0 = y / float(tex_h)
            exp_u1 = (x + w) / float(tex_w)
            exp_v1 = (y + h) / float(tex_h)
            if (
                abs(u0 - exp_u0) > EPSILON
                or abs(v0 - exp_v0) > EPSILON
                or abs(u1 - exp_u1) > EPSILON
                or abs(v1 - exp_v1) > EPSILON
            ):
                metric.pixel_sync_violations += 1

        # 5. Cocos / GPU Sampler Clamping Simulation
        corners = [(u0, v0), (u1, v0), (u0, v1), (u1, v1)]
        for cu, cv in corners:
            clamped_u = max(0.0, min(1.0, cu))
            clamped_v = max(0.0, min(1.0, cv))
            if abs(cu - clamped_u) > EPSILON or abs(cv - clamped_v) > EPSILON:
                metric.sampler_clamp_violations += 1

        # Texel center sampling
        if tex_w > 0 and tex_h > 0:
            texel_center_u = u0 + 0.5 / float(tex_w)
            texel_center_v = v0 + 0.5 / float(tex_h)
            if not (0.0 <= texel_center_u <= 1.0 and 0.0 <= texel_center_v <= 1.0):
                metric.sampler_clamp_violations += 1

        # 6. IEEE 754 Float32 precision stability
        u0_f32, u1_f32 = to_float32(u0), to_float32(u1)
        v0_f32, v1_f32 = to_float32(v0), to_float32(v1)
        if not (0.0 <= u0_f32 < u1_f32 <= 1.0 and 0.0 <= v0_f32 < v1_f32 <= 1.0):
            metric.float32_precision_violations += 1

    return metric


def run_full_stress_audit() -> Tuple[List[FrameAuditMetric], int, Dict[str, Any]]:
    """Runs stress audit across all target manifests and aggregates results."""
    manifests = collect_target_manifests()
    metrics: List[FrameAuditMetric] = []

    total_frames = 0
    total_violations = 0
    global_min_u = 1.0
    global_max_u = 0.0
    global_min_v = 1.0
    global_max_v = 0.0

    for m_path in manifests:
        metric = audit_manifest(m_path)
        metrics.append(metric)

        total_frames += metric.total_frames
        manifest_violations = (
            metric.uv_out_of_bounds_violations
            + metric.uv_degenerate_violations
            + metric.area_violations
            + metric.pivot_violations
            + metric.sampler_clamp_violations
            + metric.float32_precision_violations
            + metric.pixel_sync_violations
            + metric.pot_violations
            + metric.texture_boundary_violations
        )
        total_violations += manifest_violations

        global_min_u = min(global_min_u, metric.min_u)
        global_max_u = max(global_max_u, metric.max_u)
        global_min_v = min(global_min_v, metric.min_v)
        global_max_v = max(global_max_v, metric.max_v)

    summary = {
        "manifest_count": len(manifests),
        "total_frames": total_frames,
        "total_violations": total_violations,
        "global_min_u": global_min_u,
        "global_max_u": global_max_u,
        "global_min_v": global_min_v,
        "global_max_v": global_max_v,
        "verdict": "APPROVE" if total_violations == 0 else "REQUEST_CHANGES",
    }
    return metrics, total_violations, summary


# =========================================================================
# Pytest Integration Suite
# =========================================================================


def test_manifest_inventory_completeness() -> None:
    """Asserts that all 17 canonical manifests exist (1 VFX, 10 Monsters, 6 Characters)."""
    manifests = collect_target_manifests()
    assert len(manifests) == 17, f"Expected 17 manifests, found {len(manifests)}"
    assert VFX_MANIFEST in manifests, "VFX manifest missing"
    assert len(list(MONSTERS_DIR.glob("*/*_anim_manifest.json"))) == 10, "Expected 10 monsters"
    assert len(list(CHARACTERS_DIR.glob("*/*_anim_manifest.json"))) == 6, "Expected 6 characters"


def test_uv_geometry_stress_and_sampler_clamping() -> None:
    """Audits every frame descriptor across all manifests against geometric and sampler invariants."""
    metrics, total_violations, summary = run_full_stress_audit()

    assert summary["total_frames"] >= 5500, (
        f"Expected > 5,500 audited frame descriptors, got {summary['total_frames']}"
    )
    assert total_violations == 0, (
        f"Adversarial stress test failed with {total_violations} violations: {summary}"
    )
    assert summary["verdict"] == "APPROVE"
    assert 0.0 <= summary["global_min_u"] <= summary["global_max_u"] <= 1.0
    assert 0.0 <= summary["global_min_v"] <= summary["global_max_v"] <= 1.0


if __name__ == "__main__":
    print("=" * 80)
    print("CHALLENGER EMPIRICAL ADVERSARIAL STRESS TEST: UV GEOMETRY & CLAMPING")
    print("=" * 80)

    metrics_list, violations, summ = run_full_stress_audit()

    print(f"Manifests Audited : {summ['manifest_count']}")
    print(f"Frames Audited    : {summ['total_frames']}")
    print(f"Global U Range    : [{summ['global_min_u']:.6f}, {summ['global_max_u']:.6f}]")
    print(f"Global V Range    : [{summ['global_min_v']:.6f}, {summ['global_max_v']:.6f}]")
    print(f"Total Violations  : {summ['total_violations']}")
    print(f"Empirical Verdict : {summ['verdict']}")
    print("-" * 80)

    for m in metrics_list:
        status = "PASS" if (
            m.uv_out_of_bounds_violations
            + m.uv_degenerate_violations
            + m.area_violations
            + m.pivot_violations
            + m.sampler_clamp_violations
        ) == 0 else "FAIL"
        print(f"[{status}] {m.path.name:<40} frames={m.total_frames:4d} U=[{m.min_u:.4f}, {m.max_u:.4f}] V=[{m.min_v:.4f}, {m.max_v:.4f}]")

    print("=" * 80)
    if summ["verdict"] == "APPROVE":
        print("VERDICT: APPROVE - All UV coordinates strictly conform to [0.0, 1.0] geometry.")
    else:
        print("VERDICT: REQUEST_CHANGES - Invariant violations detected.")
        exit(1)
