from __future__ import annotations

import tempfile
import unittest
import json
import subprocess
import sys
from pathlib import Path

import fitz
from PIL import Image

from omr_service.input_validation import (
    CORRUPT_FILE,
    open_pdf,
    partition_images,
    render_pdf_pages,
    validate_image,
)


ROOT = Path(__file__).resolve().parents[1]


class InputValidationTest(unittest.TestCase):
    def setUp(self) -> None:
        self.temporary = tempfile.TemporaryDirectory()
        self.root = Path(self.temporary.name)
        self.good_jpeg = self.root / "complete.jpg"
        Image.new("RGB", (120, 80), "white").save(self.good_jpeg, quality=90)

    def tearDown(self) -> None:
        self.temporary.cleanup()

    def assert_corrupt_image(self, path: Path) -> None:
        issue = validate_image(path)
        self.assertIsNotNone(issue)
        assert issue is not None
        self.assertEqual(CORRUPT_FILE, issue.code)
        self.assertEqual(path.name, issue.filename)
        self.assertIsNone(issue.page_number)

    def test_valid_image_fully_decodes(self) -> None:
        self.assertIsNone(validate_image(self.good_jpeg))

    def test_empty_and_random_jpegs_are_rejected(self) -> None:
        empty = self.root / "empty.jpg"
        empty.write_bytes(b"")
        random = self.root / "random.jpg"
        random.write_bytes(b"this is not a jpeg")
        self.assert_corrupt_image(empty)
        self.assert_corrupt_image(random)

    def test_generated_truncated_jpegs_are_rejected(self) -> None:
        complete = self.good_jpeg.read_bytes()
        for label, length in (("header", 32), ("quarter", len(complete) // 4), ("half", len(complete) // 2)):
            damaged = self.root / f"truncated-{label}.jpg"
            damaged.write_bytes(complete[:length])
            self.assert_corrupt_image(damaged)

    def test_mixed_images_preserve_valid_items_in_order(self) -> None:
        second_good = self.root / "complete-2.jpg"
        Image.new("RGB", (100, 60), "gray").save(second_good)
        damaged = self.root / "damaged.jpg"
        damaged.write_bytes(self.good_jpeg.read_bytes()[:40])
        valid, issues = partition_images([self.good_jpeg, damaged, second_good])
        self.assertEqual([self.good_jpeg, second_good], valid)
        self.assertEqual([damaged.name], [issue.filename for issue in issues])
        self.assertTrue(all(issue.code == CORRUPT_FILE for issue in issues))

    def make_pdf(self, path: Path, pages: int = 2) -> None:
        document = fitz.open()
        for number in range(1, pages + 1):
            page = document.new_page()
            page.insert_text((72, 72), f"Page {number}")
        document.save(path)
        document.close()

    def test_truncated_pdf_is_rejected_before_rendering(self) -> None:
        complete = self.root / "complete.pdf"
        self.make_pdf(complete)
        truncated = self.root / "truncated.pdf"
        truncated.write_bytes(complete.read_bytes()[:-20])
        document, issue = open_pdf(truncated)
        self.assertIsNone(document)
        self.assertIsNotNone(issue)
        assert issue is not None
        self.assertEqual(CORRUPT_FILE, issue.code)

    def test_valid_pdf_pages_are_preserved_in_order(self) -> None:
        complete = self.root / "complete.pdf"
        self.make_pdf(complete, pages=3)
        output = self.root / "rendered"
        output.mkdir()
        page_map, issues = render_pdf_pages(complete, output, 72)
        self.assertEqual([], issues)
        self.assertEqual([1, 2, 3], [row["page"] for row in page_map])
        self.assertTrue(all((output / row["file_id"]).is_file() for row in page_map))

    def test_corrupt_pdf_does_not_invoke_omrchecker(self) -> None:
        complete = self.root / "complete.pdf"
        self.make_pdf(complete)
        damaged = self.root / "damaged.pdf"
        damaged.write_bytes(complete.read_bytes()[:-20])
        omrchecker = self.root / "fake-omrchecker"
        omrchecker.mkdir()
        marker = self.root / "omr-was-invoked"
        (omrchecker / "main.py").write_text(
            "from pathlib import Path\n"
            f"Path({str(marker)!r}).write_text('invoked')\n"
        )
        output = self.root / "batch-output"
        completed = subprocess.run(
            [
                sys.executable,
                str(ROOT / "scripts" / "run_pdf_batch.py"),
                str(damaged),
                "--omrchecker",
                str(omrchecker),
                "--output",
                str(output),
            ],
            text=True,
            capture_output=True,
        )
        self.assertEqual(2, completed.returncode)
        self.assertFalse(marker.exists())
        summary = json.loads((output / "batch-summary.json").read_text())
        self.assertEqual("rejected", summary["status"])
        self.assertEqual(CORRUPT_FILE, summary["validation_issues"][0]["code"])


if __name__ == "__main__":
    unittest.main()
