#!/usr/bin/env python3
"""
Grava HTML com Chrome DevTools Protocol Page.startScreencast.
Recebe frames JPEG do Chrome via WebSocket; monta MP4 com ffmpeg.
Muito mais rápido que page.screenshot loop (~30fps sustentável).
"""
import os, sys, time, base64, subprocess, argparse, threading
from pathlib import Path
from playwright.sync_api import sync_playwright

def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--html", required=True)
    ap.add_argument("--out", required=True)
    ap.add_argument("--width", type=int, default=1080)
    ap.add_argument("--height", type=int, default=1920)
    ap.add_argument("--fps", type=int, default=30)
    ap.add_argument("--duration-ms", type=int, default=50500)
    ap.add_argument("--frames-dir", default=None)
    ap.add_argument("--quality", type=int, default=90)
    args = ap.parse_args()

    html_path = Path(args.html).resolve()
    out_mp4 = Path(args.out).resolve()
    out_mp4.parent.mkdir(parents=True, exist_ok=True)
    frames_dir = Path(args.frames_dir) if args.frames_dir else out_mp4.parent / f"_frames_{out_mp4.stem}"
    frames_dir.mkdir(parents=True, exist_ok=True)
    for f in frames_dir.glob("*.jpg"): f.unlink()
    for f in frames_dir.glob("*.png"): f.unlink()

    duration_s = args.duration_ms / 1000.0
    max_frames = int(args.fps * duration_s * 1.4)  # margem
    print(f"[record] target {duration_s:.1f}s @ {args.fps}fps -> {args.width}x{args.height}")

    frames_captured = []  # (session_ms, path)
    stop_flag = threading.Event()

    with sync_playwright() as p:
        browser = p.chromium.launch(headless=True, args=[
            '--disable-web-security',
            '--allow-file-access-from-files',
            '--no-sandbox',
            '--disable-blink-features=AutomationControlled',
            '--font-render-hinting=none',
        ])
        ctx = browser.new_context(
            viewport={'width': args.width, 'height': args.height},
            device_scale_factor=1,
        )
        page = ctx.new_page()
        page.goto(f"file://{html_path}?record=1", wait_until="load")
        page.evaluate("() => document.fonts.ready")
        page.wait_for_load_state('networkidle', timeout=15000)
        time.sleep(2.0)

        cdp = ctx.new_cdp_session(page)

        # handler para receber frames
        def on_frame(evt):
            session_id = evt.get('sessionId')
            data = evt.get('data')
            ts = time.time()
            if data and not stop_flag.is_set():
                idx = len(frames_captured)
                if idx < max_frames:
                    fpath = frames_dir / f"frame_{idx:05d}.jpg"
                    fpath.write_bytes(base64.b64decode(data))
                    frames_captured.append((ts, fpath))
            # ack frame
            try:
                cdp.send('Page.screencastFrameAck', {'sessionId': session_id})
            except Exception:
                pass

        cdp.on('Page.screencastFrame', on_frame)

        # start screencast
        cdp.send('Page.startScreencast', {
            'format': 'jpeg',
            'quality': args.quality,
            'maxWidth': args.width,
            'maxHeight': args.height,
            'everyNthFrame': 1,
        })
        # dispara animação
        page.evaluate("() => window.__START && window.__START()")
        start_ts = time.time()

        # aguarda duração - usa page.wait_for_timeout que bombeia eventos
        end_ts = start_ts + duration_s + 0.3
        while time.time() < end_ts:
            page.wait_for_timeout(50)
        stop_flag.set()
        try:
            cdp.send('Page.stopScreencast', {})
        except Exception:
            pass
        # deixa últimos frames chegarem
        time.sleep(0.3)
        browser.close()

    n = len(frames_captured)
    if n == 0:
        print("[erro] nenhum frame capturado")
        sys.exit(1)

    # calcula duração real e fps efetivo
    first_ts = frames_captured[0][0]
    last_ts = frames_captured[-1][0]
    real_dur = last_ts - first_ts
    real_fps = (n - 1) / real_dur if real_dur > 0 else args.fps
    print(f"[record] {n} frames em {real_dur:.2f}s (real fps={real_fps:.1f})")

    # Estratégia: os frames chegam com timestamps não-uniformes.
    # Vamos gerar um concat list para ffmpeg com duração por frame.
    concat_file = frames_dir / "_concat.txt"
    with open(concat_file, "w") as f:
        for i in range(n):
            fpath = frames_captured[i][1]
            if i < n - 1:
                dur = frames_captured[i+1][0] - frames_captured[i][0]
            else:
                dur = 1.0 / args.fps
            dur = max(0.005, min(0.5, dur))
            f.write(f"file '{fpath.name}'\n")
            f.write(f"duration {dur:.5f}\n")
        # ffmpeg requer última entrada duplicada sem duration
        f.write(f"file '{frames_captured[-1][1].name}'\n")

    # limita duração final
    cmd = [
        'ffmpeg', '-y', '-f', 'concat', '-safe', '0',
        '-i', str(concat_file),
        '-t', str(duration_s),
        '-vf', f'fps={args.fps},format=yuv420p',
        '-c:v', 'libx264', '-preset', 'medium', '-crf', '18',
        '-movflags', '+faststart',
        str(out_mp4)
    ]
    print(f"[ffmpeg] concat -> mp4")
    r = subprocess.run(cmd, capture_output=True, text=True)
    if r.returncode != 0:
        print("STDERR:", r.stderr[-2000:])
        sys.exit(r.returncode)

    dur = subprocess.check_output(['ffprobe','-v','error','-show_entries','format=duration',
                                   '-of','default=noprint_wrappers=1:nokey=1', str(out_mp4)]).decode().strip()
    print(f"[done] {out_mp4} ({dur}s)")

if __name__ == '__main__':
    main()
