"""
High-Speed TensorRT FP16 Upscale Pipeline for Historical Series
Architecture: PureVision (Denoise) + SPAN (Super-Resolution) End-to-End TensorRT Engine
Engineered for NVIDIA RTX 4090 + Pinned Host Memory + Hardware NVENC Async Stream
"""

import os
import sys
import time
import json
import queue
import threading
import subprocess
import cv2
import torch
import gdown
import numpy as np
import tensorrt as trt

# 6 remaining episodes to upscale on Runpod
EPISODES = [
    ("ep3", "1iVnVxnFHeXQTfNRDAmnD_HSuPLVZhhNm"),
    ("ep4", "1B13SYAm2IPdlRzguQTltpjzrDXzIYBez"),
    ("ep5", "1lOi5dwlHRpJxepDo7dWA481pmQM5KY0M"),
    ("ep6", "1dn5_yNfgJ2YakOASqh8tlWLLL_Q-8W5d"),
    ("ep7", "1LKhuEY-5wnniEyqqUSgre5wXa1PC28as"),
    ("ep8", "1b9TVdVsRUFc61btpJ9NDop1olWp_oDWv")
]

INPUT_DIR = "/workspace/inputs"
OUTPUT_DIR = "/workspace/outputs_4k"
ENGINE_PATH = "/workspace/purescale_4k.engine"
PROGRESS_FILE = "/workspace/progress.json"
LOG_FILE = "/workspace/trt_batch.log"

os.makedirs(INPUT_DIR, exist_ok=True)
os.makedirs(OUTPUT_DIR, exist_ok=True)

def log(msg):
    ts = time.strftime("[%Y-%m-%d %H:%M:%S]")
    line = f"{ts} {msg}"
    print(line, flush=True)
    with open(LOG_FILE, "a", encoding="utf-8") as f:
        f.write(line + "\n")

def write_progress(ep_name, frame_idx, total_frames, cur_fps, completed, ep_start):
    el = time.time() - ep_start
    pct = frame_idx / max(1, total_frames) * 100
    rem_sec = (total_frames - frame_idx) / max(0.1, cur_fps)
    data = {
        "current_episode": ep_name,
        "current_frame": frame_idx,
        "total_frames": total_frames,
        "percent": round(pct, 1),
        "speed_fps": round(cur_fps, 1),
        "completed": completed,
        "eta_minutes": round(rem_sec / 60, 1),
        "elapsed_minutes": round(el / 60, 1),
        "timestamp": time.time()
    }
    tmp_path = PROGRESS_FILE + ".tmp"
    with open(tmp_path, "w", encoding="utf-8") as f:
        json.dump(data, f, indent=2)
    os.replace(tmp_path, PROGRESS_FILE)

def process_all():
    log("=== KHỞI TẠO PIPELINE PUREVISION + SPAN TENSORRT FP16 4K ===")
    log("Hardware: NVIDIA GeForce RTX 4090")
    
    # 1. Load Compiled TensorRT Engine
    log(f"Loading TensorRT Engine: {ENGINE_PATH}")
    runtime = trt.Runtime(trt.Logger(trt.Logger.WARNING))
    with open(ENGINE_PATH, "rb") as f:
        engine = runtime.deserialize_cuda_engine(f.read())
    context = engine.create_execution_context()

    stream = torch.cuda.Stream()
    # Input: 720p (1, 3, 720, 1280) float32
    d_input = torch.empty(1, 3, 720, 1280, device="cuda", dtype=torch.float32).contiguous()
    # Output: 4K UHD (1, 3, 2160, 3840) float32 (fused bicubic in TRT)
    d_output = torch.empty(1, 3, 2160, 3840, device="cuda", dtype=torch.float32).contiguous()
    
    context.set_tensor_address("input", d_input.data_ptr())
    context.set_tensor_address("output", d_output.data_ptr())

    POOL_SIZE = 16
    pinned_pool = [torch.empty(2160, 3840, 3, dtype=torch.uint8, pin_memory=True) for _ in range(POOL_SIZE)]

    completed_episodes = []
    # Check already completed files
    for ep_name, _ in EPISODES:
        out_video = os.path.join(OUTPUT_DIR, f"{ep_name}_4k.mp4")
        if os.path.exists(out_video) and os.path.getsize(out_video) > 50 * 1024 * 1024:
            completed_episodes.append(ep_name)

    log(f"Episodes already completed: {completed_episodes}")

    for ep_name, drive_id in EPISODES:
        out_video = os.path.join(OUTPUT_DIR, f"{ep_name}_4k.mp4")
        if ep_name in completed_episodes:
            log(f"Bỏ qua {ep_name}: Đã có video 4K hoàn chỉnh ({os.path.getsize(out_video)/(1024*1024):.1f} MB).")
            continue

        in_video = os.path.join(INPUT_DIR, f"{ep_name}.mp4")
        if not os.path.exists(in_video):
            log(f"Đang tải {ep_name}.mp4 từ Google Drive (ID: {drive_id})...")
            url = f"https://drive.google.com/uc?id={drive_id}"
            gdown.download(url, in_video, quiet=False)

        if not os.path.exists(in_video):
            log(f"[-] LỖI: Không tìm thấy {in_video} sau khi tải!")
            continue

        # Extract audio stream
        audio_file = os.path.join(INPUT_DIR, f"{ep_name}_audio.aac")
        subprocess.run([
            "ffmpeg", "-y", "-i", in_video,
            "-vn", "-c:a", "copy", audio_file
        ], stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, check=True)

        cap = cv2.VideoCapture(in_video)
        fps = cap.get(cv2.CAP_PROP_FPS) or 24.0
        total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))

        # Hardware NVENC stream with linear timestamps (-bf 0 -avoid_negative_ts make_zero)
        writer_cmd = [
            "ffmpeg", "-y",
            "-f", "rawvideo",
            "-pix_fmt", "bgr24",
            "-s", "3840x2160",
            "-r", str(fps),
            "-i", "pipe:0",
            "-i", audio_file,
            "-c:v", "hevc_nvenc",
            "-preset", "p4",
            "-cq", "19",
            "-bf", "0",
            "-avoid_negative_ts", "make_zero",
            "-pix_fmt", "yuv420p",
            "-c:a", "copy",
            "-shortest",
            out_video
        ]
        proc = subprocess.Popen(writer_cmd, stdin=subprocess.PIPE, stderr=subprocess.DEVNULL)

        write_queue = queue.Queue(maxsize=POOL_SIZE)
        def writer_worker():
            while True:
                data = write_queue.get()
                if data is None:
                    break
                proc.stdin.write(data)
                write_queue.task_done()

        writer_thread = threading.Thread(target=writer_worker, daemon=True)
        writer_thread.start()

        log(f"=== BẮT ĐẦU XỬ LÝ {ep_name.upper()} ({total_frames} frames) ===")
        ep_start = time.time()
        frame_idx = 0

        while True:
            ret, frame = cap.read()
            if not ret:
                break
            frame_idx += 1

            # 1. Preprocess on GPU
            t_frame = torch.from_numpy(frame).to("cuda", non_blocking=True)
            t_rgb = t_frame[:, :, [2, 1, 0]].permute(2, 0, 1).unsqueeze(0).float().div_(255.0)
            d_input.copy_(t_rgb)

            # 2. End-to-End TensorRT FP16 Inference (PureVision Denoise + SPAN Super-Resolution + Bicubic 4K)
            with torch.cuda.stream(stream):
                context.execute_async_v3(stream.cuda_stream)
            stream.synchronize()

            # 3. Output is already 4K UHD in d_output
            out_uint8 = d_output.squeeze(0).clamp_(0.0, 1.0).mul_(255.0).to(torch.uint8)
            out_bgr = out_uint8[[2, 1, 0], :, :].permute(1, 2, 0).contiguous()

            # 4. Fast 1ms DMA copy to pinned host buffer
            p_buf = pinned_pool[frame_idx % POOL_SIZE]
            p_buf.copy_(out_bgr, non_blocking=False)

            # 5. Enqueue zero-copy memoryview for async FFmpeg NVENC writer
            write_queue.put(memoryview(p_buf.numpy()))

            if frame_idx % 100 == 0 or frame_idx == total_frames:
                el = time.time() - ep_start
                cur_fps = frame_idx / max(0.1, el)
                pct = frame_idx / total_frames * 100
                rem_sec = (total_frames - frame_idx) / max(0.1, cur_fps)
                log(f"[{ep_name}] {frame_idx}/{total_frames} ({pct:.1f}%) | {cur_fps:.1f} FPS | ETA: {rem_sec/60:.1f}m")
                write_progress(ep_name, frame_idx, total_frames, cur_fps, completed_episodes, ep_start)

        write_queue.put(None)
        writer_thread.join()
        proc.stdin.close()
        proc.wait()
        cap.release()

        ep_time = time.time() - ep_start
        out_mb = os.path.getsize(out_video) / (1024 * 1024)
        completed_episodes.append(ep_name)
        log(f"✓ HOÀN THÀNH {ep_name.upper()}! Dung lượng: {out_mb:.1f} MB trong {ep_time/60:.1f} phút ({total_frames/ep_time:.1f} FPS)")

        write_progress(ep_name, total_frames, total_frames, total_frames/ep_time, completed_episodes, ep_start)

        # Cleanup input files to maintain zero disk waste
        if os.path.exists(in_video): os.remove(in_video)
        if os.path.exists(audio_file): os.remove(audio_file)

    log("🎉 TẤT CẢ 6 TẬP ĐÃ ĐƯỢC XỬ LÝ 4K THÀNH CÔNG!")
    with open("/workspace/ALL_DONE", "w") as f:
        f.write(f"Completed at {time.strftime('%Y-%m-%d %H:%M:%S')}\n")

if __name__ == "__main__":
    process_all()
