import os
import sys
import tempfile
import subprocess
import json
from pathlib import Path

# Add production pipeline to path
sys.path.insert(0, r"c:\Projects\KieuStory\05_Production_Pipeline")
from audio_continuity_engine import AudioContinuityEngine, get_ffmpeg, get_ffprobe

ffmpeg = get_ffmpeg()
ffprobe = get_ffprobe()
engine = AudioContinuityEngine()

results = []

def record(test_name, success, details=""):
    results.append({"test": test_name, "success": success, "details": details})
    status = "✓ PASS" if success else "✗ FAIL"
    print(f"[{status}] {test_name}: {details}")

print("================================================================================")
print("RUNNING ADVERSARIAL STRESS TESTS FOR AUDIO CONTINUITY ENGINE (M3)")
print("================================================================================")

with tempfile.TemporaryDirectory() as tmpdir:
    tmp_path = Path(tmpdir)

    # Helper to generate test video with tone
    def gen_video_with_audio(path, dur=3.0, freq=440, vol=0.5, channels=2, sample_rate=48000):
        # tone audio
        # stereo layout
        chan_layout = "stereo" if channels == 2 else "mono"
        cmd = [
            ffmpeg, "-y",
            "-f", "lavfi", "-i", f"color=c=blue:s=320x240:d={dur}:r=24",
            "-f", "lavfi", "-i", f"sine=frequency={freq}:duration={dur}:sample_rate={sample_rate}",
            "-af", f"volume={vol},aformat=channel_layouts={chan_layout}",
            "-c:v", "libx264", "-pix_fmt", "yuv420p", "-t", str(dur),
            "-c:a", "aac", "-b:a", "192k", "-ar", str(sample_rate),
            str(path)
        ]
        res = subprocess.run(cmd, capture_output=True, text=True)
        return res.returncode == 0

    # Helper to generate mute video
    def gen_mute_video(path, dur=3.0):
        cmd = [
            ffmpeg, "-y",
            "-f", "lavfi", "-i", f"color=c=red:s=320x240:d={dur}:r=24",
            "-c:v", "libx264", "-pix_fmt", "yuv420p", "-t", str(dur),
            "-an",
            str(path)
        ]
        res = subprocess.run(cmd, capture_output=True, text=True)
        return res.returncode == 0

    # Helper to generate audio wav
    def gen_audio(path, dur=3.0, freq=440, vol=0.5):
        cmd = [
            ffmpeg, "-y",
            "-f", "lavfi", "-i", f"sine=frequency={freq}:duration={dur}:sample_rate=48000",
            "-af", f"volume={vol},aformat=channel_layouts=stereo",
            "-c:a", "pcm_s16le", "-ar", "48000",
            str(path)
        ]
        res = subprocess.run(cmd, capture_output=True, text=True)
        return res.returncode == 0

    # -------------------------------------------------------------------------
    # TEST 1: measure_loudness on normal audio
    # -------------------------------------------------------------------------
    v_norm = tmp_path / "v_norm.mp4"
    gen_video_with_audio(v_norm, dur=3.0, freq=440, vol=0.2)
    m_norm = engine.measure_loudness(str(v_norm))
    if m_norm and not m_norm.get("is_silent") and "input_i" in m_norm and "input_tp" in m_norm:
        record("T1_measure_loudness_normal", True, f"Measured I={m_norm['input_i']}, TP={m_norm['input_tp']}, is_silent={m_norm.get('is_silent')}")
    else:
        record("T1_measure_loudness_normal", False, f"Failed measurement: {m_norm}")

    # -------------------------------------------------------------------------
    # TEST 2: measure_loudness silence guard (-inf)
    # -------------------------------------------------------------------------
    v_silence = tmp_path / "v_silence.mp4"
    cmd_sil = [
        ffmpeg, "-y",
        "-f", "lavfi", "-i", "color=c=black:s=320x240:d=3:r=24",
        "-f", "lavfi", "-i", "aevalsrc=0:d=3:s=48000:c=stereo",
        "-c:v", "libx264", "-pix_fmt", "yuv420p", "-t", "3",
        "-c:a", "aac", "-b:a", "192k",
        str(v_silence)
    ]
    subprocess.run(cmd_sil, capture_output=True, text=True)
    m_sil = engine.measure_loudness(str(v_silence))
    if m_sil and m_sil.get("is_silent") is True:
        record("T2_measure_loudness_silence_guard", True, f"Correctly detected silence: input_i={m_sil.get('input_i')}, is_silent={m_sil.get('is_silent')}")
    else:
        record("T2_measure_loudness_silence_guard", False, f"Silence guard failed: {m_sil}")

    # -------------------------------------------------------------------------
    # TEST 3: measure_loudness on nonexistent file
    # -------------------------------------------------------------------------
    m_none = engine.measure_loudness(str(tmp_path / "nonexistent.mp4"))
    record("T3_measure_loudness_nonexistent", m_none is None, f"Returned: {m_none}")

    # -------------------------------------------------------------------------
    # TEST 4: normalize_loudness Two-Pass Linear
    # -------------------------------------------------------------------------
    v_norm_out = tmp_path / "v_norm_out.mp4"
    succ4 = engine.normalize_loudness(str(v_norm), str(v_norm_out), target_lufs=-14.0, two_pass=True)
    if succ4 and v_norm_out.exists():
        # Verify measured loudness on output
        m_after = engine.measure_loudness(str(v_norm_out))
        record("T4_normalize_loudness_twopass", True, f"Output created, Post-I={m_after.get('input_i')}, Post-TP={m_after.get('input_tp')}")
    else:
        record("T4_normalize_loudness_twopass", False, "Normalization failed")

    # -------------------------------------------------------------------------
    # TEST 5: normalize_loudness on pure silence (fallback single-pass without crashing)
    # -------------------------------------------------------------------------
    v_sil_out = tmp_path / "v_sil_out.mp4"
    succ5 = engine.normalize_loudness(str(v_silence), str(v_sil_out), target_lufs=-14.0, two_pass=True)
    record("T5_normalize_loudness_silent_fallback", succ5 and v_sil_out.exists(), f"Succeeded on silence: {succ5}")

    # -------------------------------------------------------------------------
    # TEST 6: stitch_with_audio_crossfade Mode A (30ms micro-fade zero duration loss)
    # -------------------------------------------------------------------------
    v_a1 = tmp_path / "v_a1.mp4"
    v_a2 = tmp_path / "v_a2.mp4"
    v_out_mode_a = tmp_path / "v_out_mode_a.mp4"
    gen_video_with_audio(v_a1, dur=2.0, freq=440)
    gen_video_with_audio(v_a2, dur=2.0, freq=880)
    succ6 = engine.stitch_with_audio_crossfade([str(v_a1), str(v_a2)], str(v_out_mode_a), mode="boundary_smoothing")
    
    # Check duration of output video
    probe_cmd = [ffprobe, "-v", "error", "-show_entries", "format=duration", "-of", "json", str(v_out_mode_a)]
    pr = subprocess.run(probe_cmd, capture_output=True, text=True)
    dur_mode_a = float(json.loads(pr.stdout)["format"]["duration"])
    # Mode A should maintain ~4.0s (2.0 + 2.0)
    record("T6_stitch_mode_a_duration_preservation", succ6 and abs(dur_mode_a - 4.0) < 0.2, f"Success={succ6}, Duration={dur_mode_a:.2f}s (expected ~4.00s)")

    # -------------------------------------------------------------------------
    # TEST 7: stitch_with_audio_crossfade Mode B (acrossfade with apad)
    # -------------------------------------------------------------------------
    v_out_mode_b = tmp_path / "v_out_mode_b.mp4"
    succ7 = engine.stitch_with_audio_crossfade([str(v_a1), str(v_a2)], str(v_out_mode_b), crossfade_dur=0.5, mode="acrossfade")
    pr_b = subprocess.run([ffprobe, "-v", "error", "-show_entries", "format=duration", "-of", "json", str(v_out_mode_b)], capture_output=True, text=True)
    dur_mode_b = float(json.loads(pr_b.stdout)["format"]["duration"])
    # Mode B has apad and -shortest so duration matches video length (~4.0s)
    record("T7_stitch_mode_b_acrossfade_apad", succ7 and abs(dur_mode_b - 4.0) < 0.2, f"Success={succ7}, Duration={dur_mode_b:.2f}s")

    # -------------------------------------------------------------------------
    # TEST 8: stitch_with_audio_crossfade with Mute Clip (aevalsrc fallback)
    # -------------------------------------------------------------------------
    v_mute = tmp_path / "v_mute.mp4"
    v_out_mute_mix = tmp_path / "v_out_mute_mix.mp4"
    gen_mute_video(v_mute, dur=2.0)
    succ8 = engine.stitch_with_audio_crossfade([str(v_a1), str(v_mute)], str(v_out_mute_mix), mode="boundary_smoothing")
    # Check output audio stream exists and is stereo 48kHz
    pr_m = subprocess.run([ffprobe, "-v", "error", "-select_streams", "a", "-show_entries", "stream=channels,sample_rate", "-of", "json", str(v_out_mute_mix)], capture_output=True, text=True)
    streams_m = json.loads(pr_m.stdout).get("streams", [])
    has_audio_stream = len(streams_m) > 0 and streams_m[0]["channels"] == 2 and int(streams_m[0]["sample_rate"]) == 48000
    record("T8_stitch_mute_clip_fallback", succ8 and has_audio_stream, f"Success={succ8}, Audio streams: {streams_m}")

    # -------------------------------------------------------------------------
    # TEST 9: stitch_with_audio_crossfade Mono + Stereo mixed input
    # -------------------------------------------------------------------------
    v_mono = tmp_path / "v_mono.mp4"
    v_out_mono_mix = tmp_path / "v_out_mono_mix.mp4"
    gen_video_with_audio(v_mono, dur=2.0, freq=300, channels=1)
    succ9 = engine.stitch_with_audio_crossfade([str(v_mono), str(v_a1)], str(v_out_mono_mix), mode="acrossfade")
    pr_mono = subprocess.run([ffprobe, "-v", "error", "-select_streams", "a", "-show_entries", "stream=channels,sample_rate", "-of", "json", str(v_out_mono_mix)], capture_output=True, text=True)
    streams_mono = json.loads(pr_mono.stdout).get("streams", [])
    mono_ok = len(streams_mono) > 0 and streams_mono[0]["channels"] == 2
    record("T9_stitch_mono_stereo_preconditioning", succ9 and mono_ok, f"Success={succ9}, Output channels={streams_mono[0]['channels'] if streams_mono else 'none'}")

    # -------------------------------------------------------------------------
    # TEST 10: mix_four_stems 4-Stem mixing with notch filter and dynamic ducking
    # -------------------------------------------------------------------------
    v_base = tmp_path / "v_base.mp4"
    gen_video_with_audio(v_base, dur=4.0, freq=200, vol=0.1) # low background in video
    stem_bgm = tmp_path / "stem1_bgm.wav"
    stem_amb = tmp_path / "stem2_amb.wav"
    stem_fol = tmp_path / "stem3_fol.wav"
    stem_dia = tmp_path / "stem4_dia.wav"
    gen_audio(stem_bgm, dur=2.0, freq=1000, vol=0.5) # shorter than video, tests looping
    gen_audio(stem_amb, dur=2.0, freq=500, vol=0.3)  # shorter than video, tests looping
    gen_audio(stem_fol, dur=4.0, freq=2500, vol=0.4)
    gen_audio(stem_dia, dur=4.0, freq=440, vol=0.8)

    v_4stem_out = tmp_path / "v_4stem_out.mp4"
    succ10 = engine.mix_four_stems(
        str(v_base), str(v_4stem_out),
        stem1_bgm=str(stem_bgm),
        stem2_ambience=str(stem_amb),
        stem3_foley=str(stem_fol),
        stem4_dialogue=str(stem_dia),
        two_pass=True
    )
    if succ10 and v_4stem_out.exists():
        m_4s = engine.measure_loudness(str(v_4stem_out))
        record("T10_mix_four_stems_full_mix", True, f"Output size={v_4stem_out.stat().st_size}, I={m_4s.get('input_i')}, TP={m_4s.get('input_tp')}")
    else:
        record("T10_mix_four_stems_full_mix", False, "mix_four_stems failed")

    # -------------------------------------------------------------------------
    # TEST 11: mix_four_stems Partial Stems (BGM only, no dialogue stem, input video has audio)
    # -------------------------------------------------------------------------
    v_part_out = tmp_path / "v_part_out.mp4"
    succ11 = engine.mix_four_stems(
        str(v_base), str(v_part_out),
        stem1_bgm=str(stem_bgm),
        two_pass=True
    )
    record("T11_mix_four_stems_bgm_only", succ11 and v_part_out.exists(), f"Success={succ11}")

    # -------------------------------------------------------------------------
    # TEST 12: mix_four_stems with Mute Video and NO Dialogue Stem
    # -------------------------------------------------------------------------
    v_mute_base = tmp_path / "v_mute_base.mp4"
    gen_mute_video(v_mute_base, dur=3.0)
    v_mute_4s_out = tmp_path / "v_mute_4s_out.mp4"
    succ12 = engine.mix_four_stems(
        str(v_mute_base), str(v_mute_4s_out),
        stem1_bgm=str(stem_bgm),
        stem2_ambience=str(stem_amb),
        two_pass=True
    )
    record("T12_mix_four_stems_mute_video_no_dialogue", succ12 and v_mute_4s_out.exists(), f"Success={succ12}")

print("================================================================================")
all_pass = all(r["success"] for r in results)
print(f"ADVERSARIAL STRESS TEST SUMMARY: {sum(1 for r in results if r['success'])}/{len(results)} PASSED")
print(f"OVERALL STATUS: {'PASS' if all_pass else 'FAIL'}")
print("================================================================================")
