#!/usr/bin/env python3
"""Reproduce the published Apollo 13 text and speaker metrics."""

from __future__ import annotations

import argparse
import itertools
import json
import re
from dataclasses import dataclass
from pathlib import Path
from typing import Any


ROLES = ("SPACECRAFT", "HOUSTON")


@dataclass
class Turn:
    speaker: str
    text: str


def tokens(text: str) -> list[str]:
    normalized = text.lower().replace("\u2019", "'").replace("\u2014", " ").replace("\u2013", " ")
    return re.findall(r"[a-z0-9]+(?:'[a-z0-9]+)?", normalized)


def edit_distance(reference: list[Any], hypothesis: list[Any]) -> int:
    previous = list(range(len(hypothesis) + 1))
    for row, reference_item in enumerate(reference, 1):
        current = [row] + [0] * len(hypothesis)
        for column, hypothesis_item in enumerate(hypothesis, 1):
            current[column] = min(
                previous[column] + 1,
                current[column - 1] + 1,
                previous[column - 1] + (reference_item != hypothesis_item),
            )
        previous = current
    return previous[-1]


def word_error_rate(reference: list[Any], hypothesis: list[Any]) -> float:
    return edit_distance(reference, hypothesis) / max(1, len(reference))


def align_words(
    reference: list[tuple[str, str]],
    hypothesis: list[tuple[str, str]],
) -> list[tuple[tuple[str, str] | None, tuple[str, str] | None]]:
    reference_words = [word for _, word in reference]
    hypothesis_words = [word for _, word in hypothesis]
    rows, columns = len(reference_words), len(hypothesis_words)
    distances = [[0] * (columns + 1) for _ in range(rows + 1)]
    backtrack = [[""] * (columns + 1) for _ in range(rows + 1)]

    for row in range(1, rows + 1):
        distances[row][0] = row
        backtrack[row][0] = "delete"
    for column in range(1, columns + 1):
        distances[0][column] = column
        backtrack[0][column] = "insert"

    for row in range(1, rows + 1):
        for column in range(1, columns + 1):
            options = [
                (distances[row - 1][column] + 1, "delete"),
                (distances[row][column - 1] + 1, "insert"),
                (
                    distances[row - 1][column - 1]
                    + (reference_words[row - 1] != hypothesis_words[column - 1]),
                    "substitute",
                ),
            ]
            distances[row][column], backtrack[row][column] = min(
                options, key=lambda option: option[0]
            )

    alignment = []
    row, column = rows, columns
    while row or column:
        operation = backtrack[row][column]
        if operation == "substitute":
            alignment.append((reference[row - 1], hypothesis[column - 1]))
            row -= 1
            column -= 1
        elif operation == "delete":
            alignment.append((reference[row - 1], None))
            row -= 1
        else:
            alignment.append((None, hypothesis[column - 1]))
            column -= 1
    return list(reversed(alignment))


def parse_turns(path: Path, *, reference: bool) -> list[Turn]:
    if not reference and path.suffix.lower() == ".json":
        payload = json.loads(path.read_text(encoding="utf-8"))
        return [
            Turn(str(segment.get("speaker") or "SPEAKER_??"), str(segment.get("text") or ""))
            for segment in payload.get("segments", [])
            if segment.get("text")
        ]

    turns = []
    speaker = None
    text = []
    reference_label = re.compile(r"^(SPEAKER_[0-9?]+(?:\s+OR\s+SPEAKER_[0-9?]+)*):$")
    output_label = re.compile(r"^(SPEAKER_[0-9?]+)(?:\s+\[[^\]]+\])?$")

    def append_turn() -> None:
        if speaker and text:
            role = "HOUSTON" if reference and "SPEAKER_01" in speaker else speaker
            if reference and role != "HOUSTON":
                role = "SPACECRAFT"
            turns.append(Turn(role, " ".join(text)))

    for raw_line in path.read_text(encoding="utf-8").splitlines():
        line = raw_line.strip()
        match = (reference_label if reference else output_label).match(line)
        if match:
            append_turn()
            speaker = match.group(1)
            text = []
        elif line and not line.startswith("Detected Language:") and speaker:
            text.append(line)
    append_turn()
    return turns


def role_words(turns: list[Turn], mapping: dict[str, str] | None = None) -> list[tuple[str, str]]:
    return [
        ((mapping or {}).get(turn.speaker, turn.speaker), word)
        for turn in turns
        for word in tokens(turn.text)
    ]


def score(reference_path: Path, output_path: Path) -> dict[str, Any]:
    reference = role_words(parse_turns(reference_path, reference=True))
    output_turns = parse_turns(output_path, reference=False)
    output_speakers = sorted({turn.speaker for turn in output_turns})
    best = None

    for roles in itertools.product(ROLES, repeat=len(output_speakers)):
        mapping = dict(zip(output_speakers, roles))
        hypothesis = role_words(output_turns, mapping)
        text_wer = word_error_rate(
            [word for _, word in reference],
            [word for _, word in hypothesis],
        )
        speaker_wer = word_error_rate(reference, hypothesis)
        alignment = align_words(reference, hypothesis)
        matched = [
            (reference_item, output_item)
            for reference_item, output_item in alignment
            if reference_item and output_item and reference_item[1] == output_item[1]
        ]
        role_accuracy = sum(
            reference_item[0] == output_item[0]
            for reference_item, output_item in matched
        ) / max(1, len(matched))
        candidate = {
            "file": output_path.name,
            "text_wer": text_wer,
            "speaker_attributed_wer": speaker_wer,
            "matched_word_role_accuracy": role_accuracy,
            "reference_words": len(reference),
            "predicted_words": len(hypothesis),
            "matched_words_for_role_scoring": len(matched),
            "speaker_mapping": mapping,
        }
        if best is None or (speaker_wer, text_wer) < (
            best["speaker_attributed_wer"],
            best["text_wer"],
        ):
            best = candidate

    return best or {}


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("reference", type=Path)
    parser.add_argument("outputs", type=Path, nargs="+")
    args = parser.parse_args()
    print(json.dumps([score(args.reference, output) for output in args.outputs], indent=2))


if __name__ == "__main__":
    main()
