"""Stdlib tests. Temporary files stay inside this package and are removed."""

import binascii
import contextlib
import hashlib
import io
import json
from pathlib import Path
import struct
import subprocess
import sys
import tempfile
import unittest
from unittest.mock import patch
import zlib

import depth_normalize as d

BASE = Path(__file__).resolve().parent


def decode_png(blob):
    """Independent small decoder for this package's grayscale, filter-0 PNGs."""
    if blob[:8] != b"\x89PNG\r\n\x1a\n":
        raise AssertionError("signature")
    offset, chunks = 8, []
    while offset < len(blob):
        length = struct.unpack(">I", blob[offset:offset + 4])[0]
        kind = blob[offset + 4:offset + 8]
        data = blob[offset + 8:offset + 8 + length]
        crc = struct.unpack(">I", blob[offset + 8 + length:offset + 12 + length])[0]
        if crc != binascii.crc32(kind + data) & 0xffffffff:
            raise AssertionError("CRC")
        chunks.append((kind, data))
        offset += 12 + length
    if [kind for kind, _ in chunks] != [b"IHDR", b"IDAT", b"IEND"]:
        raise AssertionError("chunks")
    width, height, bits, color, comp, filt, interlace = struct.unpack(">IIBBBBB", chunks[0][1])
    if (color, comp, filt, interlace) != (0, 0, 0, 0):
        raise AssertionError("format")
    raw = zlib.decompress(chunks[1][1])
    stride = width * (bits // 8)
    if len(raw) != height * (stride + 1):
        raise AssertionError("scanline length")
    samples = []
    for y in range(height):
        start = y * (stride + 1)
        if raw[start] != 0:
            raise AssertionError("filter")
        row = raw[start + 1:start + 1 + stride]
        samples.extend(row if bits == 8 else struct.unpack(f">{width}H", row))
    return width, height, bits, samples


class NormalizationTests(unittest.TestCase):
    def norm(self, grid, **kwargs):
        return d.normalize(grid, near_is=kwargs.pop("near_is", "high"), **kwargs)

    def test_minmax(self):
        self.assertEqual(self.norm([[2, 4, 6]])[2], [0, 0.5, 1])

    def test_explicit_low_polarity(self):
        self.assertEqual(self.norm([[2, 4, 6]], near_is="low")[2], [1, 0.5, 0])

    def test_negative_values(self):
        self.assertEqual(self.norm([[-5, 0, 5]])[2], [0, 0.5, 1])

    def test_single_column(self):
        self.assertEqual(self.norm([[0], [1]])[:3], (1, 2, [0, 1]))

    def test_half_up_quantization(self):
        self.assertEqual(d.quantize([0, 0.5, 1], 8), [0, 128, 255])
        self.assertEqual(d.quantize([0, 0.5, 1], 16), [0, 32768, 65535])

    def test_quantize_clamps(self):
        self.assertEqual(d.quantize([-1, 2], 8), [0, 255])

    def test_ramp_levels(self):
        values = self.norm([list(range(1024))])[2]
        self.assertEqual(len(set(d.quantize(values, 8))), 256)
        self.assertEqual(len(set(d.quantize(values, 16))), 1024)

    def test_percentile_interpolation(self):
        self.assertEqual(d.percentile([0, 10, 20, 30], 25), 7.5)
        self.assertEqual(d.percentile([0, 10, 20, 30], 75), 22.5)

    def test_percentile_endpoints(self):
        self.assertEqual(d.percentile([0, 10, 20], 0), 0)
        self.assertEqual(d.percentile([0, 10, 20], 100), 20)

    def test_percentile_clipping(self):
        _, _, values, manifest = self.norm([[0, 10, 20, 30]], mode="percentile", low_percentile=25, high_percentile=75)
        self.assertEqual(values, [0, 1 / 6, 5 / 6, 1])
        self.assertEqual((manifest["clipped_below"], manifest["clipped_above"]), (1, 1))

    def test_equal_threshold_is_not_counted_as_clipped(self):
        manifest = self.norm([[0, 1, 2]], mode="percentile", low_percentile=0, high_percentile=100)[3]
        self.assertEqual((manifest["clipped_below"], manifest["clipped_above"]), (0, 0))

    def test_outlier(self):
        grid = [list(range(100)) + [10000]]
        a = self.norm(grid)
        b = self.norm(grid, mode="percentile", low_percentile=0, high_percentile=99)
        self.assertEqual(len(set(d.quantize(a[2][:-1], 8))), 4)
        self.assertEqual(len(set(d.quantize(b[2][:-1], 8))), 100)
        self.assertEqual(b[3]["selected_high"], 99)
        self.assertEqual(b[3]["clipped_above"], 1)

    def test_extreme_finite_span(self):
        maximum = sys.float_info.max
        self.assertEqual(self.norm([[-maximum, 0, maximum]])[2], [0, 0.5, 1])

    def test_extreme_percentile_interpolation(self):
        maximum = sys.float_info.max
        self.assertEqual(d.percentile([-maximum, maximum], 50), 0)

    def test_subnormal_span(self):
        self.assertEqual(self.norm([[0, 5e-324, 1e-323]])[2], [0, 0.5, 1])

    def test_constant_rejected(self):
        for grid in ([[7, 7]], [[1]], [[0, -0.0]], [[2**53, 2**53 + 1]]):
            with self.subTest(grid=grid), self.assertRaises(d.GridError):
                self.norm(grid)

    def test_collapsed_percentile_rejected(self):
        with self.assertRaises(d.GridError):
            self.norm([[0, 0, 0, 1]], mode="percentile", low_percentile=0, high_percentile=50)

    def test_bad_shapes_rejected(self):
        for grid in ([], [[]], {}, [0, 1], [[0, 1], [2]], [[0, 1], None]):
            with self.subTest(grid=grid), self.assertRaises(d.GridError):
                self.norm(grid)

    def test_invalid_cells_rejected(self):
        for value in (True, False, None, "1", {}, [], float("nan"), float("inf"), -float("inf"), 10**400):
            with self.subTest(value=repr(value)), self.assertRaises(d.GridError):
                self.norm([[0, value]])

    def test_invalid_options_rejected(self):
        for kwargs in ({"near_is": "unknown"}, {"mode": "unknown"}, {"low_percentile": 0},
                       {"mode": "percentile"}, {"mode": "percentile", "low_percentile": 50, "high_percentile": 50},
                       {"mode": "percentile", "low_percentile": -1, "high_percentile": 99},
                       {"mode": "percentile", "low_percentile": 1, "high_percentile": 101},
                       {"mode": "percentile", "low_percentile": float("nan"), "high_percentile": 99},
                       {"mode": "percentile", "low_percentile": 0, "high_percentile": float("inf")},
                       {"mode": "percentile", "low_percentile": True, "high_percentile": 99}):
            with self.subTest(kwargs=kwargs), self.assertRaises(d.GridError):
                self.norm([[0, 1]], **kwargs)

    def test_shape_limit_boundaries(self):
        with patch.multiple(d, MAX_WIDTH=2, MAX_HEIGHT=2, MAX_PIXELS=4):
            self.assertEqual(self.norm([[0, 1], [2, 3]])[:2], (2, 2))
            for grid in ([[0, 1, 2]], [[0], [1], [2]]):
                with self.subTest(grid=grid), self.assertRaises(d.GridError):
                    self.norm(grid)
        with patch.object(d, "MAX_PIXELS", 3), self.assertRaises(d.GridError):
            self.norm([[0, 1], [2, 3]])

    def test_real_width_and_height_limits(self):
        self.assertEqual(self.norm([list(range(d.MAX_WIDTH))])[0], d.MAX_WIDTH)
        self.assertEqual(self.norm([[i] for i in range(d.MAX_HEIGHT)])[1], d.MAX_HEIGHT)
        with self.assertRaises(d.GridError):
            self.norm([list(range(d.MAX_WIDTH + 1))])
        with self.assertRaises(d.GridError):
            self.norm([[i] for i in range(d.MAX_HEIGHT + 1)])


class PngTests(unittest.TestCase):
    def test_roundtrip_8(self):
        self.assertEqual(decode_png(d.grayscale_png(2, 2, [0, 1, 128, 255], 8)), (2, 2, 8, [0, 1, 128, 255]))

    def test_roundtrip_16_big_endian(self):
        self.assertEqual(decode_png(d.grayscale_png(2, 2, [0, 256, 32768, 65535], 16)), (2, 2, 16, [0, 256, 32768, 65535]))

    def test_stored_deflate_multiple_blocks(self):
        samples = [i % 65536 for i in range(65536)]
        self.assertEqual(decode_png(d.grayscale_png(256, 256, samples, 16))[3], samples)

    def test_byte_determinism(self):
        self.assertEqual(d.grayscale_png(2, 1, [0, 255], 8), d.grayscale_png(2, 1, [0, 255], 8))

    def test_bad_png_samples(self):
        for samples in ([], [256], [-1], [True], [0.5]):
            with self.subTest(samples=samples), self.assertRaises(d.GridError):
                d.grayscale_png(1, 1, samples, 8)

    def test_bad_bit_depth(self):
        for bits in (0, 1, 32, 8.0, True):
            with self.subTest(bits=bits), self.assertRaises(d.GridError):
                d.grayscale_png(1, 1, [0], bits)
            with self.subTest(bits=bits), self.assertRaises(d.GridError):
                d.quantize([0], bits)

    def test_bad_png_dimensions(self):
        for width, height in ((0, 1), (-1, 1), (1, 0), (True, 1), (4097, 1)):
            with self.subTest(size=(width, height)), self.assertRaises(d.GridError):
                d.grayscale_png(width, height, [], 8)


class FileAndCliTests(unittest.TestCase):
    def setUp(self):
        self.temporary = tempfile.TemporaryDirectory(prefix=".test-", dir=BASE)
        self.addCleanup(self.temporary.cleanup)
        self.root = Path(self.temporary.name)
        self.input = self.root / "input.json"
        self.input.write_text("[[0,1,2]]\n", encoding="utf-8")

    def test_manifest_hashes_and_decoded_levels(self):
        out = self.root / "output"
        manifest = d.export_grid(self.input, out, near_is="high")
        self.assertEqual(manifest["input_sha256"], hashlib.sha256(self.input.read_bytes()).hexdigest())
        self.assertEqual(json.loads((out / "manifest.json").read_text()), manifest)
        for bits in (8, 16):
            blob = (out / f"depth-{bits}.png").read_bytes()
            self.assertEqual(manifest["png"][str(bits)]["sha256"], hashlib.sha256(blob).hexdigest())
            self.assertEqual(decode_png(blob)[3], [0, 128, 255] if bits == 8 else [0, 32768, 65535])

    def test_existing_directory_refused_unchanged(self):
        out = self.root / "output"
        out.mkdir()
        marker = out / "keep.txt"
        marker.write_text("unchanged")
        with self.assertRaises(FileExistsError):
            d.export_grid(self.input, out, near_is="high")
        self.assertEqual([p.name for p in out.iterdir()], ["keep.txt"])
        self.assertEqual(marker.read_text(), "unchanged")

    def test_existing_file_refused(self):
        with self.assertRaises(FileExistsError):
            d.export_grid(self.input, self.input, near_is="high")
        self.assertEqual(self.input.read_text(), "[[0,1,2]]\n")

    def test_symlink_output_refused(self):
        out = self.root / "link"
        out.symlink_to(self.root / "missing", target_is_directory=True)
        with self.assertRaises(FileExistsError):
            d.export_grid(self.input, out, near_is="high")

    def test_missing_parent_refused(self):
        with self.assertRaises(FileNotFoundError):
            d.export_grid(self.input, self.root / "absent" / "output", near_is="high")

    def test_invalid_input_creates_no_output(self):
        self.input.write_text("[[1,1]]")
        out = self.root / "output"
        with self.assertRaises(d.GridError):
            d.export_grid(self.input, out, near_is="high")
        self.assertFalse(out.exists())

    def test_json_rejections(self):
        for raw in (b"[[0,NaN]]", b"[[0,Infinity]]", b"[[0,-Infinity]]", b"[[0,1e309]]", b"[[0,1]] {}", b"\xff", b"["):
            self.input.write_bytes(raw)
            with self.subTest(raw=raw), self.assertRaises(d.GridError):
                grid, _ = d.load_grid(self.input)
                d.validate_grid(grid)

    def test_input_byte_limit(self):
        raw = self.input.read_bytes()
        with patch.object(d, "MAX_INPUT_BYTES", len(raw)):
            d.load_grid(self.input)
        with patch.object(d, "MAX_INPUT_BYTES", len(raw) - 1), self.assertRaises(d.GridError):
            d.load_grid(self.input)

    def test_non_file_input(self):
        with self.assertRaises(d.GridError):
            d.load_grid(self.root)

    def test_cli_requires_explicit_polarity_and_output(self):
        for args in ([str(self.input)], [str(self.input), "--out-dir", str(self.root / "out")],
                     [str(self.input), "--near-is", "high"]):
            with self.subTest(args=args), contextlib.redirect_stderr(io.StringIO()), self.assertRaises(SystemExit) as raised:
                d.main(args)
            self.assertEqual(raised.exception.code, 2)

    def test_cli_subprocess_success_and_duplicate_failure(self):
        command = [sys.executable, "-B", str(BASE / "depth_normalize.py"), str(self.input), "--out-dir", str(self.root / "out"), "--near-is", "low"]
        first = subprocess.run(command, capture_output=True, text=True)
        self.assertEqual(first.returncode, 0, first.stderr)
        self.assertEqual(json.loads(first.stdout)["near_is"], "low")
        second = subprocess.run(command, capture_output=True, text=True)
        self.assertEqual(second.returncode, 2)
        self.assertIn("error:", second.stderr)

    def test_two_exports_are_identical(self):
        for name in ("a", "b"):
            d.export_grid(self.input, self.root / name, near_is="high")
        for name in ("depth-8.png", "depth-16.png", "manifest.json"):
            self.assertEqual((self.root / "a" / name).read_bytes(), (self.root / "b" / name).read_bytes())


if __name__ == "__main__":
    unittest.main(verbosity=2)
