#!/usr/bin/env python3
"""Make the upscaler test's reference images and their 4x-reduced inputs.

For every sample with `use: upscale` in data/photo/samples.yaml:

  1. open the downloaded file and apply its EXIF orientation;
  2. convert to sRGB if it embeds a different ICC profile (two files are ProPhoto RGB), then drop
     the profile, so every tool sees plain sRGB pixels;
  3. crop to `reference.crop` when given (film borders on the two FSA transparencies);
  4. resize so the long side is `reference.long_side` (2000 px) with Pillow's LANCZOS filter,
     which averages over the whole footprint when shrinking, so sensor noise and film grain are
     mostly averaged away and the reference is sharp at the pixel level;
  5. trim the right and bottom edge (at most 3 px) so both sides are multiples of 4.

That is the reference ("ground truth") image. The low-resolution input is the reference reduced
by exactly 4x with Pillow's BICUBIC filter (Keys cubic, a = -0.5, kernel widened by the scale
factor so it anti-aliases, like MATLAB's imresize), rounded to 8 bits. Upscaling that input 4x
gives back the reference's exact size, so the two can be compared pixel by pixel.

Outputs go to data/private/photo-bench/<id>/{reference,lr}.png (git-ignored); the hashes, sizes
and library versions go to data/photo/bench/prepared.json so a re-run can be checked.

    data/private/photo-bench-venv/bin/python data/photo/bench/prepare.py
"""
from __future__ import annotations

import hashlib
import io
import json
import sys
from pathlib import Path

import PIL
import yaml
from PIL import Image, ImageCms, ImageOps

REPO = Path(__file__).resolve().parents[3]
sys.path.insert(0, str(REPO / "data" / "photo"))
from fetch_samples import load_samples, local_path  # noqa: E402

OUT = REPO / "data" / "private" / "photo-bench"
MANIFEST = REPO / "data" / "photo" / "bench" / "prepared.json"
SCALE = 4
Image.MAX_IMAGE_PIXELS = None  # the Library of Congress masters are up to 49 megapixels


def to_srgb(im: Image.Image) -> tuple[Image.Image, str | None]:
    icc = im.info.get("icc_profile")
    desc = None
    if icc:
        src = ImageCms.ImageCmsProfile(io.BytesIO(icc))
        desc = ImageCms.getProfileDescription(src).strip()
        if "srgb" not in desc.lower():
            im = ImageCms.profileToProfile(im.convert("RGB"), src, ImageCms.createProfile("sRGB"),
                                           renderingIntent=ImageCms.Intent.PERCEPTUAL, outputMode="RGB")
    return im.convert("RGB"), desc


def make_reference(sample: dict) -> tuple[Image.Image, dict]:
    im = Image.open(local_path(sample))
    im.load()
    im = ImageOps.exif_transpose(im)
    im, icc = to_srgb(im)
    steps = {"icc": icc, "converted_to_srgb": bool(icc and "srgb" not in icc.lower())}
    crop = (sample.get("reference") or {}).get("crop")
    if crop:
        im = im.crop(tuple(crop))
        steps["crop"] = crop
    long_side = sample["reference"]["long_side"]
    w, h = im.size
    if max(w, h) > long_side:
        s = long_side / max(w, h)
        im = im.resize((round(w * s), round(h * s)), Image.LANCZOS)
    w, h = im.size
    im = im.crop((0, 0, w - w % SCALE, h - h % SCALE))
    return im, steps


def sha256_png(path: Path) -> str:
    return hashlib.sha256(path.read_bytes()).hexdigest()


def main() -> int:
    samples = [s for s in load_samples() if s["use"] == "upscale"]
    records = {}
    for s in samples:
        ref, steps = make_reference(s)
        lr = ref.resize((ref.width // SCALE, ref.height // SCALE), Image.BICUBIC)
        d = OUT / s["id"]
        d.mkdir(parents=True, exist_ok=True)
        ref.save(d / "reference.png", optimize=False)
        lr.save(d / "lr.png", optimize=False)
        # Hash the decoded pixels, not the PNG bytes, so zlib versions don't matter.
        records[s["id"]] = {
            "reference_size": list(ref.size),
            "lr_size": list(lr.size),
            "reference_pixels_sha256": hashlib.sha256(ref.tobytes()).hexdigest(),
            "lr_pixels_sha256": hashlib.sha256(lr.tobytes()).hexdigest(),
            **steps,
        }
        print(f"{s['id']:30s} reference {ref.size[0]}x{ref.size[1]}  input {lr.size[0]}x{lr.size[1]}  {'(sRGB from ' + steps['icc'] + ')' if steps['converted_to_srgb'] else ''}")
    MANIFEST.write_text(json.dumps({
        "pillow": PIL.__version__,
        "scale": SCALE,
        "reference_filter": "PIL.Image.LANCZOS (long side 2000 px)",
        "downscale_filter": "PIL.Image.BICUBIC (exact 1/4 size)",
        "samples": records,
    }, indent=1) + "\n", encoding="utf-8")
    return 0


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