import os
import sys
import time
import subprocess
import cv2
import torch
import numpy as np
import tensorrt as trt

ENGINE_PATH = "/workspace/srvgg_x4.engine"
INPUT_DIR = "/workspace/inputs"
OUTPUT_DIR = "/workspace/outputs_4k"

in_video = os.path.join(INPUT_DIR, "ep3.mp4")
audio_file = os.path.join(INPUT_DIR, "ep3_audio.aac")
test_out = os.path.join(OUTPUT_DIR, "test_sync_60f.mp4")

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()
d_input = torch.empty(1, 3, 720, 1280, device="cuda", dtype=torch.float32).contiguous()
d_output = torch.empty(1, 3, 2880, 5120, device="cuda", dtype=torch.float32).contiguous()
context.set_tensor_address("input", d_input.data_ptr())
context.set_tensor_address("output", d_output.data_ptr())

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

# 1 Single dedicated host buffer (Pinned memory for fast PCIe DMA transfer)
host_buf = torch.empty(2160, 3840, 3, dtype=torch.uint8, pin_memory=True)

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

print("Starting STRICTLY SYNCHRONOUS 60-frame test (zero race conditions)...")
t0 = time.time()
N = 60
for i in range(N):
    ret, frame = cap.read()
    if not ret:
        break

    # 1. GPU Preprocess
    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. TensorRT Inference
    with torch.cuda.stream(stream):
        context.execute_async_v3(stream.cuda_stream)
    stream.synchronize()

    # 3. GPU Bicubic downscale to 4K & Contiguous BGR
    out_4k = torch.nn.functional.interpolate(d_output, size=(2160, 3840), mode="bicubic", align_corners=False)
    out_uint8 = out_4k.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. DMA transfer to host buffer
    host_buf.copy_(out_bgr, non_blocking=False)

    # 5. STRICT SYNCHRONOUS WRITE: FFmpeg receives frame i before frame i+1 is ever read!
    # Using bytes() guarantees a complete, immutable snapshot of the frame
    proc.stdin.write(bytes(host_buf.numpy()))

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

dt = time.time() - t0
fps_achieved = N / dt
print(f"DONE: {N} frames in {dt:.2f}s -> {fps_achieved:.1f} FPS!")
print(f"Output video size: {os.path.getsize(test_out)} bytes, exit code: {proc.returncode}")
