#!/usr/bin/env python3
from __future__ import annotations

import argparse
import json
from pathlib import Path

from PIL import Image


def composite_rgb(path: Path) -> Image.Image:
    image = Image.open(path)
    image.load()
    if image.mode == 'RGBA':
        white = Image.new('RGBA', image.size, (255, 255, 255, 255))
        image = Image.alpha_composite(white, image)
    elif image.mode != 'RGB':
        image = image.convert('RGBA')
        white = Image.new('RGBA', image.size, (255, 255, 255, 255))
        image = Image.alpha_composite(white, image)
    return image.convert('RGB')


def page_metric(before_path: Path, after_path: Path, page: int) -> dict[str, object]:
    before = composite_rgb(before_path)
    after = composite_rgb(after_path)
    if before.size != after.size:
        raise RuntimeError(f'page {page}: image-size mismatch {before.size} != {after.size}')
    total_abs = 0
    changed_pixels = 0
    max_channel_delta = 0
    for a, b in zip(before.getdata(), after.getdata()):
        dr = abs(a[0] - b[0])
        dg = abs(a[1] - b[1])
        db = abs(a[2] - b[2])
        total_abs += dr + dg + db
        max_channel_delta = max(max_channel_delta, dr, dg, db)
        if dr or dg or db:
            changed_pixels += 1
    pixels = before.width * before.height
    mae = total_abs / (3 * 255 * pixels)
    return {
        'page': page,
        'width': before.width,
        'height': before.height,
        'pixel_similarity': 1.0 - mae,
        'mae': mae,
        'changed_pixels': changed_pixels,
        'changed_pixel_fraction': changed_pixels / pixels,
        'max_channel_delta': max_channel_delta,
    }


def main() -> int:
    ap = argparse.ArgumentParser(description='Measure CV-METRIC-001 page pixel similarity for semantic before/after screenshots.')
    ap.add_argument('evidence_dir', type=Path)
    ap.add_argument('--output', type=Path)
    args = ap.parse_args()
    baseline = args.evidence_dir / 'baseline'
    candidate = args.evidence_dir / 'candidate'
    before = sorted(baseline.glob('page-*.png'))
    after = sorted(candidate.glob('page-*.png'))
    if len(before) != len(after) or not before:
        raise RuntimeError(f'screenshot set mismatch: baseline={len(before)} candidate={len(after)}')
    pages = []
    for page, (before_path, after_path) in enumerate(zip(before, after), start=1):
        if before_path.name != after_path.name:
            raise RuntimeError(f'page filename mismatch: {before_path.name} != {after_path.name}')
        pages.append(page_metric(before_path, after_path, page))
    similarities = [float(row['pixel_similarity']) for row in pages]
    result = {
        'schema': 'pdf2html-pixel-similarity-evidence/v1',
        'status': 'POC-OBSERVED',
        'metric': 'CV-METRIC-001',
        'page_count': len(pages),
        'document_pixel_similarity_mean': sum(similarities) / len(similarities),
        'page_pixel_similarity_minimum': min(similarities),
        'pages_with_changed_pixels': sum(int(row['changed_pixels']) > 0 for row in pages),
        'total_changed_pixels': sum(int(row['changed_pixels']) for row in pages),
        'maximum_channel_delta': max(int(row['max_channel_delta']) for row in pages),
        'pages': pages,
    }
    output = args.output or args.evidence_dir / 'pixels.json'
    output.write_text(json.dumps(result, indent=2, sort_keys=False) + '\n', encoding='utf-8')
    print(json.dumps({key: result[key] for key in result if key != 'pages'}, sort_keys=True))
    return 0


if __name__ == '__main__':
    raise SystemExit(main())
