import os
import sys
import time
import json
import subprocess
import cv2
import torch
import gdown
from basicsr.archs.rrdbnet_arch import RRDBNet

EPISODES = [
    ("ep2", "1tQCOjzdZT9q1GHftgOcMknOvEvMW5K3n"),
    ("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"
LOG_FILE = "/workspace/batch.log"
PROGRESS_FILE = "/workspace/progress.json"

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 update_progress(current_ep, frame_idx, total_frames, speed_fps, completed_list):
    data = {
        "current_episode": current_ep,
        "current_frame": frame_idx,
        "total_frames": total_frames,
        "percent": round(frame_idx / max(1, total_frames) * 100, 1),
        "speed_fps": round(speed_fps, 2),
        "completed": completed_list,
        "timestamp": time.time()
    }
    with open(PROGRESS_FILE, "w", encoding="utf-8") as f:
        json.dump(data, f, indent=2)

def main():
    log("==================================================")
    log("STARTING HISTORICAL 4K BATCH UPSCALE ON RTX 4090")
    log("==================================================")

    # 1. Start HTTP Server in background for live downloads
    log("Starting background HTTP server on port 8888...")
    subprocess.Popen(
        ["python3", "-m", "http.server", "8888", "--directory", OUTPUT_DIR],
        stdout=subprocess.DEVNULL,
        stderr=subprocess.DEVNULL
    )

    # 2. Load and compile RealESRGAN_x4plus
    log("Loading RealESRGAN_x4plus model weights...")
    model = RRDBNet(num_in_ch=3, num_out_ch=3, num_feat=64, num_block=23, num_grow_ch=32, scale=4).cuda().half().eval()
    loadnet = torch.load("/workspace/Real-ESRGAN/experiments/pretrained_models/RealESRGAN_x4plus.pth", map_location="cuda")
    model.load_state_dict(loadnet.get("params_ema", loadnet.get("params", loadnet)), strict=True)

    log("Compiling model with torch.compile (reduce-overhead)...")
    compiled_model = torch.compile(model, mode="reduce-overhead")
    dummy = torch.zeros(1, 3, 720, 1280, device="cuda", dtype=torch.float16)
    with torch.no_grad():
        _ = compiled_model(dummy)
    torch.cuda.synchronize()
    log("Model compiled and verified on NVIDIA RTX 4090!")

    completed_eps = []

    for ep_name, drive_id 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:
            log(f"Skipping {ep_name}: Already finished ({os.path.getsize(out_video)/(1024*1024):.1f} MB)")
            completed_eps.append(ep_name)
            continue

        in_video = os.path.join(INPUT_DIR, f"{ep_name}.mp4")
        if not os.path.exists(in_video):
            log(f"Downloading {ep_name}.mp4 from Google Drive (ID: {drive_id})...")
            gdown.download(id=drive_id, output=in_video, quiet=True)
            log(f"Downloaded {ep_name}.mp4 ({os.path.getsize(in_video)/(1024*1024):.1f} MB)")

        log(f"\n--- PROCESSING {ep_name.upper()} ---")
        audio_file = os.path.join(INPUT_DIR, f"{ep_name}_audio.aac")
        log("Extracting audio track...")
        subprocess.run(
            ["ffmpeg", "-y", "-i", in_video, "-vn", "-c:a", "aac", "-b:a", "192k", audio_file],
            check=True,
            stdout=subprocess.DEVNULL,
            stderr=subprocess.DEVNULL
        )

        cap = cv2.VideoCapture(in_video)
        total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
        fps = cap.get(cv2.CAP_PROP_FPS) or 24.0
        width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
        height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
        log(f"{ep_name}: {width}x{height} @ {fps} fps, Total frames: {total_frames}")

        writer_cmd = [
            "ffmpeg", "-y",
            "-f", "rawvideo",
            "-pix_fmt", "bgr24",
            "-s", "3840x2160",
            "-r", str(fps),
            "-i", "pipe:0",
            "-i", audio_file,
            "-c:v", "libx265",
            "-preset", "ultrafast",
            "-crf", "20",
            "-pix_fmt", "yuv420p",
            "-c:a", "copy",
            "-shortest",
            out_video
        ]
        proc = subprocess.Popen(writer_cmd, stdin=subprocess.PIPE, stderr=subprocess.DEVNULL)

        ep_start_time = time.time()
        frame_idx = 0

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

            img = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
            tensor = torch.from_numpy(img).permute(2, 0, 1).unsqueeze(0).to("cuda", dtype=torch.float16).div_(255.0)

            with torch.no_grad():
                out = compiled_model(tensor)
                out = torch.nn.functional.interpolate(out, size=(2160, 3840), mode="bicubic", align_corners=False)
                out = out.squeeze(0).clamp_(0.0, 1.0).mul_(255.0).to(torch.uint8)
                out = out[[2, 1, 0], :, :]
                out_bytes = out.permute(1, 2, 0).cpu().numpy().tobytes()
                proc.stdin.write(out_bytes)

            if frame_idx % 100 == 0 or frame_idx == total_frames:
                elapsed = time.time() - ep_start_time
                cur_fps = frame_idx / max(0.1, elapsed)
                rem_frames = total_frames - frame_idx
                rem_min = (rem_frames / max(0.1, cur_fps)) / 60
                pct = frame_idx / total_frames * 100
                log(f"[{ep_name}] Frame {frame_idx}/{total_frames} ({pct:.1f}%) | Speed: {cur_fps:.2f} FPS | ETA: {rem_min:.1f} mins")
                update_progress(ep_name, frame_idx, total_frames, cur_fps, completed_eps)

        proc.stdin.close()
        proc.wait()
        cap.release()

        ep_total_time = time.time() - ep_start_time
        out_mb = os.path.getsize(out_video) / (1024 * 1024)
        log(f"FINISHED {ep_name}! Output: {out_video} ({out_mb:.1f} MB) in {ep_total_time/60:.1f} mins")
        completed_eps.append(ep_name)
        update_progress(None, 0, 0, 0, completed_eps)

        # Cleanup input mp4 and audio to save container disk
        if os.path.exists(in_video):
            os.remove(in_video)
        if os.path.exists(audio_file):
            os.remove(audio_file)
        log(f"Cleaned up input files for {ep_name} to maintain zero storage footprint.")

    log("\n==================================================")
    log("ALL 7 EPISODES UPSCALE TO 4K COMPLETED SUCCESSFULLY!")
    log("==================================================")

if __name__ == "__main__":
    main()
