#!/usr/bin/env python3
"""Normalize a PDF, run MJ-RC5, retry bad orientations, and report timing."""

from __future__ import annotations

import argparse
import csv
import json
import shutil
import subprocess
import sys
import time
from pathlib import Path

ROOT = Path(__file__).resolve().parents[1]
PROFILE = ROOT / "profiles" / "mj-rc5"
sys.path.insert(0, str(ROOT))

from omr_service.input_validation import render_pdf_pages  # noqa: E402
from omr_service.omr_results import has_valid_controls, read_omr_results  # noqa: E402
from omr_service.page_normalization import normalize_scanner_page  # noqa: E402


RESULT_FIELDS = [
    "source_page", "file_id", "orientation", "route", "control_left", "control_right",
    "player_id_1", "player_id_2", "player_id_3", "round_tens", "round_ones",
    "game_1", "game_2", "game_3", "game_4",
]


def prepare_input(source_dir: Path, page_map: list[dict[str, object]], destination: Path, angle: int) -> float:
    started = time.perf_counter()
    destination.mkdir(parents=True)
    for name in ("template.json", "config.json", "omr_marker.jpg"):
        shutil.copy2(PROFILE / name, destination / name)
    for item in page_map:
        filename = str(item["file_id"])
        normalize_scanner_page(source_dir / filename, destination / filename, angle)
    return time.perf_counter() - started


def run_omr(omr_root: Path, input_dir: Path, output_dir: Path) -> tuple[subprocess.CompletedProcess[str], float]:
    output_dir.mkdir(parents=True)
    started = time.perf_counter()
    completed = subprocess.run(
        [sys.executable, str(omr_root / "main.py"), "-i", str(input_dir), "-o", str(output_dir)],
        cwd=omr_root,
        text=True,
        capture_output=True,
    )
    return completed, time.perf_counter() - started


def main() -> int:
    parser = argparse.ArgumentParser()
    parser.add_argument("pdf", type=Path)
    parser.add_argument("--omrchecker", required=True, type=Path)
    parser.add_argument("--output", required=True, type=Path)
    parser.add_argument("--dpi", type=int, default=300)
    args = parser.parse_args()

    pdf = args.pdf.resolve()
    output = args.output.resolve()
    omr_root = args.omrchecker.resolve()
    if not pdf.is_file() or not (omr_root / "main.py").is_file():
        raise FileNotFoundError(pdf if not pdf.is_file() else omr_root / "main.py")
    if output.exists():
        raise FileExistsError(f"Refusing to replace existing output: {output}")

    rendered_dir = output / "rendered-pages"
    rendered_dir.mkdir(parents=True)
    render_started = time.perf_counter()
    page_map, validation_issues = render_pdf_pages(pdf, rendered_dir, args.dpi)
    render_seconds = time.perf_counter() - render_started

    chosen: dict[str, dict[str, str]] = {}
    pending = list(page_map)
    normalization_seconds = 0.0
    omr_seconds = 0.0
    return_codes: list[int] = []
    console: list[str] = []
    for angle in (0, 180, 90, 270):
        if not pending:
            break
        stage = f"orientation-{angle:03d}"
        input_dir = output / "normalized" / stage
        stage_output = output / "omr-output" / stage
        normalization_seconds += prepare_input(rendered_dir, pending, input_dir, angle)
        completed, elapsed = run_omr(omr_root, input_dir, stage_output)
        omr_seconds += elapsed
        return_codes.append(completed.returncode)
        console.append(f"===== {stage} stdout =====\n{completed.stdout}")
        console.append(f"===== {stage} stderr =====\n{completed.stderr}")
        results = read_omr_results(stage_output)
        retry: list[dict[str, object]] = []
        for item in pending:
            filename = str(item["file_id"])
            row = results.get(filename)
            if row is not None:
                row["orientation"] = str(angle)
                chosen[filename] = row
            if not has_valid_controls(row):
                retry.append(item)
        pending = retry

    (output / "console.txt").write_text("\n".join(console))
    unresolved_ids = {str(item["file_id"]) for item in pending}
    page_lookup = {str(item["file_id"]): item["page"] for item in page_map}
    with (output / "batch-results.csv").open("w", newline="") as handle:
        writer = csv.DictWriter(handle, fieldnames=RESULT_FIELDS, extrasaction="ignore")
        writer.writeheader()
        for item in page_map:
            filename = str(item["file_id"])
            row = chosen.get(filename, {"file_id": filename, "route": "error"}).copy()
            row["source_page"] = page_lookup[filename]
            if filename in unresolved_ids:
                row["route"] = "error"
            writer.writerow(row)

    summary = {
        "source_pdf": str(pdf),
        "pages": len(page_map),
        "dpi": args.dpi,
        "render_seconds": round(render_seconds, 3),
        "normalization_seconds": round(normalization_seconds, 3),
        "omr_seconds": round(omr_seconds, 3),
        "total_seconds": round(render_seconds + normalization_seconds + omr_seconds, 3),
        "seconds_per_page": round((render_seconds + normalization_seconds + omr_seconds) / max(len(page_map), 1), 3),
        "orientation_retries": sum(row.get("orientation", "0") != "0" for row in chosen.values()),
        "unresolved_pages": [item["page"] for item in pending],
        "omr_return_codes": return_codes,
        "status": "rejected" if validation_issues and not page_map else "partial" if validation_issues or pending else "ok",
        "validation_issues": [issue.as_dict() for issue in validation_issues],
        "page_map": page_map,
    }
    (output / "batch-summary.json").write_text(json.dumps(summary, indent=2) + "\n")
    with (output / "page-map.csv").open("w", newline="") as handle:
        writer = csv.DictWriter(handle, fieldnames=["page", "file_id"])
        writer.writeheader()
        writer.writerows(page_map)
    print(json.dumps(summary, indent=2))
    if validation_issues or pending:
        return 2
    return next((code for code in return_codes if code), 0)


if __name__ == "__main__":
    raise SystemExit(main())
