#!/usr/bin/env python3
"""Generate pilot keyframes through NVIDIA hosted image APIs."""

from __future__ import annotations

import argparse
import json
import sys
from pathlib import Path
from typing import Any

from PIL import Image

PROJECT_ROOT = Path(__file__).resolve().parents[1]
if str(PROJECT_ROOT) not in sys.path:
    sys.path.insert(0, str(PROJECT_ROOT))

from autoshorts.nvidia_image_client import NvidiaImageClient, NvidiaImageRequest

DEFAULT_REQUESTS = Path(
    "data/keyframe_requests/mechanism_pilots/hidden-money-systems-warum-ein-billiges-auto-teuer-sein-kann.json"
)
DEFAULT_REPORT = Path(
    "data/keyframe_generation_reports/mechanism_pilots/hidden-money-systems-warum-ein-billiges-auto-teuer-sein-kann.json"
)


def main() -> int:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--requests", type=Path, default=DEFAULT_REQUESTS)
    parser.add_argument("--report", type=Path, default=DEFAULT_REPORT)
    parser.add_argument(
        "--models",
        default="qwen-image,stable-diffusion-3.5-large,flux.1-dev",
        help="Comma-separated NVIDIA model order. The first working model is used per scene.",
    )
    parser.add_argument("--width", type=int, default=768)
    parser.add_argument("--height", type=int, default=1344)
    parser.add_argument("--steps", type=int, default=28)
    parser.add_argument("--cfg-scale", type=float, default=4.0)
    parser.add_argument("--seed-base", type=int, default=7000)
    parser.add_argument("--overwrite", action="store_true")
    args = parser.parse_args()

    package = json.loads(args.requests.read_text(encoding="utf-8"))
    model_order = [item.strip() for item in args.models.split(",") if item.strip()]
    client = NvidiaImageClient()
    results: list[dict[str, Any]] = []

    for index, item in enumerate(package["requests"], start=1):
        output_path = Path(item["target_path"])
        if output_path.exists() and not args.overwrite:
            results.append(
                {
                    "scene_id": item["scene_id"],
                    "status": "skipped",
                    "reason": "output_exists",
                    "output_path": str(output_path),
                }
            )
            continue

        prompt = _build_prompt(item)
        request = NvidiaImageRequest(
            prompt=prompt,
            output_path=output_path,
            width=args.width,
            height=args.height,
            seed=args.seed_base + index,
            steps=args.steps,
            cfg_scale=args.cfg_scale,
        )
        result = client.generate(request, model_order=model_order)
        row: dict[str, Any] = {
            "scene_id": item["scene_id"],
            "status": result.status,
            "model": result.model,
            "endpoint": result.endpoint,
            "http_status": result.http_status,
            "output_path": str(result.output_path) if result.output_path else str(output_path),
            "bytes_written": result.bytes_written,
        }
        if result.error:
            row["error"] = result.error
        if result.status == "success":
            required = item.get("required_format", {})
            row["image_check"] = _inspect_image(
                output_path,
                target_size=(int(required.get("width", args.width)), int(required.get("height", args.height))),
            )
        results.append(row)

    report = {
        "schema_version": "nvidia_keyframe_generation_report.v1",
        "source_requests": str(args.requests),
        "model_order": model_order,
        "requested_size": {"width": args.width, "height": args.height},
        "total": len(results),
        "succeeded": sum(1 for item in results if item["status"] == "success"),
        "skipped": sum(1 for item in results if item["status"] == "skipped"),
        "failed": sum(1 for item in results if item["status"] == "failed"),
        "results": results,
    }
    args.report.parent.mkdir(parents=True, exist_ok=True)
    args.report.write_text(json.dumps(report, indent=2, ensure_ascii=False) + "\n", encoding="utf-8")
    print(json.dumps(report, indent=2, ensure_ascii=False))
    return 0 if report["failed"] == 0 else 1


def _build_prompt(item: dict[str, Any]) -> str:
    prompt = item["generator_prompt"]
    negative = item.get("negative_requirements", "")
    scene_id = item.get("scene_id", "")
    if scene_id == "scene_02":
        prompt += " Slightly closer composition than the first shot, with subtle investigative mood and clean caption space."
    return (
        f"{prompt}\n\nStrict requirements: vertical 9:16 photorealistic keyframe, generic unbranded used compact car, "
        f"smooth badge-free grille and hood, blank neutral license plate area without blue EU strip, "
        f"no readable text, no brand logos, no manufacturer emblem, no watermark, no fake UI, no captions. "
        f"Avoid: {negative}"
    )


def _inspect_image(path: Path, *, target_size: tuple[int, int]) -> dict[str, Any]:
    with Image.open(path) as image:
        original = {"width": image.width, "height": image.height, "mode": image.mode, "format": image.format}
        converted = image.convert("RGB").resize(target_size, Image.Resampling.LANCZOS)
        converted.save(path, format="PNG")
        original["normalized_format"] = "PNG"
        original["normalized_width"] = target_size[0]
        original["normalized_height"] = target_size[1]
        return original


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