"""tkttsライブラリファミリー向けのQwen3-TTSバックエンド。

このモジュールは、公式の ``qwen-tts`` Pythonパッケージを介してQwen3-TTSをローカルで実行しながら、
``tktts_voicevox.py`` と互換性のあるインターフェースを提供します。
モデルは遅延ロードされキャッシュされるため、連続した呼び出しで再ロードされることはありません。

関連リンク:
    :doc:`tktts_base`
"""

from __future__ import annotations

import os
import re
from pathlib import Path
from typing import Any


missing = []
for lib, import_name in [
    ("torch", "torch"),
    ("soundfile", "soundfile"),
    ("qwen-tts", "qwen_tts"),
]:
    try:
        __import__(import_name)
    except ImportError:
        missing.append(lib)

if missing:
    raise ImportError(
        "Error: Missing libraries: "
        + ", ".join(missing)
        + "\n  install: pip install qwen-tts soundfile"
    )

import soundfile as sf
import torch
from qwen_tts import Qwen3TTSModel

try:
    from tktts_base import apply_replacements, normalize_speaker, split_dialogue
except ImportError as exc:
    raise ImportError(
        "tktts_qwen3 requires tktts_base.py to be importable."
    ) from exc


TTS_ENGINE_NAME = "qwen3"
DEFAULT_QWEN3_MODEL = "Qwen/Qwen3-TTS-12Hz-0.6B-CustomVoice"
DEFAULT_QWEN3_VOICE = "Ono_Anna"
DEFAULT_LANGUAGE = "Japanese"
DEFAULT_DEVICE = "auto"
DEFAULT_DTYPE = "auto"


# Official speakers included in the Qwen3-TTS CustomVoice checkpoints.
_AVAILABLE_VOICES = [
    {
        "name": "Vivian",
        "native_language": "Chinese",
        "description": "Bright, slightly edgy young female voice",
    },
    {
        "name": "Serena",
        "native_language": "Chinese",
        "description": "Warm, gentle young female voice",
    },
    {
        "name": "Uncle_Fu",
        "native_language": "Chinese",
        "description": "Seasoned male voice with a low, mellow timbre",
    },
    {
        "name": "Dylan",
        "native_language": "Chinese",
        "description": "Youthful Beijing male voice with a clear timbre",
    },
    {
        "name": "Eric",
        "native_language": "Chinese",
        "description": "Lively Chengdu male voice with husky brightness",
    },
    {
        "name": "Ryan",
        "native_language": "English",
        "description": "Dynamic male voice with strong rhythmic drive",
    },
    {
        "name": "Aiden",
        "native_language": "English",
        "description": "Sunny American male voice with a clear midrange",
    },
    {
        "name": "Ono_Anna",
        "native_language": "Japanese",
        "description": "Playful Japanese female voice with a light timbre",
    },
    {
        "name": "Sohee",
        "native_language": "Korean",
        "description": "Warm Korean female voice with rich emotion",
    },
]


_MODEL_CACHE: dict[tuple[str, str, str], Qwen3TTSModel] = {}


def _resolve_device(device: str | None = DEFAULT_DEVICE) -> str:
    """デバイス設定を解決し、利用可能なCUDAデバイスまたはCPUを決定します。

    ``auto`` が指定された場合、CUDAが利用可能であれば ``cuda:0`` を、
    そうでなければ ``cpu`` を返します。

    :param device: 使用するデバイス名（例: "cuda:0", "cpu", "auto"）。
                   Noneまたは"auto"の場合、自動的に決定されます。
    :type device: str | None
    :returns: 解決されたデバイス名。
    :rtype: str
    """
    if device is None or str(device).strip().lower() == "auto":
        return "cuda:0" if torch.cuda.is_available() else "cpu"
    return str(device)


def _resolve_dtype(device: str, dtype: str | torch.dtype | None) -> torch.dtype:
    """データ型名をTorchのデータ型に解決します。

    指定されたデバイスに適したデータ型を決定します。
    ``auto`` が指定された場合、CUDAデバイスでは ``bfloat16`` または ``float16`` を、
    CPUでは ``float32`` を使用します。

    :param device: ターゲットデバイス名。
    :type device: str
    :param dtype: 使用するデータ型名（例: "auto", "bfloat16", "float16", "float32"）
                  または ``torch.dtype`` オブジェクト。
    :type dtype: str | torch.dtype | None
    :returns: 解決された ``torch.dtype`` オブジェクト。
    :rtype: torch.dtype
    :raises ValueError: サポートされていないデータ型が指定された場合。
    """
    if isinstance(dtype, torch.dtype):
        return dtype

    name = "auto" if dtype is None else str(dtype).strip().lower()
    if name == "auto":
        if device.startswith("cuda"):
            return torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16
        return torch.float32

    dtype_map = {
        "bfloat16": torch.bfloat16,
        "bf16": torch.bfloat16,
        "float16": torch.float16,
        "fp16": torch.float16,
        "float32": torch.float32,
        "fp32": torch.float32,
    }
    if name not in dtype_map:
        raise ValueError(f"Unsupported dtype: {dtype}")
    return dtype_map[name]


def load_model(
    model_id: str = DEFAULT_QWEN3_MODEL,
    device: str = DEFAULT_DEVICE,
    dtype: str | torch.dtype = DEFAULT_DTYPE,
    *,
    force_reload: bool = False,
) -> Qwen3TTSModel:
    """Qwen3-TTSモデルをロードし、キャッシュします。

    FlashAttentionは意図的に要求されていません。これにより、個別にCUDA SDKをインストールせずに
    Windows上でも通常のPyTorchアテンションパスが動作します。
    モデルは、``model_id``, ``device``, ``dtype`` の組み合わせをキーとしてキャッシュされます。
    既にキャッシュされているモデルは再ロードされません。

    :param model_id: ロードするQwen3-TTSモデルのID（例: "Qwen/Qwen3-TTS-12Hz-0.6B-CustomVoice"）。
    :type model_id: str
    :param device: モデルをロードするデバイス（例: "cuda:0", "cpu", "auto"）。
    :type device: str
    :param dtype: モデルのデータ型（例: "auto", "bfloat16", "float16", "float32"）。
    :type dtype: str | torch.dtype
    :param force_reload: Trueの場合、キャッシュを無視してモデルを強制的に再ロードします。
    :type force_reload: bool
    :returns: ロードされたQwen3TTSModelインスタンス。
    :rtype: Qwen3TTSModel
    """
    resolved_device = _resolve_device(device)
    resolved_dtype = _resolve_dtype(resolved_device, dtype)
    cache_key = (model_id, resolved_device, str(resolved_dtype))

    if force_reload:
        _MODEL_CACHE.pop(cache_key, None)

    if cache_key not in _MODEL_CACHE:
        print(
            "tktts_qwen3.load_model(): "
            f"model={model_id}, device={resolved_device}, dtype={resolved_dtype}"
        )
        _MODEL_CACHE[cache_key] = Qwen3TTSModel.from_pretrained(
            model_id,
            device_map=resolved_device,
            dtype=resolved_dtype,
        )

    return _MODEL_CACHE[cache_key]


def unload_models() -> None:
    """このモジュールによってキャッシュされたすべてのモデルを解放します。

    GPUメモリを使用している場合、TorchのCUDAキャッシュもクリアします。

    :returns: なし。
    :rtype: None
    """
    _MODEL_CACHE.clear()
    if torch.cuda.is_available():
        torch.cuda.empty_cache()


def get_available_voices_info(model: Qwen3TTSModel | None = None):
    """利用可能な話者名とその説明を返します。

    モデルインスタンスが提供された場合、そのモデルが報告する話者リストが使用されます。
    これにより、将来のCustomVoiceチェックポイントとの互換性が保たれます。
    モデルが提供されない場合、または ``get_supported_speakers`` メソッドがない場合は、
    このモジュールにハードコードされたデフォルトのリストが使用されます。

    :param model: 問い合わせるQwen3TTSModelインスタンス。Noneの場合はデフォルトリストを使用。
    :type model: Qwen3TTSModel | None
    :returns: 各話者の名前、母国語、説明を含む辞書のリスト。
    :rtype: list[dict]
    """
    if model is None:
        return [voice.copy() for voice in _AVAILABLE_VOICES]

    get_speakers = getattr(model, "get_supported_speakers", None)
    if not callable(get_speakers):
        return [voice.copy() for voice in _AVAILABLE_VOICES]

    known = {voice["name"].lower(): voice for voice in _AVAILABLE_VOICES}
    voices = []
    for speaker in get_speakers():
        info = known.get(str(speaker).lower(), {})
        voices.append(
            {
                "name": str(speaker),
                "native_language": info.get("native_language", ""),
                "description": info.get("description", ""),
            }
        )
    return voices


def get_available_voices(model: Qwen3TTSModel | None = None):
    """利用可能な話者名をリストで返します。

    これは ``get_available_voices_info()`` から話者名のみを抽出するヘルパー関数です。

    :param model: 問い合わせるQwen3TTSModelインスタンス。Noneの場合はデフォルトリストを使用。
    :type model: Qwen3TTSModel | None
    :returns: 利用可能な話者名のリスト。
    :rtype: list[str]
    """
    return [voice["name"] for voice in get_available_voices_info(model)]


def list_available_voices(model: Qwen3TTSModel | None = None) -> bool:
    """利用可能な話者をコンソールに出力します。

    :param model: 問い合わせるQwen3TTSModelインスタンス。Noneの場合はデフォルトリストを使用。
    :type model: Qwen3TTSModel | None
    :returns: 利用可能な話者が出力された場合はTrue、そうでなければFalse。
    :rtype: bool
    """
    print(f"=== 利用可能な {TTS_ENGINE_NAME} voices ===")
    voices = get_available_voices_info(model)
    if not voices:
        return False

    for voice in voices:
        language = voice.get("native_language") or "-"
        description = voice.get("description") or "-"
        print(
            f"  Name: {voice['name']}, "
            f"Native language: {language}, Description: {description}"
        )
    return True


def _speaker_key(name: str) -> str:
    """話者名を正規化し、検索用のキーを生成します。

    スペース、アンダースコア、ハイフン、括弧を削除し、すべて小文字に変換します。

    :param name: 処理する話者名。
    :type name: str
    :returns: 正規化された話者キー。
    :rtype: str
    """
    normalized = normalize_speaker(str(name))
    return re.sub(r"[\s_\-（）()]+", "", normalized).lower()


def resolve_speaker_id(
    speaker_name: str,
    voices_info=None,
    model: Qwen3TTSModel | None = None,
):
    """フルまたは部分的な話者名をQwenの話者名に解決します。

    指定された話者名に基づいて、最も一致する話者名を検索します。
    まず正規化された完全一致を優先し、次に部分一致を探します。

    :param speaker_name: 解決したい話者名（フル名または部分名）。
    :type speaker_name: str
    :param voices_info: 利用可能な話者情報のリスト。Noneの場合、 ``get_available_voices_info`` を呼び出します。
    :type voices_info: list[dict] | None
    :param model: 話者情報を取得するために使用されるQwen3TTSModelインスタンス。
                  ``voices_info`` が提供されている場合は使用されません。
    :type model: Qwen3TTSModel | None
    :returns: (更新されたvoices_info, 解決された話者名) のタプル。
    :rtype: tuple[list[dict], str]
    :raises ValueError: 話者名が空であるか、見つからない場合。
    """
    if voices_info is None:
        voices_info = get_available_voices_info(model)

    query = _speaker_key(speaker_name)
    if not query:
        raise ValueError("話者名が空です")

    # Prefer an exact normalized match before accepting a partial match.
    for voice in voices_info:
        if query == _speaker_key(voice["name"]):
            return voices_info, voice["name"]
    for voice in voices_info:
        if query in _speaker_key(voice["name"]):
            return voices_info, voice["name"]

    raise ValueError(
        "❌ Error in tktts_qwen3.resolve_speaker_id(): "
        f"話者 [{speaker_name}] が見つかりませんでした"
    )


def speak(
    outfile,
    text,
    voice=DEFAULT_QWEN3_VOICE,
    speak_rate=None,
    speak_pitch=None,
    *,
    language: str = DEFAULT_LANGUAGE,
    model: Qwen3TTSModel | None = None,
    model_id: str = DEFAULT_QWEN3_MODEL,
    device: str = DEFAULT_DEVICE,
    dtype: str | torch.dtype = DEFAULT_DTYPE,
    instruct: str | None = None,
):
    """Qwen3-TTS CustomVoiceモデルを使用して単一の音声ファイルを生成します。

    ``speak_rate`` と ``speak_pitch`` は ``tktts_voicevox`` とのインターフェース互換性のために
    残されていますが、CustomVoiceモデルはこれらのパラメータを公開していないため、
    非デフォルト値は現在無視されます。

    :param outfile: 生成された音声ファイルを保存するパス。
    :type outfile: str | PathLike
    :param text: 読み上げるテキスト。
    :type text: str
    :param voice: 使用する話者名。
    :type voice: str
    :param speak_rate: 音声の速さ（現在は無視されます）。
    :type speak_rate: float | None
    :param speak_pitch: 音声のピッチ（現在は無視されます）。
    :type speak_pitch: float | None
    :param language: 読み上げに使用する言語（例: "Japanese"）。
    :type language: str
    :param model: 既存のQwen3TTSModelインスタンス。Noneの場合、 ``load_model`` を呼び出します。
    :type model: Qwen3TTSModel | None
    :param model_id: ``model`` がNoneの場合にロードするモデルID。
    :type model_id: str
    :param device: ``model`` がNoneの場合にモデルをロードするデバイス。
    :type device: str
    :param dtype: ``model`` がNoneの場合にモデルをロードするデータ型。
    :type dtype: str | torch.dtype
    :param instruct: 音声生成のための追加の指示テキスト。
    :type instruct: str | None
    :returns: 生成されたファイルのパス、または失敗した場合はNone。
    :rtype: str | None
    """
    text = str(text).strip()
    if not text:
        print("❌ tktts_qwen3.speak(): 読み上げテキストが空です")
        return None

    if speak_rate is not None or speak_pitch is not None:
        print(
            "    ** Warning: Qwen3-TTS CustomVoiceでは "
            "speak_rate/speak_pitchを直接指定できないため無視します"
        )

    if model is None:
        model = load_model(model_id=model_id, device=device, dtype=dtype)

    kwargs: dict[str, Any] = {
        "text": text,
        "language": language,
        "speaker": voice,
    }
    if instruct:
        kwargs["instruct"] = instruct

    try:
        with torch.inference_mode():
            wavs, sample_rate = model.generate_custom_voice(**kwargs)
    except Exception as exc:
        print(f"❌ Qwen3-TTS生成エラー: {exc}")
        return None

    output_path = Path(outfile)
    output_path.parent.mkdir(parents=True, exist_ok=True)
    sf.write(str(output_path), wavs[0], sample_rate)

    if output_path.exists():
        print(f"    ** 一時ファイル [{output_path}] を保存しました")
        return str(output_path)

    print(f"    ** Error: ファイル [{output_path}] の出力に失敗しました")
    return None


def speak_dialogue(
    dialogue,
    replacements,
    target_voices,
    speakers=None,
    temp_dir=".",
    outfile=None,
    ext="wav",
    cfg=None,
    *,
    language: str = DEFAULT_LANGUAGE,
    model: Qwen3TTSModel | None = None,
    model_id: str = DEFAULT_QWEN3_MODEL,
    device: str = DEFAULT_DEVICE,
    dtype: str | torch.dtype = DEFAULT_DTYPE,
):
    """tktts対話シーケンス用の一時オーディオファイルを生成します。

    この関数は、与えられた対話リストを処理し、各セグメントに対して個別の音声ファイルを生成します。
    ``outfile`` は ``tktts_voicevox`` との互換性のために署名に残されていますが、この関数では使用されません。

    :param dialogue: 対話アイテムのリスト。各アイテムはテキストまたは話者とテキストの組み合わせ。
    :type dialogue: list
    :param replacements: テキストに適用される置換規則の辞書。
    :type replacements: dict
    :param target_voices: 対話で話す話者の名前、または話者のリスト。
    :type target_voices: str | list
    :param speakers: 話者名と話者IDのマッピング辞書（オプション）。
    :type speakers: dict | None
    :param temp_dir: 一時音声ファイルを保存するディレクトリ。
    :type temp_dir: str
    :param outfile: この関数では使用されません（互換性のため残されています）。
    :type outfile: Any
    :param ext: 生成される音声ファイルの拡張子（例: "wav"）。
    :type ext: str
    :param cfg: 設定オブジェクト（ ``fspeak_rate``, ``fspeak_pitch``, ``qwen3_instruct``, ``monologue`` などの属性を含む場合があります）。
    :type cfg: Any | None
    :param language: 音声生成に使用する言語（例: "Japanese"）。
    :type language: str
    :param model: 既存のQwen3TTSModelインスタンス。Noneの場合、 ``load_model`` を呼び出します。
    :type model: Qwen3TTSModel | None
    :param model_id: ``model`` がNoneの場合にロードするモデルID。
    :type model_id: str
    :param device: ``model`` がNoneの場合にモデルをロードするデバイス。
    :type device: str
    :param dtype: ``model`` がNoneの場合にモデルをロードするデータ型。
    :type dtype: str | torch.dtype
    :returns: (成功したかどうかのブール値, 生成された一時ファイルのパスのリスト) のタプル。
    :rtype: tuple[bool, list[str]]
    """
    del outfile  # Kept in the signature for compatibility with tktts_voicevox.
    speakers = {} if speakers is None else speakers

    print("tktts_qwen3.speak_dialogue(): target_voices:", target_voices)

    if model is None:
        try:
            model = load_model(model_id=model_id, device=device, dtype=dtype)
        except Exception as exc:
            print(f"❌ Qwen3-TTSモデル読み込みエラー: {exc}")
            return False, []

    tmpfiles = []
    voices_info = get_available_voices_info(model)
    idx = 1
    is_monologue = bool(getattr(cfg, "monologue", False))

    for i, dialogue_item in enumerate(dialogue):
        print()
        print(f"Dialogue {i:04d}:")
        dialogue_list = split_dialogue(
            dialogue_item,
            target_voices,
            speakers=speakers,
            default_voice=DEFAULT_QWEN3_VOICE,
            is_monologue=is_monologue,
        )

        for speaker, text in dialogue_list:
            tmpfile = os.path.join(temp_dir, f"tmp_{idx:03d}.{ext}")
            text = apply_replacements(text, replacements)
            if isinstance(target_voices, str):
                speaker = target_voices

            try:
                voices_info, target_voice = resolve_speaker_id(
                    speaker,
                    voices_info=voices_info,
                    model=model,
                )
            except ValueError as exc:
                print(exc)
                return False, tmpfiles

            print(f"  {idx:04d}: voice={speaker} ({target_voice}): {text}")

            speak_rate = getattr(cfg, "fspeak_rate", None)
            speak_pitch = getattr(cfg, "fspeak_pitch", None)
            instruct = getattr(cfg, "qwen3_instruct", None)

            generated = speak(
                outfile=tmpfile,
                text=text,
                voice=target_voice,
                speak_rate=speak_rate,
                speak_pitch=speak_pitch,
                language=language,
                model=model,
                instruct=instruct,
            )
            if generated is None:
                return False, tmpfiles

            tmpfiles.append(tmpfile)
            idx += 1

    return True, tmpfiles


if __name__ == "__main__":
    list_available_voices()
