#!/usr/bin/env python3

import cv2
import numpy as np
import subprocess
import os
import random
import sys
from pathlib import Path

VID1 = os.path.expanduser("VID00003.mov")
VID2 = os.path.expanduser("VID00004.mov")
AUDIO = os.path.expanduser("background.mp3")
OUTPUT = os.path.expanduser("output.mp4")

TARGET_DURATION = 2*60+41 # 2:41
VID1_START = 9.0
VID1_END   = 209.0
VID2_START = 126.0
VID2_END   = None

MIN_CLIP_LEN = 1.5   # floor to avoid zero/negative-length clips after bounds clamping
CHUNK_LENGTH = 10.0  # analysis window size, also used to keep clips out of the intro/outro
FADE_DURATION = 0.5  # seconds, fade-to-black between clips


def analyze_fpv_motion(video_path, start_offset, end_offset, chunk_length=CHUNK_LENGTH,
                        sample_stride=5, flow_size=(480, 270), max_corners=200):
    """
    Scores chunks by residual motion after compensating for global camera
    motion (a RANSAC-fit affine transform between tracked features between
    consecutive sampled frames). Steady forward flight produces large but
    coherent flow that gets cancelled out; close obstacle passes, erratic
    maneuvers, and parallax against nearby objects produce large residual
    flow and score higher.
    """
    print(f"Analyzing motion in {Path(video_path).name} (range: {start_offset}s to {end_offset if end_offset else 'EOF'})...")
    cap = cv2.VideoCapture(video_path)
    fps = cap.get(cv2.CAP_PROP_FPS)
    if not fps or fps <= 0:
        fps = 30.0

    cap.set(cv2.CAP_PROP_POS_MSEC, start_offset * 1000)
    actual_start = cap.get(cv2.CAP_PROP_POS_MSEC) / 1000.0
    print(f"  requested {start_offset}s, got {actual_start:.2f}s")

    frames_per_chunk = max(1, int(fps * chunk_length))
    feature_params = dict(maxCorners=max_corners, qualityLevel=0.01, minDistance=7, blockSize=7)
    lk_params = dict(winSize=(21, 21), maxLevel=3,
                      criteria=(cv2.TERM_CRITERIA_EPS | cv2.TERM_CRITERIA_COUNT, 30, 0.01))

    chunks = []
    current_start = actual_start
    residual_sum = 0.0
    samples_in_chunk = 0
    frame_count = 0

    ret, prev_frame = cap.read()
    if not ret:
        cap.release()
        return chunks

    prev_gray = cv2.cvtColor(cv2.resize(prev_frame, flow_size), cv2.COLOR_BGR2GRAY)

    def close_chunk():
        if samples_in_chunk > 0:
            chunks.append({
                'start': current_start,
                'score': residual_sum / samples_in_chunk,  # normalized, not summed
            })

    while True:
        current_pos_sec = cap.get(cv2.CAP_PROP_POS_MSEC) / 1000.0
        if end_offset and current_pos_sec >= end_offset:
            break

        ret, frame = cap.read()
        if not ret:
            break

        if frame_count % sample_stride == 0:
            gray = cv2.cvtColor(cv2.resize(frame, flow_size), cv2.COLOR_BGR2GRAY)
            residual = None

            prev_pts = cv2.goodFeaturesToTrack(prev_gray, mask=None, **feature_params)
            if prev_pts is not None and len(prev_pts) >= 8:
                curr_pts, status, _ = cv2.calcOpticalFlowPyrLK(prev_gray, gray, prev_pts, None, **lk_params)
                status = status.reshape(-1).astype(bool)
                good_prev, good_curr = prev_pts[status], curr_pts[status]

                if len(good_prev) >= 8:
                    M, _ = cv2.estimateAffinePartial2D(good_prev, good_curr, method=cv2.RANSAC,
                                                         ransacReprojThreshold=3.0)
                    if M is not None:
                        warped_prev = cv2.warpAffine(prev_gray, M, flow_size)
                        b = 12  # ignore border pixels warpAffine drags in from outside the frame
                        diff = cv2.absdiff(warped_prev[b:-b, b:-b], gray[b:-b, b:-b])
                        residual = float(np.mean(diff))

            if residual is None:
                # too few trackable features (e.g. featureless sky/whiteout) -
                # fall back to raw frame diff for this sample
                residual = float(np.mean(cv2.absdiff(prev_gray, gray)))

            residual_sum += residual
            samples_in_chunk += 1
            prev_gray = gray

        frame_count += 1

        if frame_count >= frames_per_chunk or (end_offset and current_pos_sec >= end_offset):
            close_chunk()
            current_start = current_pos_sec
            residual_sum = 0.0
            samples_in_chunk = 0
            frame_count = 0
            if end_offset and current_pos_sec >= end_offset:
                break

    cap.release()
    return sorted(chunks, key=lambda x: x['score'], reverse=True)


def per_clip_fade(duration, base=FADE_DURATION):
    # keep the fade from swallowing very short clips entirely
    return max(0.05, min(base, duration / 4.0))


def main():
    vid1_chunks = analyze_fpv_motion(VID1, VID1_START, VID1_END)
    vid2_chunks = analyze_fpv_motion(VID2, VID2_START, VID2_END)

    if not vid1_chunks or not vid2_chunks:
        print("Error: Could not extract valid chunks within the specified time bounds.")
        sys.exit(1)

    # --- reserve the intro (start of VID2) and outro (end of VID1) up front ---
    intro_duration = max(MIN_CLIP_LEN, random.uniform(8.0, 12.0))
    if VID2_END:
        intro_duration = min(intro_duration, VID2_END - VID2_START)
    intro_clip = {'file': VID2, 'start': VID2_START, 'duration': intro_duration, 'video': 2}
    intro_end = VID2_START + intro_duration

    outro_duration = max(MIN_CLIP_LEN, min(random.uniform(8.0, 12.0), VID1_END - VID1_START))
    outro_start = VID1_END - outro_duration
    outro_clip = {'file': VID1, 'start': outro_start, 'duration': outro_duration, 'video': 1}

    # keep the general selection pool from reusing that same footage elsewhere:
    # VID1 chunks are only usable up to outro_start, VID2 chunks only from intro_end onward
    vid1_middle_end = outro_start
    vid1_chunks = [c for c in vid1_chunks if c['start'] + CHUNK_LENGTH <= vid1_middle_end]
    vid2_chunks = [c for c in vid2_chunks if c['start'] >= intro_end]

    if not vid1_chunks or not vid2_chunks:
        print("Error: no chunks left after excluding the intro/outro regions "
              "(try a shorter intro/outro or a wider VID1_START/VID1_END/VID2_START/VID2_END range).")
        sys.exit(1)

    vid1_pool = list(vid1_chunks)
    vid2_pool = list(vid2_chunks)

    middle_target = max(0.0, TARGET_DURATION - intro_duration - outro_duration)
    playlist = []
    current_time = 0.0
    turn = 1
    stall_guard = 0

    while current_time < middle_target:
        clip_len = random.uniform(8.0, 12.0)
        if current_time + clip_len > middle_target:
            clip_len = middle_target - current_time

        source_vid = VID1 if turn == 1 else VID2
        active_pool = vid1_pool if turn == 1 else vid2_pool
        backup_pool = vid1_chunks if turn == 1 else vid2_chunks
        end_bound = vid1_middle_end if turn == 1 else VID2_END

        if not active_pool:
            active_pool.extend(backup_pool)

        best_chunk = active_pool.pop(0)

        if end_bound:
            clip_len = min(clip_len, end_bound - best_chunk['start'])
            if clip_len <= MIN_CLIP_LEN:
                stall_guard += 1
                if stall_guard > 500:
                    print("Warning: couldn't fill target duration within bounds, stopping early.")
                    break
                continue
        stall_guard = 0

        playlist.append({
            'file': source_vid,
            'start': best_chunk['start'],
            'duration': clip_len,
            'video': turn,
        })

        current_time += clip_len
        turn = 2 if turn == 1 else 1

        # reorder each video's clips chronologically, preserving alternating slots
        for vid_id, vid_start, vid_end in ((1, VID1_START, vid1_middle_end), (2, VID2_START, VID2_END)):
            idxs = [i for i, c in enumerate(playlist) if c['video'] == vid_id]
            for i, s in zip(idxs, sorted(playlist[i]['start'] for i in idxs)):
                playlist[i]['start'] = max(vid_start, s)
                if vid_end:
                    playlist[i]['duration'] = max(MIN_CLIP_LEN, min(playlist[i]['duration'], vid_end - playlist[i]['start']))

    playlist = [intro_clip] + playlist + [outro_clip]

    print("Building filtergraph with crossfade-to-black transitions...")
    input_args = []
    filter_parts = []
    concat_labels = []

    for idx, clip in enumerate(playlist):
        input_args += ["-ss", f"{clip['start']:.3f}", "-t", f"{clip['duration']:.3f}", "-i", clip['file']]
        fd = per_clip_fade(clip['duration'])
        label = f"v{idx}"
        filter_parts.append(
            f"[{idx}:v]setpts=PTS-STARTPTS,"
            f"fade=t=in:st=0:d={fd:.3f}:c=black,"
            f"fade=t=out:st={max(0.0, clip['duration'] - fd):.3f}:d={fd:.3f}:c=black[{label}]"
        )
        concat_labels.append(f"[{label}]")

    audio_input_index = len(playlist)
    input_args += ["-i", AUDIO]

    filter_parts.append(f"{''.join(concat_labels)}concat=n={len(playlist)}:v=1:a=0[vcat]")
    filter_parts.append("[vcat]format=nv12,hwupload[vout]")
    filter_complex = ";".join(filter_parts)

    print("Rendering final video via AMD VAAPI...")
    ffmpeg_cmd = [
        "ffmpeg",
        "-init_hw_device", "vaapi=foo:/dev/dri/renderD128",
        "-filter_hw_device", "foo",
        *input_args,
        "-filter_complex", filter_complex,
        "-map", "[vout]",
        "-map", f"{audio_input_index}:a",
        "-c:v", "h264_vaapi",
        "-qp", "20",
        "-c:a", "aac",
        "-b:a", "192k",
        "-shortest",
        "-movflags", "+faststart",
        "-y", OUTPUT
    ]

    subprocess.run(ffmpeg_cmd, check=True)
    print(f"Done! Saved to {OUTPUT}")


if __name__ == "__main__":
    main()