import json
from pathlib import Path
from PIL import Image
import numpy as np
import cv2

PROJECT_ROOT = Path("c:/Projects/FreeExile")
vfx_dir = PROJECT_ROOT / "client" / "cocos" / "assets" / "resources" / "vfx"
diffuse_path = vfx_dir / "savage_primal_skills_vfx_atlas.png"
normal_path = vfx_dir / "savage_primal_skills_vfx_atlas_normal.png"
manifest_path = vfx_dir / "savage_primal_skills_vfx_atlas.json"
astc_path = vfx_dir / "savage_primal_skills_vfx_atlas.astc"

print("=== 1. FILE EXISTENCE & SIZES ===")
for p in [diffuse_path, normal_path, manifest_path, astc_path]:
    print(f"{p.name}: {p.stat().st_size:,} bytes")

print("\n=== 2. DIFFUSE ATLAS INSPECTION ===")
diff_img = Image.open(diffuse_path)
diff_arr = np.array(diff_img)
print(f"Dimensions: {diff_img.size}, Mode: {diff_img.mode}")
assert diff_img.size == (2048, 2048)
alpha = diff_arr[:, :, 3]
opaque_pixels = int(np.sum(alpha > 0))
print(f"Opaque pixels (alpha > 0): {opaque_pixels:,} ({opaque_pixels / (2048*2048)*100:.2f}%)")

print("\n=== 3. SEQUENCE VARIATION & ENTROPY ===")
with open(manifest_path, "r", encoding="utf-8") as f:
    manifest = json.load(f)

unique_fingerprints = set()
clip_stats = []
for clip_name, clip_data in manifest.get("clips", {}).items():
    if "atlas_clip" in clip_data:
        continue  # skip alias
    frames = clip_data.get("frames", [])
    clip_pixels = 0
    for f_idx, f_info in enumerate(frames):
        x, y, w, h = f_info["x"], f_info["y"], f_info["w"], f_info["h"]
        sub = diff_arr[y:y+h, x:x+w]
        mean_rgb = tuple(sub[:, :, :3].mean(axis=(0, 1)).round(2))
        std_rgb = tuple(sub[:, :, :3].std(axis=(0, 1)).round(2))
        non_zero = int(np.sum(sub[:, :, 3] > 0))
        clip_pixels += non_zero
        fp = (mean_rgb, std_rgb, non_zero)
        unique_fingerprints.add(fp)
    clip_stats.append((clip_name, len(frames), clip_pixels))

print(f"Total clips: {len(manifest.get('clips', {}))}")
print(f"Total unique frame visual fingerprints: {len(unique_fingerprints)} across 72 frames (18 sequences * 4 frames)")
for cname, fcnt, px in clip_stats:
    print(f"  - {cname}: {fcnt} frames, {px:,} opaque pixels")

print("\n=== 4. NORMAL MAP MATHEMATICAL VALIDATION ===")
norm_img = Image.open(normal_path)
norm_arr = np.array(norm_img)
r_mean = float(norm_arr[:, :, 0].mean())
g_mean = float(norm_arr[:, :, 1].mean())
b_mean = float(norm_arr[:, :, 2].mean())
print(f"Normal channels: R(Nx) mean={r_mean:.2f}, G(Ny) mean={g_mean:.2f}, B(Nz) mean={b_mean:.2f}")
assert b_mean > 128.0

op_mask = alpha > 0
if np.any(op_mask):
    nx = (norm_arr[op_mask, 0].astype(float) / 255.0) * 2.0 - 1.0
    ny = (norm_arr[op_mask, 1].astype(float) / 255.0) * 2.0 - 1.0
    nz = (norm_arr[op_mask, 2].astype(float) / 255.0) * 2.0 - 1.0
    lens = np.sqrt(nx**2 + ny**2 + nz**2)
    print(f"Normal vector lengths on opaque texels: mean={lens.mean():.4f}, min={lens.min():.4f}, max={lens.max():.4f}")

print("\n=== 5. DILATION PADDING VERIFICATION ===")
kernel = cv2.getStructuringElement(cv2.MORPH_RECT, (3, 3))
dilated_alpha = cv2.dilate((alpha > 0).astype(np.uint8), kernel)
fringe = (dilated_alpha == 1) & (alpha == 0)
fringe_count = int(np.sum(fringe))
fringe_rgb = diff_arr[fringe, :3]
fringe_nonzero = int(np.sum(np.any(fringe_rgb > 0, axis=-1)))
print(f"Boundary fringe texels (alpha==0, adjacent to alpha>0): {fringe_count:,}")
print(f"Fringe texels with non-zero RGB (dilated colors): {fringe_nonzero:,} ({fringe_nonzero/fringe_count*100:.2f}%)")

print("\n=== 6. ASTC BINARY VALIDATION ===")
with open(astc_path, "rb") as f:
    astc_data = f.read()
header = astc_data[:16]
magic = header[:4]
bx, by, bz = header[4], header[5], header[6]
w = int.from_bytes(header[7:10], "little")
h = int.from_bytes(header[10:13], "little")
d = int.from_bytes(header[13:16], "little")
print(f"Magic: {magic.hex()} (expected: 13aba15c)")
print(f"Block size: {bx}x{by}x{bz}")
print(f"Dimensions: {w}x{h}x{d}")
expected_blocks = ((w + bx - 1) // bx) * ((h + by - 1) // by) * d
expected_size = 16 + expected_blocks * 16
print(f"Total file size: {len(astc_data):,} (expected: {expected_size:,})")
assert len(astc_data) == expected_size

# Check payload block entropy (are all blocks identical/zero or varied?)
payload = astc_data[16:]
block_set = set()
for i in range(0, min(len(payload), 16 * 1000), 16):
    block_set.add(payload[i:i+16])
print(f"Unique blocks sampled in first 1,000 blocks: {len(block_set)}")

print("\n=== ALL FORENSIC CHECKS EXECUTED SUCCESSFULLY ===")
