#!/usr/bin/env python3
"""Normalize a bounded scalar JSON grid; no inference or engine integration."""

import argparse
import binascii
import hashlib
import json
import math
import os
from pathlib import Path
import stat
import struct
import sys
import zlib

MAX_INPUT_BYTES = 16 * 1024 * 1024
MAX_WIDTH = 4096
MAX_HEIGHT = 4096
MAX_PIXELS = 1024 * 1024
PNG_SIGNATURE = b"\x89PNG\r\n\x1a\n"


class GridError(ValueError):
    """Invalid scalar input or normalization parameters."""


def validate_grid(grid):
    if not isinstance(grid, list) or not grid:
        raise GridError("grid must be a nonempty JSON array of rows")
    if len(grid) > MAX_HEIGHT:
        raise GridError(f"height exceeds {MAX_HEIGHT}")
    if not isinstance(grid[0], list) or not grid[0]:
        raise GridError("rows must be nonempty arrays")
    width = len(grid[0])
    if width > MAX_WIDTH:
        raise GridError(f"width exceeds {MAX_WIDTH}")
    if width * len(grid) > MAX_PIXELS:
        raise GridError(f"pixel count exceeds {MAX_PIXELS}")
    values = []
    for y, row in enumerate(grid):
        if not isinstance(row, list) or len(row) != width:
            raise GridError(f"row {y} is not a rectangular array of width {width}")
        for x, value in enumerate(row):
            if type(value) not in (int, float):
                raise GridError(f"cell [{y}][{x}] must be a number, not bool/null/text")
            try:
                value = float(value)
            except OverflowError as exc:
                raise GridError(f"cell [{y}][{x}] is outside finite binary64 range") from exc
            if not math.isfinite(value):
                raise GridError(f"cell [{y}][{x}] must be finite in binary64")
            values.append(value)
    return width, len(grid), values


def load_grid(path):
    if not Path(path).is_file():
        raise GridError("input must be a regular file")
    with Path(path).open("rb") as stream:
        if not stat.S_ISREG(os.fstat(stream.fileno()).st_mode):
            raise GridError("input must be a regular file")
        raw = stream.read(MAX_INPUT_BYTES + 1)
    if len(raw) > MAX_INPUT_BYTES:
        raise GridError(f"input exceeds {MAX_INPUT_BYTES} bytes")

    def reject_constant(token):
        raise GridError(f"non-standard JSON numeric token: {token}")

    try:
        grid = json.loads(raw.decode("utf-8"), parse_constant=reject_constant)
    except (UnicodeDecodeError, json.JSONDecodeError, RecursionError, ValueError) as exc:
        raise GridError(f"invalid JSON grid: {exc}") from exc
    return grid, hashlib.sha256(raw).hexdigest()


def percentile(sorted_values, q):
    """Linear order-statistic interpolation: h=(n-1)*q/100, zero-based."""
    h = (len(sorted_values) - 1) * (q / 100.0)
    lower = math.floor(h)
    upper = min(lower + 1, len(sorted_values) - 1)
    a, b = sorted_values[lower], sorted_values[upper]
    fraction = h - lower
    # Avoid overflow in b-a for opposite-sign extreme finite values.
    if a < 0 < b:
        return a * (1.0 - fraction) + b * fraction
    return a + (b - a) * fraction


def validate_options(near_is, mode, low_percentile, high_percentile):
    if near_is not in ("high", "low"):
        raise GridError("near-is must be high or low")
    if mode not in ("minmax", "percentile"):
        raise GridError("mode must be minmax or percentile")
    if mode == "minmax":
        if low_percentile is not None or high_percentile is not None:
            raise GridError("percentile bounds are only allowed in percentile mode")
        return
    if (type(low_percentile) not in (int, float) or
            type(high_percentile) not in (int, float)):
        raise GridError("percentile mode requires both percentile bounds")
    if not (0 <= low_percentile < high_percentile <= 100):
        raise GridError("require 0 <= low-percentile < high-percentile <= 100")


def normalize(grid, *, near_is, mode="minmax", low_percentile=None, high_percentile=None):
    validate_options(near_is, mode, low_percentile, high_percentile)
    width, height, values = validate_grid(grid)
    minimum, maximum = min(values), max(values)
    if minimum == maximum:
        raise GridError("constant grid has no usable normalization span")
    if mode == "percentile":
        ordered = sorted(values)
        low = percentile(ordered, low_percentile)
        high = percentile(ordered, high_percentile)
    else:
        low, high = minimum, maximum
    if not (math.isfinite(low) and math.isfinite(high) and low < high):
        raise GridError("selected bounds have no finite positive normalization span")

    span = high - low
    scale = max(abs(low), abs(high)) if not math.isfinite(span) else None
    normalized = []
    for value in values:
        if value <= low:
            t = 0.0
        elif value >= high:
            t = 1.0
        elif scale is None:
            t = (value - low) / span
        else:
            t = (value / scale - low / scale) / (high / scale - low / scale)
        t = min(1.0, max(0.0, t))
        normalized.append(t if near_is == "high" else 1.0 - t)
    manifest = {
        "schema_version": 1,
        "purpose": "scalar export experiment; not model inference or an engine test",
        "width": width, "height": height, "samples": len(values),
        "input_min": minimum, "input_max": maximum,
        "mode": mode, "near_is": near_is,
        "output_convention": "0 is far; maximum integer is near",
        "selected_low": low, "selected_high": high,
        "low_percentile": low_percentile, "high_percentile": high_percentile,
        "percentile_definition": "sorted ascending; h=(n-1)*(q/100); linear interpolation at h",
        "clipped_below": sum(value < low for value in values),
        "clipped_above": sum(value > high for value in values),
        "quantization": "clamp to [0,1], then floor(t*(2**bits-1)+0.5)",
        "number_representation": "Python binary64 float; no metric distance calibration",
    }
    return width, height, normalized, manifest


def quantize(values, bits):
    if type(bits) is not int or bits not in (8, 16):
        raise GridError("PNG bit depth must be 8 or 16")
    maximum = (1 << bits) - 1
    return [math.floor(min(1.0, max(0.0, value)) * maximum + 0.5) for value in values]


def _chunk(kind, data):
    return struct.pack(">I", len(data)) + kind + data + struct.pack(">I", binascii.crc32(kind + data) & 0xffffffff)


def _stored_zlib(raw):
    """Fixed stored DEFLATE blocks: reproducible without compressor heuristics."""
    result = bytearray(b"\x78\x01")
    for offset in range(0, len(raw), 65535):
        block = raw[offset:offset + 65535]
        result.append(1 if offset + len(block) == len(raw) else 0)
        result.extend(struct.pack("<HH", len(block), len(block) ^ 0xffff))
        result.extend(block)
    result.extend(struct.pack(">I", zlib.adler32(raw) & 0xffffffff))
    return bytes(result)


def grayscale_png(width, height, samples, bits):
    if (type(width) is not int or type(height) is not int or
            not 1 <= width <= MAX_WIDTH or not 1 <= height <= MAX_HEIGHT or
            width * height > MAX_PIXELS):
        raise GridError("invalid PNG dimensions")
    if type(bits) is not int or bits not in (8, 16) or len(samples) != width * height:
        raise GridError("invalid PNG bit depth or sample count")
    maximum = (1 << bits) - 1
    if any(type(value) is not int or not 0 <= value <= maximum for value in samples):
        raise GridError("invalid integer PNG sample")
    raw = bytearray()
    for y in range(height):
        raw.append(0)  # PNG filter 0: None.
        row = samples[y * width:(y + 1) * width]
        raw.extend(bytes(row) if bits == 8 else struct.pack(f">{width}H", *row))
    header = struct.pack(">IIBBBBB", width, height, bits, 0, 0, 0, 0)
    return PNG_SIGNATURE + _chunk(b"IHDR", header) + _chunk(b"IDAT", _stored_zlib(raw)) + _chunk(b"IEND", b"")


def export_grid(input_path, output_dir, *, near_is, mode="minmax", low_percentile=None, high_percentile=None):
    grid, source_hash = load_grid(input_path)
    width, height, normalized, manifest = normalize(
        grid, near_is=near_is, mode=mode,
        low_percentile=low_percentile, high_percentile=high_percentile)
    payloads = {}
    manifest["input_sha256"] = source_hash
    manifest["png"] = {}
    for bits in (8, 16):
        samples = quantize(normalized, bits)
        name = f"depth-{bits}.png"
        payloads[name] = grayscale_png(width, height, samples, bits)
        manifest["png"][str(bits)] = {
            "filename": name, "bit_depth": bits, "color_type": 0,
            "minimum": min(samples), "maximum": max(samples),
            "unique_levels": len(set(samples)),
            "sha256": hashlib.sha256(payloads[name]).hexdigest(),
        }
    payloads["manifest.json"] = (json.dumps(manifest, indent=2, sort_keys=True, allow_nan=False) + "\n").encode("utf-8")
    # Existing directories, files, and symlinks are refused. No overwrite flag.
    # The caller must supply an existing parent directory.
    output_dir = Path(output_dir)
    output_dir.mkdir(exist_ok=False, parents=False)
    for name, data in payloads.items():
        with (output_dir / name).open("xb") as stream:
            stream.write(data)
    return manifest


def main(argv=None):
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("input", help="UTF-8 JSON rectangular array of finite numbers")
    parser.add_argument("--out-dir", required=True, help="new directory; parent must already exist")
    parser.add_argument("--near-is", required=True, choices=("high", "low"))
    parser.add_argument("--mode", choices=("minmax", "percentile"), default="minmax")
    parser.add_argument("--low-percentile", type=float)
    parser.add_argument("--high-percentile", type=float)
    args = parser.parse_args(argv)
    try:
        manifest = export_grid(args.input, args.out_dir, near_is=args.near_is, mode=args.mode,
                               low_percentile=args.low_percentile, high_percentile=args.high_percentile)
    except (OSError, GridError) as exc:
        print(f"error: {exc}", file=sys.stderr)
        return 2
    print(json.dumps({"samples": manifest["samples"], "near_is": manifest["near_is"],
                      "unique_levels_8": manifest["png"]["8"]["unique_levels"],
                      "unique_levels_16": manifest["png"]["16"]["unique_levels"]}, sort_keys=True))
    return 0


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