#!/usr/bin/env python3
"""Score run.py output: word error rate per file and pooled, plus speed.

  score.py DATA_DIR OUT_DIR  -> writes OUT_DIR/results.json and prints a table

Both reference and hypothesis are passed through the OpenAI Whisper English text
normalizer (lowercase, punctuation removed, numbers and contractions standardised,
filler words "um"/"uh" dropped) before alignment, so vendors are not penalised
for formatting choices. WER = (substitutions + deletions + insertions) / reference words,
computed by jiwer. "Pooled" WER counts errors across all three files together.

pip install jiwer whisper-normalizer
"""
import json
import os
import sys

import jiwer
from whisper_normalizer.english import EnglishTextNormalizer

normalize = EnglishTextNormalizer()


def main():
    data_dir, out_dir = sys.argv[1], sys.argv[2]
    refs = {}
    for p in sorted(os.listdir(data_dir)):
        if p.endswith(".ref.txt"):
            with open(os.path.join(data_dir, p)) as f:
                refs[p.replace(".ref.txt", "")] = normalize(f.read())

    results = {}
    for provider in sorted(os.listdir(out_dir)):
        pdir = os.path.join(out_dir, provider)
        if not os.path.isdir(pdir):
            continue
        per_file = {}
        hyps, ref_list = [], []
        audio_total = wall_total = 0.0
        version = None
        for p in sorted(os.listdir(pdir)):
            if not p.endswith(".json"):
                continue
            with open(os.path.join(pdir, p)) as f:
                r = json.load(f)
            stem = p[:-5]
            hyp = normalize(r["hypothesis"])
            ref = refs[stem]
            m = jiwer.process_words(ref, hyp)
            per_file[stem] = {
                "wer": round(m.wer, 4),
                "substitutions": m.substitutions,
                "deletions": m.deletions,
                "insertions": m.insertions,
                "reference_words": len(ref.split()),
                "wall_seconds": r["wall_seconds"],
                "audio_seconds": round(r["audio_seconds"], 1),
                "run_at": r["run_at"],
            }
            hyps.append(hyp)
            ref_list.append(ref)
            audio_total += r["audio_seconds"]
            wall_total += r["wall_seconds"]
            version = r["version"]
        if not per_file:
            continue
        pooled = jiwer.process_words(ref_list, hyps)
        results[provider] = {
            "version": version,
            "pooled_wer": round(pooled.wer, 4),
            "audio_seconds": round(audio_total, 1),
            "wall_seconds": round(wall_total, 1),
            "speed_x_realtime": round(audio_total / wall_total, 1) if wall_total else None,
            "files": per_file,
        }

    with open(os.path.join(out_dir, "results.json"), "w") as f:
        json.dump(results, f, indent=1)

    print(f"{'provider':14} {'pooled WER':>10} " + " ".join(f"{k[:18]:>18}" for k in refs) + "   speed")
    for name, r in sorted(results.items(), key=lambda kv: kv[1]["pooled_wer"]):
        row = " ".join(f"{r['files'].get(k, {}).get('wer', float('nan'))*100:>17.1f}%" for k in refs)
        print(f"{name:14} {r['pooled_wer']*100:>9.1f}% {row}   {r['speed_x_realtime']}x")


if __name__ == "__main__":
    main()
