#!/usr/bin/env python3
"""
概要:
    ローカルのQwen3-ASRサーバーを使用して音声ファイルを文字起こしします。

詳細説明:
    入力ファイルをFFmpegでデコードし、重複を持たせたWAVチャンクに分割して、OpenAI互換の音声認識APIエンドポイントに送信します。チャンク処理が完了するごとにチェックポイントを保存するため、長時間の文字起こしを中断後に再開することができます。
"""

from __future__ import annotations

import argparse
import json
import math
import subprocess
import sys
import tempfile
import time
import urllib.error
import urllib.request
import uuid
from pathlib import Path
from typing import Any

DEFAULT_API_BASE = "http://192.168.27.18:8001/v1"
DEFAULT_MODEL = "Qwen/Qwen3-ASR-1.7B"


def run_command(command: list[str]) -> str:
    """
    概要:
        外部コマンドを実行し標準出力を返します。

    詳細説明:
        subprocessを使用してコマンドを実行します。失敗した場合はエラーメッセージを含めた例外を発生させます。

    引数:
        :param command: 実行するコマンドとその引数のリスト
        :type command: list

    戻り値:
        :returns: コマンドの標準出力文字列
        :rtype: str

    例外:
        :raises RuntimeError: コマンドが見つからないか実行に失敗した場合
    """
    try:
        completed = subprocess.run(
            command, check=True, text=True, stdout=subprocess.PIPE, stderr=subprocess.PIPE
        )
    except FileNotFoundError as exc:
        raise RuntimeError(f"'{command[0]}' が見つかりません。FFmpeg を PATH に追加してください。") from exc
    except subprocess.CalledProcessError as exc:
        details = exc.stderr.strip() or exc.stdout.strip() or "詳細なエラー出力はありません。"
        rendered = subprocess.list2cmdline(command)
        raise RuntimeError(f"FFmpeg コマンドが失敗しました:\n{rendered}\n\n{details}") from exc
    return completed.stdout.strip()


def media_duration_seconds(path: Path) -> float:
    """
    概要:
        音声ファイルの長さを秒単位で取得します。

    詳細説明:
        ffprobeコマンドを使用してメディアの長さを取得します。

    引数:
        :param path: 音声ファイルのパス
        :type path: pathlib.Path

    戻り値:
        :returns: メディアの長さの秒数
        :rtype: float

    例外:
        :raises RuntimeError: 長さの取得に失敗した場合
    """
    text = run_command(
        [
            "ffprobe", "-v", "error", "-show_entries", "format=duration",
            "-of", "default=noprint_wrappers=1:nokey=1", str(path),
        ]
    )
    try:
        return float(text)
    except ValueError as exc:
        raise RuntimeError(f"音声長を取得できませんでした: {text!r}") from exc


def extract_chunk(source: Path, start: float, duration: float, destination: Path) -> None:
    """
    概要:
        音声ファイルの指定区間を抽出してWAV形式で保存します。

    詳細説明:
        ffmpegを使用して、指定した開始位置から指定した長さの音声を16kHzのモノラルWAVファイルとして切り出します。

    引数:
        :param source: 入力音声ファイルのパス
        :type source: pathlib.Path
        :param start: 切り出し開始位置の秒数
        :type start: float
        :param duration: 切り出す長さの秒数
        :type duration: float
        :param destination: 出力WAVファイルのパス
        :type destination: pathlib.Path

    戻り値:
        :returns: なし
        :rtype: None
    """
    run_command(
        [
            "ffmpeg", "-hide_banner", "-loglevel", "error", "-y",
            "-ss", f"{start:.3f}", "-t", f"{duration:.3f}", "-i", str(source),
            "-vn", "-map", "0:a:0", "-ac", "1", "-ar", "16000", "-c:a", "pcm_s16le",
            str(destination),
        ]
    )


def merge_overlap(previous: str, current: str, max_chars: int = 240) -> str:
    """
    概要:
        重複する音声チャンクによって生じたテキストの重複部分を削除して結合します。

    詳細説明:
        前回のテキストの末尾と今回のテキストの先頭を比較し、一致する重複部分を取り除いた今回のテキストを返します。

    引数:
        :param previous: 前回の文字起こしテキスト
        :type previous: str
        :param current: 今回の文字起こしテキスト
        :type current: str
        :param max_chars: 重複を判定する最大文字数
        :type max_chars: int

    戻り値:
        :returns: 重複部分を削除した今回のテキスト
        :rtype: str
    """
    if not previous:
        return current.strip()
    current = current.strip()
    if not current:
        return ""

    # Text is often Japanese without spaces, so compare characters rather than words.
    limit = min(len(previous), len(current), max_chars)
    for length in range(limit, 5, -1):
        if previous[-length:] == current[:length]:
            return current[length:].lstrip()
    return current


def save_json(path: Path, payload: dict[str, Any]) -> None:
    """
    概要:
        辞書データをJSONファイルとして安全に保存します。

    詳細説明:
        一度一時ファイルとして書き出してから目的のパスにリネームすることで、保存中の中断によるデータ破損を防ぎます。

    引数:
        :param path: 保存先のJSONファイルのパス
        :type path: pathlib.Path
        :param payload: 保存する辞書データ
        :type payload: dict

    戻り値:
        :returns: なし
        :rtype: None
    """
    temporary = path.with_suffix(path.suffix + ".tmp")
    temporary.write_text(json.dumps(payload, ensure_ascii=False, indent=2), encoding="utf-8")
    temporary.replace(path)


def transcribe_chunk(api_base: str, model: str, wav_path: Path, timeout: float) -> tuple[str, dict[str, Any]]:
    """
    概要:
        WAVファイルを音声認識APIに送信し文字起こし結果を取得します。

    詳細説明:
        multipart/form-data形式でWAVファイルとモデル名をPOSTリクエストで送信し、結果をJSONとして受け取ります。

    引数:
        :param api_base: APIのエンドポイントのベースURL
        :type api_base: str
        :param model: 使用する音声認識モデルの名前
        :type model: str
        :param wav_path: 送信するWAVファイルのパス
        :type wav_path: pathlib.Path
        :param timeout: APIリクエストのタイムアウト秒数
        :type timeout: float

    戻り値:
        :returns: 文字起こしテキストとAPI応答の辞書のタプル
        :rtype: tuple

    例外:
        :raises RuntimeError: APIへの接続や応答処理に失敗した場合
    """
    endpoint = api_base.rstrip("/") + "/audio/transcriptions"
    boundary = "----QwenASR" + uuid.uuid4().hex
    with wav_path.open("rb") as audio:
        audio_bytes = audio.read()
    body = b"".join(
        [
            f"--{boundary}\r\nContent-Disposition: form-data; name=\"model\"\r\n\r\n{model}\r\n".encode(),
            f"--{boundary}\r\nContent-Disposition: form-data; name=\"file\"; filename=\"{wav_path.name}\"\r\nContent-Type: audio/wav\r\n\r\n".encode(),
            audio_bytes,
            f"\r\n--{boundary}--\r\n".encode(),
        ]
    )
    request = urllib.request.Request(
        endpoint, data=body,
        headers={"Content-Type": f"multipart/form-data; boundary={boundary}"}, method="POST",
    )
    try:
        with urllib.request.urlopen(request, timeout=timeout) as response:
            response_text = response.read().decode("utf-8", errors="replace")
    except urllib.error.HTTPError as exc:
        detail = exc.read().decode("utf-8", errors="replace")[:1000]
        raise RuntimeError(f"ASR API エラー ({exc.code}): {detail}") from exc
    except urllib.error.URLError as exc:
        raise RuntimeError(f"ASR API に接続できません: {exc.reason}") from exc
    try:
        payload = json.loads(response_text)
        text = payload["text"]
    except (ValueError, KeyError, TypeError) as exc:
        raise RuntimeError(f"ASR API の応答形式が不正です: {response_text[:1000]}") from exc
    return str(text).strip(), payload


def parse_args() -> argparse.Namespace:
    """
    概要:
        コマンドライン引数を解析します。

    詳細説明:
        入力ファイルやAPI設定、チャンク設定などの引数を解析して返します。

    戻り値:
        :returns: 解析されたコマンドライン引数
        :rtype: argparse.Namespace
    """
    parser = argparse.ArgumentParser(description="Qwen3-ASR API で WAV/MP3/MP4 を文字起こしします。")
    parser.add_argument("input", type=Path, help="入力 WAV / MP3 / MP4 ファイル")
    parser.add_argument("--api-base", default=DEFAULT_API_BASE, help=f"API の /v1 URL (既定: {DEFAULT_API_BASE})")
    parser.add_argument("--model", default=DEFAULT_MODEL, help=f"ASR モデル名 (既定: {DEFAULT_MODEL})")
    parser.add_argument("--output", type=Path, help="出力 TXT。省略時は入力名.qwen.txt")
    parser.add_argument("--chunk-seconds", type=float, default=240, help="1 チャンクの長さ [s] (既定: 240)")
    parser.add_argument("--overlap-seconds", type=float, default=10, help="隣接チャンクの重なり [s] (既定: 10)")
    parser.add_argument("--timeout", type=float, default=1800, help="各 API 要求のタイムアウト [s] (既定: 1800)")
    parser.add_argument("--resume", action="store_true", help="残っている .partial.json から再開")
    parser.add_argument("--keep-wav", action="store_true", help="分割した一時 WAV を出力先の .chunks に残す")
    return parser.parse_args()


def main() -> int:
    """
    概要:
        文字起こし処理のメインフローを実行します。

    詳細説明:
        コマンドライン引数を解析し、入力全体の長さを取得してチャンクごとに分割と文字起こしを実行します。途中で中断された場合でもチェックポイントファイルから再開できます。

    戻り値:
        :returns: 終了コード
        :rtype: int
    """
    args = parse_args()
    source = args.input.expanduser().resolve()
    if not source.is_file():
        print(f"入力ファイルがありません: {source}", file=sys.stderr)
        return 2
    if args.chunk_seconds <= 0 or args.overlap_seconds < 0 or args.overlap_seconds >= args.chunk_seconds:
        print("--chunk-seconds は正、--overlap-seconds は 0 以上かつ chunk より小さくしてください。", file=sys.stderr)
        return 2

    output = (args.output or source.with_suffix(".qwen.txt")).expanduser().resolve()
    output.parent.mkdir(parents=True, exist_ok=True)
    checkpoint_path = output.with_suffix(output.suffix + ".partial.json")
    segments_path = output.with_suffix(output.suffix + ".segments.json")

    try:
        duration = media_duration_seconds(source)
    except RuntimeError as exc:
        print(f"エラー: {exc}", file=sys.stderr)
        return 1

    starts: list[float] = []
    step = args.chunk_seconds - args.overlap_seconds
    start = 0.0
    while start < duration - 0.01:
        starts.append(start)
        start += step

    state: dict[str, Any] = {
        "source": str(source), "duration_seconds": duration, "api_base": args.api_base,
        "model": args.model, "chunk_seconds": args.chunk_seconds,
        "overlap_seconds": args.overlap_seconds, "segments": [],
    }
    if args.resume and checkpoint_path.exists():
        try:
            saved = json.loads(checkpoint_path.read_text(encoding="utf-8"))
            if Path(saved.get("source", "")).resolve() != source:
                raise RuntimeError("チェックポイントの入力ファイルが異なります。")
            state = saved
            print(f"再開: {len(state['segments'])}/{len(starts)} チャンクは完了済みです。")
        except (json.JSONDecodeError, RuntimeError) as exc:
            print(f"エラー: チェックポイントを読めません: {exc}", file=sys.stderr)
            return 1

    completed_count = len(state["segments"])
    if completed_count > len(starts):
        print("エラー: チェックポイントのチャンク数が入力条件と一致しません。", file=sys.stderr)
        return 1

    chunk_dir: Path | None = output.with_suffix(output.suffix + ".chunks") if args.keep_wav else None
    if chunk_dir:
        chunk_dir.mkdir(exist_ok=True)

    temporary_directory: tempfile.TemporaryDirectory[str] | None = None
    if chunk_dir is None:
        temporary_directory = tempfile.TemporaryDirectory(prefix="qwen-asr-")
        work_dir = Path(temporary_directory.name)
    else:
        work_dir = chunk_dir

    try:
        for index, start in enumerate(starts[completed_count:], start=completed_count):
            length = min(args.chunk_seconds, duration - start)
            wav_path = work_dir / f"chunk_{index:04d}_{start:010.3f}.wav"
            print(f"[{index + 1}/{len(starts)}] {start:.1f}–{start + length:.1f} s: WAV 変換中...", flush=True)
            extract_chunk(source, start, length, wav_path)
            print(f"[{index + 1}/{len(starts)}] ASR 実行中...", flush=True)
            began = time.monotonic()
            text, api_response = transcribe_chunk(args.api_base, args.model, wav_path, args.timeout)
            elapsed = time.monotonic() - began
            print(f"[{index + 1}/{len(starts)}] 完了 ({elapsed:.1f} s): {text[:100]}", flush=True)
            state["segments"].append(
                {
                    "index": index, "start_seconds": start, "end_seconds": start + length,
                    "text": text, "api_response": api_response,
                }
            )
            save_json(checkpoint_path, state)
            if not args.keep_wav:
                wav_path.unlink(missing_ok=True)
    except RuntimeError as exc:
        print(f"エラー: {exc}\n途中結果は {checkpoint_path} に保存されています。--resume で再開できます。", file=sys.stderr)
        return 1
    finally:
        if temporary_directory is not None:
            temporary_directory.cleanup()

    merged = ""
    for segment in state["segments"]:
        addition = merge_overlap(merged, segment["text"])
        if addition:
            merged += ("" if not merged else "\n") + addition
    output.write_text(merged + ("\n" if merged else ""), encoding="utf-8")
    save_json(segments_path, state)
    print(f"\n完了: {output}")
    print(f"区間・API 応答: {segments_path}")
    return 0


if __name__ == "__main__":
    raise SystemExit(main())