import os, struct, json
from PIL import Image
import numpy as np

vfx_dir = 'client/cocos/assets/resources/vfx'
diff_png = os.path.join(vfx_dir, 'savage_primal_skills_vfx_atlas.png')
norm_png = os.path.join(vfx_dir, 'savage_primal_skills_vfx_atlas_normal.png')
astc_file = os.path.join(vfx_dir, 'savage_primal_skills_vfx_atlas.astc')
json_file = os.path.join(vfx_dir, 'savage_primal_skills_vfx_atlas.json')

print('--- ASSET SIZES ---')
for p in [diff_png, norm_png, astc_file, json_file]:
    print(f'{p}: size={os.path.getsize(p)} bytes')

# PNG checks
img_diff = Image.open(diff_png)
print(f'Diffuse size: {img_diff.size}, mode: {img_diff.mode}')
assert img_diff.size == (2048, 2048)

img_norm = Image.open(norm_png)
print(f'Normal size: {img_norm.size}, mode: {img_norm.mode}')
assert img_norm.size == (2048, 2048)

# Normal map channels (PIL loads as RGB)
arr_norm = np.array(img_norm)
mean_r = float(np.mean(arr_norm[:, :, 0]))
mean_g = float(np.mean(arr_norm[:, :, 1]))
mean_b = float(np.mean(arr_norm[:, :, 2]))
print(f'Normal channels mean: R={mean_r:.2f}, G={mean_g:.2f}, B={mean_b:.2f}')
assert mean_b > 128.0, f'Blue mean {mean_b} <= 128.0'

# ASTC checks
astc_size = os.path.getsize(astc_file)
with open(astc_file, 'rb') as f:
    hdr = f.read(16)
magic = hdr[:4]
block_x, block_y, block_z = hdr[4], hdr[5], hdr[6]
w = int.from_bytes(hdr[7:10], 'little')
h = int.from_bytes(hdr[10:13], 'little')
d = int.from_bytes(hdr[13:16], 'little')
print(f'ASTC Header: magic={magic.hex()}, block=({block_x}x{block_y}x{block_z}), dim=({w}x{h}x{d}), total_size={astc_size}')
assert magic == bytes([0x13, 0xab, 0xa1, 0x5c]), f'Invalid ASTC magic {magic.hex()}'
assert block_x == 4 and block_y == 4 and block_z == 1
assert w == 2048 and h == 2048 and d == 1
assert astc_size == 4194320, f'Expected 4194320 bytes, got {astc_size}'

# JSON Manifest checks
with open(json_file, 'r', encoding='utf-8') as f:
    manifest = json.load(f)
name = manifest.get('name')
clips_count = len(manifest.get('clips', {}))
frames_count = len(manifest.get('frames', {}))
default_pivot = manifest.get('defaultPivot')
print(f'Manifest name: {name}, clips={clips_count}, frames={frames_count}')
print(f'Default pivot: {default_pivot}')
assert default_pivot == [0.5, 0.90] or default_pivot == [0.5, 0.9]

# Check UV bounds in manifest
out_of_bounds = 0
total_uvs = 0
for f_id, f_data in manifest.get('frames', {}).items():
    uv = f_data.get('uv')
    if uv:
        total_uvs += 1
        for val in uv:
            if val < 0.0 or val > 1.0:
                out_of_bounds += 1
print(f'Manifest UV check: total={total_uvs}, out_of_bounds={out_of_bounds}')
assert out_of_bounds == 0

# Check skill coverage in manifest clips
required_skills = [
    1001, 1002, 1003, 1004, 1005, 1006, 1007, 1008, 1009, 1010, 1011, 1012,
    1099, 2001, 2002, 2004, 2005, 2006
]
clips = manifest.get('clips', {})
for sk in required_skills:
    assert str(sk) in clips or any(str(sk) in k for k in clips), f'Missing skill {sk} in manifest clips'

print(f'All {len(required_skills)} required skill/sigil IDs present in manifest!')
print('ALL ASSET FORENSIC CHECKS PASSED!')
