import argparse
import base64
import json
import math
import os
import shutil
import subprocess
import tempfile
import urllib.error
import urllib.parse
import urllib.request
import warnings

os.environ["NUMBA_DISABLE_JIT"] = "1"
os.environ["NUMBA_CACHE_DIR"] = os.path.join(tempfile.gettempdir(), "numba-cache")

import numpy as np
import soundfile as sf
from scipy.signal import resample_poly

try:
    import parselmouth
except ImportError:
    parselmouth = None

try:
    import webrtcvad
except ImportError:
    webrtcvad = None

try:
    from vosk import KaldiRecognizer, Model, SetLogLevel
except ImportError:
    KaldiRecognizer = None
    Model = None

    def SetLogLevel(*_args, **_kwargs):
        return None


TARGET_SAMPLE_RATE = 16000
MIN_PITCH_HZ = 60.0
MAX_PITCH_HZ = 350.0
TARGET_RANGE_MIN_HZ = 50.0
TARGET_RANGE_MAX_HZ = 150.0
PACE_TARGET_MIN_WPM = 120.0
PACE_TARGET_MAX_WPM = 180.0
TRANSCRIPTION_CHUNK_SECONDS = 45.0
FILLER_WORDS = {"uh", "um", "er", "erm", "hmm", "mmm"}


def resolve_word_recognition_provider():
    raw = os.environ.get("SPEECH_AUDIO_INSIGHTS_WORDS_PROVIDER", "").strip().lower()
    if raw in {"vosk", "google", "none", "disabled"}:
        return raw
    return "google"


def detect_container(path):
    try:
        with open(path, "rb") as handle:
            header = handle.read(64)
    except OSError:
        return "unknown"

    lowered = header.lower()
    if header.startswith(b"RIFF") and header[8:12] == b"WAVE":
        return "wav"
    if header.startswith(b"\x1a\x45\xdf\xa3"):
        return "webm" if b"webm" in lowered else "matroska"
    if header.startswith(b"OggS"):
        return "ogg"
    if header.startswith(b"fLaC"):
        return "flac"
    if header.startswith(b"ID3") or header[:2] in {b"\xff\xfb", b"\xff\xf3", b"\xff\xf2"}:
        return "mp3"
    return "unknown"


def normalize_audio(audio):
    audio = np.asarray(audio)
    if audio.ndim > 1:
        audio = np.mean(audio, axis=1)
    if np.issubdtype(audio.dtype, np.integer):
        info = np.iinfo(audio.dtype)
        scale = float(max(abs(info.min), info.max))
        audio = audio.astype(np.float32) / scale
    else:
        audio = audio.astype(np.float32, copy=False)

    audio = np.nan_to_num(audio, nan=0.0, posinf=0.0, neginf=0.0)
    peak = float(np.max(np.abs(audio))) if audio.size else 0.0
    if peak > 1.5:
        audio = audio / peak
    return audio.astype(np.float32, copy=False)


def resample_audio(audio, sr, target_sr=TARGET_SAMPLE_RATE):
    if sr == target_sr:
        return audio.astype(np.float32, copy=False), sr
    gcd = math.gcd(int(sr), int(target_sr))
    up = target_sr // gcd
    down = sr // gcd
    resampled = resample_poly(audio, up, down)
    return resampled.astype(np.float32, copy=False), target_sr


def resolve_ffmpeg_binary():
    candidates = [
        os.environ.get("SPEECHSUPER_FFMPEG_BINARY", "").strip(),
        os.environ.get("FFMPEG_BINARY", "").strip(),
        "ffmpeg",
    ]
    for candidate in candidates:
        if candidate and (os.path.isfile(candidate) or shutil.which(candidate)):
            return candidate
    return ""


def load_audio_with_ffmpeg(path):
    ffmpeg_binary = resolve_ffmpeg_binary()
    if ffmpeg_binary == "":
        raise RuntimeError("ffmpeg is required to decode this audio format.")

    command = [
        ffmpeg_binary,
        "-v",
        "error",
        "-i",
        path,
        "-f",
        "s16le",
        "-acodec",
        "pcm_s16le",
        "-ac",
        "1",
        "-ar",
        str(TARGET_SAMPLE_RATE),
        "-",
    ]
    process = subprocess.run(command, stdout=subprocess.PIPE, stderr=subprocess.PIPE, check=False)
    if process.returncode != 0:
        stderr = process.stderr.decode("utf-8", "ignore").strip()
        raise RuntimeError(stderr or "ffmpeg failed to decode the input audio.")

    audio = np.frombuffer(process.stdout, dtype=np.int16)
    if audio.size == 0:
        raise RuntimeError("ffmpeg returned empty PCM data.")
    return normalize_audio(audio), TARGET_SAMPLE_RATE, "ffmpeg"


def load_audio(path):
    if not os.path.isfile(path):
        raise FileNotFoundError(f"Audio file not found: {path}")

    container = detect_container(path)
    errors = []
    try:
        data, sr = sf.read(path, always_2d=False)
        audio = normalize_audio(data)
        audio, sr = resample_audio(audio, sr)
        return audio, sr, "soundfile", container
    except Exception as exc:
        errors.append(f"soundfile: {exc}")

    try:
        audio, sr, backend = load_audio_with_ffmpeg(path)
        return audio, sr, backend, container
    except Exception as exc:
        errors.append(f"ffmpeg: {exc}")

    raise RuntimeError("; ".join(errors))


def estimate_activity_threshold(audio):
    if audio.size == 0:
        return 0.01
    percentile = float(np.percentile(np.abs(audio), 65))
    return max(0.003, min(0.03, percentile * 0.2 if percentile > 0 else 0.01))


def simple_trim(audio, threshold=None):
    threshold = estimate_activity_threshold(audio) if threshold is None else float(threshold)
    mask = np.abs(audio) > threshold
    if not np.any(mask):
        return audio
    start = int(np.argmax(mask))
    end = int(len(audio) - np.argmax(mask[::-1]))
    return audio[start:end]


def audio_to_pcm16(audio):
    clipped = np.clip(np.asarray(audio, dtype=np.float32), -1.0, 1.0)
    return (clipped * 32767.0).astype(np.int16)


def resolve_vosk_model_path():
    configured = os.environ.get("SPEECH_VOSK_MODEL_PATH", "").strip() or os.environ.get("VOSK_MODEL_PATH", "").strip()
    script_dir = os.path.dirname(os.path.abspath(__file__))
    candidates = [
        configured,
        os.path.abspath(os.path.join(script_dir, "..", "models", "vosk", "vosk-model-small-en-us-0.15")),
        os.path.abspath(os.path.join(script_dir, "..", "models", "vosk", "vosk-model-small-en-in-0.4")),
        os.path.join(os.path.expanduser("~"), "AppData", "Local", "vosk", "vosk-model-small-en-us-0.15"),
    ]
    for candidate in candidates:
        if candidate and os.path.isdir(candidate):
            return candidate
    return ""


def extract_pitch_contour(audio, sr, duration):
    if parselmouth is None:
        raise RuntimeError("Parselmouth is not installed.")

    time_step = max(0.05, min(0.15, duration / 120.0 if duration > 0 else 0.1))
    sound = parselmouth.Sound(audio, sampling_frequency=sr)
    pitch = sound.to_pitch(time_step=time_step, pitch_floor=MIN_PITCH_HZ, pitch_ceiling=MAX_PITCH_HZ)
    frequencies = np.asarray(pitch.selected_array["frequency"], dtype=np.float32)
    times = np.asarray(pitch.xs(), dtype=np.float32)
    voiced_mask = np.isfinite(frequencies) & (frequencies > 0.0)
    return times[voiced_mask], frequencies[voiced_mask]


def reject_pitch_outliers(times, frequencies):
    if frequencies.size < 20:
        return times, frequencies
    median = float(np.median(frequencies))
    mask = (frequencies >= median * 0.6) & (frequencies <= median * 1.6)
    filtered_times = times[mask]
    filtered_freqs = frequencies[mask]
    if filtered_freqs.size < 20:
        return times, frequencies
    return filtered_times, filtered_freqs


def bucket_points(times, values, duration, max_points=18):
    if values.size == 0:
        return []
    bucket_count = min(max_points, max(8, int(round(duration / 1.1))))
    edges = np.linspace(0.0, max(duration, float(times[-1])), bucket_count + 1)
    points = []
    for index in range(bucket_count):
        start = edges[index]
        end = edges[index + 1]
        mask = (times >= start) & (times <= end) if index == bucket_count - 1 else (times >= start) & (times < end)
        if not np.any(mask):
            continue
        points.append(
            {
                "time_sec": round(float(np.mean(times[mask])), 2),
                "pitch_hz": round(float(np.median(values[mask])), 2),
            }
        )
    return points


def build_pitch_chart(points, duration):
    observed = [float(point["pitch_hz"]) for point in points if np.isfinite(float(point["pitch_hz"]))]
    upper_candidate = TARGET_RANGE_MAX_HZ + 40.0
    if observed:
        upper_candidate = max(upper_candidate, max(observed) * 1.12)
    y_max = max(300.0, math.ceil(upper_candidate / 50.0) * 50.0)
    return {
        "target_min_hz": TARGET_RANGE_MIN_HZ,
        "target_max_hz": TARGET_RANGE_MAX_HZ,
        "y_min_hz": 0.0,
        "y_max_hz": float(y_max),
        "duration_seconds": round(float(duration), 2),
        "observed_min_hz": round(min(observed), 2) if observed else None,
        "observed_max_hz": round(max(observed), 2) if observed else None,
        "points": points,
    }


def detect_vad_segments(audio, sr, aggressiveness=2, frame_ms=20):
    if webrtcvad is None or sr not in {8000, 16000, 32000, 48000}:
        return [], []

    frame_samples = int(sr * frame_ms / 1000.0)
    pcm16 = audio_to_pcm16(audio)
    frame_count = len(pcm16) // frame_samples
    if frame_count <= 0:
        return [], []

    vad = webrtcvad.Vad(int(aggressiveness))
    frame_duration = frame_ms / 1000.0
    raw_flags = []
    for index in range(frame_count):
        start = index * frame_samples
        end = start + frame_samples
        raw_flags.append(bool(vad.is_speech(pcm16[start:end].tobytes(), sr)))

    flags = []
    for index in range(len(raw_flags)):
        start = max(0, index - 2)
        end = min(len(raw_flags), index + 3)
        votes = raw_flags[start:end]
        flags.append(sum(1 for item in votes if item) >= math.ceil(len(votes) / 2.0))

    segments = []
    current_state = flags[0]
    current_start = 0.0
    for index, state in enumerate(flags[1:], start=1):
        if state == current_state:
            continue
        end_time = index * frame_duration
        segments.append(
            {
                "speech": current_state,
                "start_sec": round(current_start, 3),
                "end_sec": round(end_time, 3),
                "duration_sec": round(end_time - current_start, 3),
            }
        )
        current_state = state
        current_start = end_time

    final_end = round(frame_count * frame_duration, 3)
    segments.append(
        {
            "speech": current_state,
            "start_sec": round(current_start, 3),
            "end_sec": final_end,
            "duration_sec": round(final_end - current_start, 3),
        }
    )

    speech_segments = [segment for segment in segments if segment["speech"] and segment["duration_sec"] >= 0.06]
    silence_segments = [segment for segment in segments if not segment["speech"] and segment["duration_sec"] >= 0.08]
    return speech_segments, silence_segments


def transcribe_words_vosk(audio, sr):
    provider = resolve_word_recognition_provider()
    if provider != "vosk":
        return [], provider, "", ""

    if Model is None or KaldiRecognizer is None:
        raise RuntimeError("Vosk is not installed for word timing analysis.")

    model_path = resolve_vosk_model_path()
    if model_path == "":
        raise RuntimeError("No Vosk model directory was found. Set SPEECH_VOSK_MODEL_PATH or VOSK_MODEL_PATH.")

    SetLogLevel(-1)
    model = Model(model_path)
    recognizer = KaldiRecognizer(model, float(sr))
    recognizer.SetWords(True)
    recognizer.SetPartialWords(False)
    recognizer.SetMaxAlternatives(0)

    pcm_bytes = audio_to_pcm16(audio).tobytes()
    results = []
    for start in range(0, len(pcm_bytes), 4000):
        chunk = pcm_bytes[start : start + 4000]
        if recognizer.AcceptWaveform(chunk):
            results.append(json.loads(recognizer.Result()))
    results.append(json.loads(recognizer.FinalResult()))

    words = []
    for result in results:
        for item in result.get("result", []):
            word = str(item.get("word", "")).strip()
            start_sec = float(item.get("start", 0.0))
            end_sec = float(item.get("end", 0.0))
            if word and end_sec > start_sec:
                words.append(
                    {
                        "word": word,
                        "start_sec": round(start_sec, 3),
                        "end_sec": round(end_sec, 3),
                        "confidence": round(float(item.get("conf", 0.0)), 3),
                    }
                )
    return words, provider, os.path.basename(model_path), model_path


def overlap_duration(start_a, end_a, start_b, end_b):
    return max(0.0, min(end_a, end_b) - max(start_a, start_b))


def classify_pause(duration_sec):
    if duration_sec >= 0.85:
        return "bad"
    if duration_sec >= 0.35:
        return "good"
    if duration_sec >= 0.15:
        return "optional"
    return ""


def merge_pause_and_words(words, silence_segments):
    if not silence_segments:
        return []

    events = []
    if words:
        for index in range(len(words) - 1):
            gap_start = float(words[index]["end_sec"])
            gap_end = float(words[index + 1]["start_sec"])
            if gap_end <= gap_start:
                continue

            overlaps = []
            total_overlap = 0.0
            for pause in silence_segments:
                overlap = overlap_duration(gap_start, gap_end, float(pause["start_sec"]), float(pause["end_sec"]))
                if overlap <= 0.0:
                    continue
                overlaps.append(pause)
                total_overlap += overlap

            if total_overlap < 0.15:
                continue

            event_start = max(gap_start, min(float(item["start_sec"]) for item in overlaps))
            event_end = min(gap_end, max(float(item["end_sec"]) for item in overlaps))
            duration_sec = max(total_overlap, event_end - event_start)
            pause_type = classify_pause(duration_sec)
            if pause_type == "":
                continue

            previous_word = str(words[index]["word"]).lower()
            next_word = str(words[index + 1]["word"]).lower()
            hesitation_type = ""
            if previous_word in FILLER_WORDS or next_word in FILLER_WORDS:
                hesitation_type = "filler_pause"
            elif pause_type == "bad":
                hesitation_type = "bad_pause"

            events.append(
                {
                    "start_sec": round(event_start, 3),
                    "end_sec": round(event_end, 3),
                    "duration_sec": round(duration_sec, 3),
                    "time_sec": round((event_start + event_end) / 2.0, 2),
                    "type": pause_type,
                    "previous_word": words[index]["word"],
                    "next_word": words[index + 1]["word"],
                    "hesitation_type": hesitation_type,
                }
            )
    else:
        for pause in silence_segments:
            duration_sec = float(pause["duration_sec"])
            pause_type = classify_pause(duration_sec)
            if pause_type == "":
                continue
            events.append(
                {
                    "start_sec": round(float(pause["start_sec"]), 3),
                    "end_sec": round(float(pause["end_sec"]), 3),
                    "duration_sec": round(duration_sec, 3),
                    "time_sec": round((float(pause["start_sec"]) + float(pause["end_sec"])) / 2.0, 2),
                    "type": pause_type,
                    "previous_word": "",
                    "next_word": "",
                    "hesitation_type": "bad_pause" if pause_type == "bad" else "",
                }
            )
    return events


def build_pace_chart(words, duration):
    if duration <= 0:
        return {
            "target_min_wpm": PACE_TARGET_MIN_WPM,
            "target_max_wpm": PACE_TARGET_MAX_WPM,
            "y_min_wpm": 0.0,
            "y_max_wpm": 250.0,
            "duration_seconds": 0.0,
            "observed_min_wpm": None,
            "observed_max_wpm": None,
            "points": [],
        }

    window_sec = 10.0 if duration >= 60.0 else 5.0
    bucket_count = max(1, int(math.ceil(duration / window_sec)))
    points = []
    for bucket in range(bucket_count):
        start = bucket * window_sec
        end = min(duration, (bucket + 1) * window_sec)
        count = sum(1 for word in words if start <= float(word["start_sec"]) < end)
        wpm = (count * 60.0) / max(0.001, (end - start))
        points.append({"time_sec": round(start + ((end - start) / 2.0), 2), "wpm": round(wpm, 2)})

    observed = [point["wpm"] for point in points]
    y_max = max(250.0, math.ceil((max(observed) if observed else 250.0) / 25.0) * 25.0)
    return {
        "target_min_wpm": PACE_TARGET_MIN_WPM,
        "target_max_wpm": PACE_TARGET_MAX_WPM,
        "y_min_wpm": 0.0,
        "y_max_wpm": float(y_max),
        "duration_seconds": round(float(duration), 2),
        "observed_min_wpm": round(min(observed), 2) if observed else None,
        "observed_max_wpm": round(max(observed), 2) if observed else None,
        "points": points,
    }


def build_pause_chart(events, duration):
    durations = [float(event["duration_sec"]) for event in events]
    y_max = max(4.0, math.ceil((max(durations) if durations else 4.0) * 2.0) / 2.0)
    return {
        "duration_seconds": round(float(duration), 2),
        "y_min_seconds": 0.0,
        "y_max_seconds": float(y_max),
        "events": events,
    }


def clamp(value, low, high):
    return max(low, min(high, value))


def compute_pitch_stats(frequencies):
    if frequencies.size == 0:
        return {
            "mean": 0.0,
            "median": 0.0,
            "std": 0.0,
            "variation": 0.0,
            "min": 0.0,
            "max": 0.0,
        }

    p10 = float(np.percentile(frequencies, 10))
    p90 = float(np.percentile(frequencies, 90))
    return {
        "mean": float(np.mean(frequencies)),
        "median": float(np.median(frequencies)),
        "std": float(np.std(frequencies)),
        "variation": max(0.0, p90 - p10),
        "min": float(np.min(frequencies)),
        "max": float(np.max(frequencies)),
    }


def analyze_segments(times, frequencies, duration):
    if frequencies.size == 0 or duration <= 0:
        return []

    bucket_count = min(8, max(4, int(math.ceil(duration / 15.0))))
    edges = np.linspace(0.0, duration, bucket_count + 1)
    segments = []
    for index in range(bucket_count):
        start = float(edges[index])
        end = float(edges[index + 1])
        mask = (times >= start) & (times <= end) if index == bucket_count - 1 else (times >= start) & (times < end)
        if not np.any(mask):
            continue

        bucket_values = frequencies[mask]
        segments.append(
            {
                "start_sec": round(start, 3),
                "end_sec": round(end, 3),
                "duration": round(end - start, 3),
                "mean_pitch": round(float(np.mean(bucket_values)), 2),
                "median_pitch": round(float(np.median(bucket_values)), 2),
                "min_pitch": round(float(np.min(bucket_values)), 2),
                "max_pitch": round(float(np.max(bucket_values)), 2),
                "voiced_frames": int(bucket_values.size),
            }
        )

    return segments


def classify_pace_state(words_per_minute):
    if words_per_minute < PACE_TARGET_MIN_WPM:
        return "Slow"
    if words_per_minute > PACE_TARGET_MAX_WPM:
        return "Fast"
    return "Natural"


def score_pace(words_per_minute):
    wpm = float(words_per_minute)
    midpoint = (PACE_TARGET_MIN_WPM + PACE_TARGET_MAX_WPM) / 2.0
    if PACE_TARGET_MIN_WPM <= wpm <= PACE_TARGET_MAX_WPM:
        return round(clamp(100.0 - (abs(wpm - midpoint) * 0.5), 82.0, 100.0))
    if wpm < PACE_TARGET_MIN_WPM:
        penalty = min(75.0, (PACE_TARGET_MIN_WPM - wpm) * 0.7)
        return round(clamp(86.0 - penalty, 15.0, 81.0))
    penalty = min(75.0, (wpm - PACE_TARGET_MAX_WPM) * 0.7)
    return round(clamp(86.0 - penalty, 15.0, 81.0))


def classify_pause_state(score):
    if score >= 80:
        return "Natural"
    if score >= 55:
        return "Needs tuning"
    return "Choppy"


def classify_hesitation_state(score):
    if score >= 80:
        return "Natural"
    if score >= 55:
        return "Noticeable"
    return "Frequent"


def build_hesitation_events(words, pause_events):
    events = []
    seen = set()

    for event in pause_events:
        hesitation_type = str(event.get("hesitation_type") or "").strip()
        if not hesitation_type:
            continue
        key = ("pause", round(float(event["time_sec"]), 2), hesitation_type)
        if key in seen:
            continue
        seen.add(key)
        events.append(
            {
                "time_sec": round(float(event["time_sec"]), 2),
                "type": hesitation_type,
                "duration_sec": round(float(event["duration_sec"]), 3),
                "label": f'{event.get("previous_word", "")} ... {event.get("next_word", "")}'.strip(" ."),
            }
        )

    for word in words:
        normalized = str(word.get("word", "")).lower().strip()
        if normalized not in FILLER_WORDS:
            continue
        midpoint = (float(word["start_sec"]) + float(word["end_sec"])) / 2.0
        key = ("filler", round(midpoint, 2), normalized)
        if key in seen:
            continue
        seen.add(key)
        events.append(
            {
                "time_sec": round(midpoint, 2),
                "type": "filler_word",
                "duration_sec": round(float(word["end_sec"]) - float(word["start_sec"]), 3),
                "label": normalized,
            }
        )

    return sorted(events, key=lambda item: (float(item["time_sec"]), str(item["type"])))


def build_fluency_payload(words, speech_segments, pause_events, duration, recognition_backend, model_name, model_path, recognition_error=""):
    word_count = len(words)
    speech_seconds = sum(float(segment["duration_sec"]) for segment in speech_segments)
    speaking_ratio = (speech_seconds / duration) if duration > 0 else 0.0
    overall_wpm = (word_count * 60.0 / duration) if duration > 0 else 0.0
    articulation_rate_wpm = (word_count * 60.0 / speech_seconds) if speech_seconds > 0 else 0.0
    recognition_confidence = (
        float(np.mean([float(word["confidence"]) for word in words])) if words else 0.0
    )

    good_pause_count = sum(1 for event in pause_events if event["type"] == "good")
    optional_pause_count = sum(1 for event in pause_events if event["type"] == "optional")
    bad_pause_count = sum(1 for event in pause_events if event["type"] == "bad")
    pause_count = len(pause_events)
    average_pause_duration = float(np.mean([float(event["duration_sec"]) for event in pause_events])) if pause_events else 0.0
    longest_pause_duration = max((float(event["duration_sec"]) for event in pause_events), default=0.0)

    hesitation_events = build_hesitation_events(words, pause_events)
    filler_hesitation_count = sum(1 for event in hesitation_events if event["type"] == "filler_word")
    hesitation_count = len(hesitation_events)

    pace_score = score_pace(overall_wpm)

    pause_score = 100.0
    pause_score -= bad_pause_count * 18.0
    pause_score -= max(0.0, longest_pause_duration - 1.2) * 22.0
    pause_score -= max(0.0, average_pause_duration - 0.55) * 20.0
    pause_score += min(good_pause_count * 2.5, 10.0)
    pause_score = clamp(pause_score, 0.0, 100.0)

    hesitation_score = 100.0
    hesitation_score -= filler_hesitation_count * 14.0
    hesitation_score -= bad_pause_count * 16.0
    hesitation_score -= max(0.0, hesitation_count - max(1.0, duration / 35.0)) * 6.0
    hesitation_score = clamp(hesitation_score, 0.0, 100.0)

    if recognition_backend == "vosk":
        analysis_method = "webrtcvad_vosk_v1"
    elif recognition_backend == "google":
        analysis_method = "webrtcvad_google_stt_v1"
    else:
        analysis_method = "webrtcvad_pause_only_v1"

    tips = []
    if overall_wpm < PACE_TARGET_MIN_WPM:
        tips.append("Increase your speaking energy slightly and avoid stretching simple ideas for too long.")
    elif overall_wpm > PACE_TARGET_MAX_WPM:
        tips.append("Slow down a little around key ideas so the listener can process them more easily.")
    else:
        tips.append("Your overall pace is within the target range. Keep using that steady delivery.")

    if bad_pause_count > 0:
        tips.append("Reduce very long pauses between ideas. Plan the next phrase earlier before you stop speaking.")
    elif good_pause_count > 0:
        tips.append("Your pauses are mostly natural. Keep pausing at logical idea boundaries.")
    else:
        tips.append("Add short pauses after important ideas so your speech sounds more controlled.")

    if filler_hesitation_count > 0:
        tips.append("Replace filler words like 'um' or 'uh' with a silent pause and a clean restart.")
    else:
        tips.append("You used few filler hesitations. Maintain that clean delivery.")

    return {
        "available": bool(words or speech_segments or pause_events),
        "analysis_method": analysis_method,
        "recognition_backend": recognition_backend or "",
        "model_name": model_name or "",
        "model_path": model_path or "",
        "recognition_error": recognition_error or "",
        "word_count": int(word_count),
        "words_per_minute": round(overall_wpm, 2),
        "articulation_rate_wpm": round(articulation_rate_wpm, 2),
        "speech_seconds": round(float(speech_seconds), 3),
        "speaking_ratio": round(float(speaking_ratio), 3),
        "recognition_confidence": round(float(recognition_confidence), 3),
        "pace_score": int(round(pace_score)),
        "pause_score": int(round(pause_score)),
        "hesitation_score": int(round(hesitation_score)),
        "pace_state": classify_pace_state(overall_wpm),
        "pause_state": classify_pause_state(pause_score),
        "hesitation_state": classify_hesitation_state(hesitation_score),
        "pause_count": int(pause_count),
        "good_pause_count": int(good_pause_count),
        "optional_pause_count": int(optional_pause_count),
        "bad_pause_count": int(bad_pause_count),
        "hesitation_count": int(hesitation_count),
        "filler_hesitation_count": int(filler_hesitation_count),
        "average_pause_duration": round(float(average_pause_duration), 3),
        "longest_pause_duration": round(float(longest_pause_duration), 3),
        "speech_segments": speech_segments,
        "pause_chart": build_pause_chart(pause_events, duration),
        "pace_chart": build_pace_chart(words, duration),
        "hesitation_events": hesitation_events,
        "pause_events": pause_events,
        "words": words,
        "tips": tips[:3],
    }


def detect_emotion(pitch_stats, fluency):
    variation = float(pitch_stats.get("variation", 0.0))
    pace_state = str(fluency.get("pace_state", "Natural"))
    hesitation_score = float(fluency.get("hesitation_score", 0.0))

    if 55.0 <= variation <= 170.0 and hesitation_score >= 75.0:
        confidence = "high"
    elif 35.0 <= variation <= 210.0 and hesitation_score >= 50.0:
        confidence = "medium"
    else:
        confidence = "low"

    if hesitation_score < 50.0:
        hesitation = "high"
    elif hesitation_score < 80.0:
        hesitation = "medium"
    else:
        hesitation = "low"

    if variation < 50.0:
        stress = "flat"
    elif variation > 170.0 or pace_state == "Fast":
        stress = "high"
    else:
        stress = "balanced"

    return {
        "confidence": confidence,
        "hesitation": hesitation,
        "stress": stress,
        "pause_count": int(fluency.get("pause_count", 0)),
        "avg_pause_duration": round(float(fluency.get("average_pause_duration", 0.0)), 3),
    }


def build_error(message):
    return {"ok": False, "error": str(message)}


def main():
    parser = argparse.ArgumentParser(description="Analyze pitch, pace, pauses, and hesitations from speech audio.")
    parser.add_argument("--input", required=True, help="Path to the input audio file.")
    args = parser.parse_args()

    warnings.filterwarnings("ignore")

    try:
        audio, sr, decoder_backend, input_format = load_audio(args.input)
        silence_threshold = estimate_activity_threshold(audio)
        trimmed_audio = simple_trim(audio, silence_threshold)
        if trimmed_audio.size < int(sr * 0.3):
            trimmed_audio = audio

        duration = float(len(trimmed_audio) / sr) if sr > 0 else 0.0
        raw_duration = float(len(audio) / sr) if sr > 0 else 0.0

        speech_segments, silence_segments = detect_vad_segments(trimmed_audio, sr)
        recognition_error = ""
        recognition_backend = resolve_word_recognition_provider()
        try:
            words, recognition_backend, model_name, model_path = transcribe_words_vosk(trimmed_audio, sr)
        except Exception as exc:
            words, model_name, model_path = [], "", ""
            recognition_error = str(exc)

        pause_events = merge_pause_and_words(words, silence_segments)
        fluency = build_fluency_payload(
            words,
            speech_segments,
            pause_events,
            duration,
            recognition_backend,
            model_name,
            model_path,
            recognition_error=recognition_error,
        )
        pitch_error = ""
        pitch_stats = None
        pitch_chart = None
        pitch_segments = []

        try:
            times, frequencies = extract_pitch_contour(trimmed_audio, sr, duration)
            times, frequencies = reject_pitch_outliers(times, frequencies)
            if frequencies.size == 0:
                pitch_error = "Pitch contour could not be extracted from this audio."
            else:
                pitch_stats = compute_pitch_stats(frequencies)
                pitch_points = bucket_points(times, frequencies, duration)
                pitch_chart = build_pitch_chart(pitch_points, duration)
                pitch_segments = analyze_segments(times, frequencies, duration)
        except Exception as exc:
            pitch_error = str(exc)

        emotion = detect_emotion(pitch_stats or {"variation": 0.0}, fluency)

        payload = {
            "ok": True,
            "duration": round(duration, 3),
            "audio": {
                "input_format": input_format,
                "decoder_backend": decoder_backend,
                "sample_rate": int(sr),
                "original_duration": round(raw_duration, 3),
                "trimmed_duration": round(duration, 3),
                "silence_threshold": round(float(silence_threshold), 5),
            },
            "pitch": {
                "backend": "parselmouth",
                "available": pitch_stats is not None,
                "error": pitch_error,
                "mean": round(float(pitch_stats["mean"]), 3) if pitch_stats is not None else None,
                "median": round(float(pitch_stats["median"]), 3) if pitch_stats is not None else None,
                "std": round(float(pitch_stats["std"]), 3) if pitch_stats is not None else None,
                "variation": round(float(pitch_stats["variation"]), 3) if pitch_stats is not None else None,
                "min": round(float(pitch_stats["min"]), 3) if pitch_stats is not None else None,
                "max": round(float(pitch_stats["max"]), 3) if pitch_stats is not None else None,
            },
            "emotion": emotion,
            "segments": pitch_segments,
            "chart": pitch_chart,
            "fluency": fluency,
        }

        print(json.dumps(payload, ensure_ascii=False))
    except Exception as exc:
        print(json.dumps(build_error(exc), ensure_ascii=False))


if __name__ == "__main__":
    main()
