#!/usr/bin/env python3
"""Send four sequential CosmosFull requests to an already running vLLM-Omni server.

Python standard library only. This helper does not deploy a server, load weights,
run a critic, select a winner, edit a film, or publish anything.
"""

from __future__ import annotations

import argparse
import datetime as dt
import hashlib
from http.client import HTTPException
import json
import struct
import sys
import time
import urllib.error
import urllib.parse
import urllib.request
import uuid
from pathlib import Path

SEEDS = (7, 21, 42, 99)
MODEL = "nvidia/Cosmos3-Super-Image2Video"
REVISION = "580f3f28e33ba93c8d464768876da8c322619aad"
IMAGE = "vllm/vllm-omni@sha256:970dee6658ea223f615b2438ce41e47f1d5322225482546e6e6bc5d8134f757c"
MAX_RESPONSE_BYTES = 128 * 1024 * 1024
SETTINGS = {
    "size": "720x1280", "num_frames": "121", "fps": "24",
    "num_inference_steps": "50", "guidance_scale": "6.0", "flow_shift": "5.0",
    "extra_params": json.dumps({"use_resolution_template": False,
                                "use_duration_template": False, "guardrails": True}),
}


def sha256(data: bytes) -> str:
    return hashlib.sha256(data).hexdigest()


def utc_now() -> str:
    return dt.datetime.now(dt.timezone.utc).isoformat()


def write_json(path: Path, value: dict) -> None:
    # Exclusive creation: never replace an earlier run or receipt.
    with path.open("x", encoding="utf-8") as handle:
        json.dump(value, handle, indent=2, ensure_ascii=False)
        handle.write("\n")


def no_duplicate_keys(pairs):
    result = {}
    for key, value in pairs:
        if key in result:
            raise ValueError(f"duplicate JSON key: {key}")
        result[key] = value
    return result


def read_inputs(keyframe: Path, prompt_file: Path) -> tuple[bytes, str, str, bytes]:
    image = keyframe.read_bytes()
    if len(image) < 33 or image[:8] != b"\x89PNG\r\n\x1a\n" or image[12:16] != b"IHDR":
        raise ValueError("keyframe must be a PNG file")
    if struct.unpack(">II", image[16:24]) != (720, 1280):
        raise ValueError("prepare a 720 x 1280 PNG keyframe; this helper does not resize it")
    raw = prompt_file.read_bytes()
    config = json.loads(raw, object_pairs_hook=no_duplicate_keys)
    if not isinstance(config, dict) or set(config) != {"prompt", "negative_prompt"}:
        raise ValueError("prompt JSON must contain exactly prompt and negative_prompt")
    prompt = config["prompt"]
    if isinstance(prompt, dict):
        if not isinstance(prompt.get("temporal_caption"), str) or not prompt["temporal_caption"].strip():
            raise ValueError("prompt object needs a nonempty temporal_caption")
        prompt = json.dumps(prompt, ensure_ascii=False, allow_nan=False)
    if not isinstance(prompt, str) or not prompt.strip():
        raise ValueError("prompt must be a nonempty string or temporal-caption object")
    negative = config["negative_prompt"]
    if not isinstance(negative, str):
        raise ValueError("negative_prompt must be a string")
    return image, prompt, negative, raw


def endpoint_url(endpoint: str) -> str:
    parts = urllib.parse.urlsplit(endpoint)
    if (parts.scheme not in {"http", "https"} or not parts.hostname or parts.username
            or parts.password or parts.query or parts.fragment or parts.path not in {"", "/"}):
        raise ValueError("endpoint must be an HTTP(S) server origin, without credentials or a path")
    return endpoint.rstrip("/") + "/v1/videos/sync"


class NoRedirect(urllib.request.HTTPRedirectHandler):
    def redirect_request(self, req, fp, code, msg, headers, newurl):
        # Do not silently send a keyframe to another server.
        return None


def multipart(fields: dict, image: bytes) -> tuple[bytes, str]:
    boundary = "cosmos-shot-" + uuid.uuid4().hex
    pieces = []
    for name, value in fields.items():
        pieces.append((f'--{boundary}\r\nContent-Disposition: form-data; name="{name}"'
                       f"\r\n\r\n{value}\r\n").encode("utf-8"))
    pieces.extend([
        (f'--{boundary}\r\nContent-Disposition: form-data; name="input_reference"; '
         'filename="input.png"\r\nContent-Type: image/png\r\n\r\n').encode("ascii"),
        image, f"\r\n--{boundary}--\r\n".encode("ascii"),
    ])
    return b"".join(pieces), f"multipart/form-data; boundary={boundary}"


def run(args) -> int:
    image, prompt, negative, raw_config = read_inputs(args.keyframe, args.prompt)
    url = endpoint_url(args.endpoint)
    if not 0 < args.timeout_seconds <= 1800:
        raise ValueError("timeout must be greater than zero and at most 1800 seconds")
    args.out.mkdir(parents=True, exist_ok=False)
    write_json(args.out / "run.json", {
        "schema_version": 1, "started_at_utc": utc_now(), "endpoint": url,
        "expected_server": {"model": MODEL, "revision": REVISION, "image": IMAGE,
                            "attested_by_helper": False},
        "settings": SETTINGS, "seeds": list(SEEDS),
        "keyframe_sha256": sha256(image), "keyframe_bytes": len(image),
        "prompt_file_sha256": sha256(raw_config), "prompt": prompt, "negative_prompt": negative,
        "validation": "HTTP status, MIME and MP4 header only; video decoding is not verified",
    })
    opener = urllib.request.build_opener(NoRedirect())
    completed = []
    for seed in SEEDS:
        target = args.out / f"take-s{seed}.mp4"
        receipt_path = args.out / f"take-s{seed}.json"
        if target.exists() or receipt_path.exists():
            raise FileExistsError("a take output or receipt already exists; refusing request")
        fields = {"prompt": prompt, "negative_prompt": negative, **SETTINGS, "seed": str(seed)}
        payload, content_type = multipart(fields, image)
        request = urllib.request.Request(url, data=payload, method="POST", headers={
            "Content-Type": content_type, "Accept": "video/mp4",
        })
        record = {"seed": seed, "started_at_utc": utc_now(), "request_fields": fields,
                  "request_body_sha256": sha256(payload)}
        started = time.monotonic()
        try:
            with opener.open(request, timeout=args.timeout_seconds) as response:
                status = response.status
                mime = response.headers.get_content_type()
                record.update({"http_status": status, "content_type": mime})
                if status != 200 or mime != "video/mp4":
                    raise ValueError("expected HTTP 200 with Content-Type video/mp4")
                video = response.read(MAX_RESPONSE_BYTES + 1)
                if len(video) > MAX_RESPONSE_BYTES:
                    raise ValueError("response exceeds the 128 MiB limit")
                if len(video) < 16 or video[4:8] != b"ftyp":
                    raise ValueError("response does not have an MP4 file-type header")
            with target.open("xb") as handle:
                handle.write(video)
            record.update({"ok": True, "output": target.name,
                           "output_sha256": sha256(video), "output_bytes": len(video)})
            completed.append(seed)
        except (OSError, ValueError, urllib.error.URLError, HTTPException) as error:
            record.update({"ok": False, "error_type": type(error).__name__})
            if isinstance(error, urllib.error.HTTPError):
                record["http_status"] = error.code
            # Do not print server bodies, which can contain private diagnostics.
        record.update({"finished_at_utc": utc_now(),
                       "request_elapsed_seconds": round(time.monotonic() - started, 6)})
        write_json(receipt_path, record)
        if not record["ok"]:
            write_json(args.out / "summary.json", {"ok": False, "completed_seeds": completed,
                                                   "failed_seed": seed, "finished_at_utc": utc_now()})
            print(f"Seed {seed} failed ({record['error_type']}); stopped without retries. See {receipt_path.name}.",
                  file=sys.stderr)
            return 1
        print(f"Seed {seed}: {target.name}, {len(video)} bytes, {record['request_elapsed_seconds']:.2f}s")
    write_json(args.out / "summary.json", {"ok": True, "completed_seeds": completed,
                                           "finished_at_utc": utc_now()})
    return 0


def main() -> int:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--keyframe", type=Path, required=True, help="your prepared 720 x 1280 PNG")
    parser.add_argument("--prompt", type=Path, required=True, help="JSON with prompt and negative_prompt")
    parser.add_argument("--out", type=Path, required=True, help="a new output directory; must not exist")
    parser.add_argument("--endpoint", default="http://127.0.0.1:8000", help="already running server origin")
    parser.add_argument("--timeout-seconds", type=float, default=900, help="socket timeout per request")
    try:
        return run(parser.parse_args())
    except (OSError, ValueError) as error:
        print(f"Stopped: {error}", file=sys.stderr)
        return 1
    except KeyboardInterrupt:
        print("Interrupted. No further requests will be sent; the server may still finish the current request.",
              file=sys.stderr)
        return 130


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