#!/usr/bin/env python3
"""LiteLLMを使ってOpenAI/Geminiのモデル一覧取得と応答テストを行う。

使用例:
    python list_litellm_models.py mode=list
    python list_litellm_models.py mode=list api=openai
    python list_litellm_models.py mode=test model=gpt-4o-mini
    python list_litellm_models.py mode=test model=gemini-2.0-flash
    python list_litellm_models.py --mode test --model openai/gpt-4o-mini
"""

from __future__ import annotations

import argparse
import os
import sys
from collections.abc import Sequence
from typing import Any

from litellm import completion, get_valid_models

try:
    from tkai_lib_litellm import read_ai_config
except ImportError:
    try:
        from tkai_lib import read_ai_config
    except ImportError:
        read_ai_config = None  # type: ignore[assignment]


PROVIDERS = ("openai", "gemini")
DEFAULT_TEST_PROMPT = (
    "これはAPI接続テストです。日本語で、あなたのモデル名を名乗らずに、"
    "『接続テストに成功しました』と短く回答してください。"
)


def parse_key_value_args(argv: Sequence[str]) -> list[str]:
    """mode=list のような指定を argparse 形式へ変換する。"""
    supported = {"mode", "api", "model", "prompt", "config"}
    converted: list[str] = []

    for arg in argv:
        if "=" in arg and not arg.startswith("--"):
            key, value = arg.split("=", 1)
            if key in supported:
                converted.extend((f"--{key}", value))
                continue
        converted.append(arg)

    return converted


def parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace:
    parser = argparse.ArgumentParser(
        description="LiteLLMでOpenAI/Geminiのモデル一覧取得と応答テストを行います。"
    )
    parser.add_argument(
        "--mode",
        choices=("list", "test"),
        default="list",
        help="list: モデル一覧、test: 指定モデルへテスト送信",
    )
    parser.add_argument(
        "--api",
        choices=("all", *PROVIDERS),
        default="all",
        help="listモードで対象にするプロバイダー。既定値: all",
    )
    parser.add_argument(
        "--model",
        help="testモードで使用するモデル名。例: gpt-4o-mini, gemini-2.0-flash",
    )
    parser.add_argument(
        "--prompt",
        default=DEFAULT_TEST_PROMPT,
        help="testモードで送信するプロンプト",
    )
    parser.add_argument(
        "--config",
        default="ai.env",
        help="read_ai_config() に渡す設定ファイル。既定値: ai.env",
    )
    return parser.parse_args(parse_key_value_args(argv or sys.argv[1:]))


def load_api_keys(config_path: str) -> None:
    """添付ライブラリを使ってAPIキー設定を読み込む。"""
    if read_ai_config is None:
        print(
            "Warning: tkai_lib_litellm または tkai_lib が見つからないため、"
            "現在の環境変数だけを使用します。",
            file=sys.stderr,
        )
    else:
        read_ai_config(config_path)

    # 既存環境の GOOGLE_API_KEY をLiteLLM用にも利用する。
    if not os.getenv("GEMINI_API_KEY") and os.getenv("GOOGLE_API_KEY"):
        os.environ["GEMINI_API_KEY"] = os.environ["GOOGLE_API_KEY"]


def required_key_exists(provider: str) -> bool:
    if provider == "openai":
        return bool(os.getenv("OPENAI_API_KEY"))
    if provider == "gemini":
        return bool(os.getenv("GEMINI_API_KEY") or os.getenv("GOOGLE_API_KEY"))
    raise ValueError(f"Unknown provider: {provider}")


def list_provider_models(provider: str) -> list[str]:
    """プロバイダーAPIへ問い合わせ、利用可能モデル名を返す。"""
    models = get_valid_models(
        check_provider_endpoint=True,
        custom_llm_provider=provider,
    )
    return sorted({str(model) for model in models}, key=str.casefold)


def print_provider_models(provider: str) -> bool:
    title = "OpenAI" if provider == "openai" else "Gemini (Google AI)"
    print(f"\n--- {title}: 利用可能モデル一覧 ---")

    if not required_key_exists(provider):
        key_name = (
            "OPENAI_API_KEY"
            if provider == "openai"
            else "GOOGLE_API_KEY / GEMINI_API_KEY"
        )
        print(f"SKIP: {key_name} が設定されていません", file=sys.stderr)
        return False

    try:
        models = list_provider_models(provider)
    except Exception as exc:
        print(f"ERROR: {title} のモデル一覧取得に失敗しました", file=sys.stderr)
        print(f"       [{type(exc).__name__}] {exc}", file=sys.stderr)
        return False

    for model in models:
        print(f"ID: {model}")

    print(f"--- 完了: {len(models)} 件 ---")
    return True


def normalize_model_name(model: str) -> tuple[str, str]:
    """モデル名からプロバイダーを判定し、LiteLLM形式へ正規化する。"""
    normalized = model.strip()
    if not normalized:
        raise ValueError("model が空です")

    if "/" in normalized:
        provider = normalized.split("/", 1)[0].lower()
        if provider == "google":
            provider = "gemini"
            normalized = "gemini/" + normalized.split("/", 1)[1]
        if provider not in PROVIDERS:
            raise ValueError(
                f"対応していないプロバイダーです: {provider} "
                f"(対応: {', '.join(PROVIDERS)})"
            )
        return provider, normalized

    lower_name = normalized.lower()
    if lower_name.startswith("gemini"):
        return "gemini", f"gemini/{normalized}"

    # OpenAIモデルは gpt-, o1, o3, o4 など複数の接頭辞があるため、
    # Gemini以外はOpenAIとして扱う。
    return "openai", f"openai/{normalized}"


def extract_response_text(response: Any) -> str:
    """LiteLLMの応答オブジェクトから本文を取り出す。"""
    try:
        content = response.choices[0].message.content
    except (AttributeError, IndexError, TypeError):
        try:
            content = response["choices"][0]["message"]["content"]
        except (KeyError, IndexError, TypeError) as exc:
            raise ValueError("応答本文を取得できませんでした") from exc

    if content is None:
        return ""
    if isinstance(content, str):
        return content
    return str(content)


def test_model(model: str, prompt: str) -> bool:
    """指定モデルへプロンプトを送信し、回答を表示する。"""
    try:
        provider, litellm_model = normalize_model_name(model)
    except ValueError as exc:
        print(f"ERROR: {exc}", file=sys.stderr)
        return False

    if not required_key_exists(provider):
        key_name = (
            "OPENAI_API_KEY"
            if provider == "openai"
            else "GOOGLE_API_KEY / GEMINI_API_KEY"
        )
        print(f"ERROR: {key_name} が設定されていません", file=sys.stderr)
        return False

    print("--- モデル応答テスト ---")
    print(f"Provider : {provider}")
    print(f"Model    : {litellm_model}")
    print(f"Prompt   : {prompt}")
    print("--- 応答 ---")

    try:
        response = completion(
            model=litellm_model,
            messages=[{"role": "user", "content": prompt}],
        )
        text = extract_response_text(response)
    except Exception as exc:
        print("ERROR: モデルからの回答取得に失敗しました", file=sys.stderr)
        print(f"       [{type(exc).__name__}] {exc}", file=sys.stderr)
        print(
            "       APIキー、モデル名、利用権限、クォータを確認してください。",
            file=sys.stderr,
        )
        return False

    print(text if text else "(空の応答)")
    print("--- テスト完了 ---")
    return True


def main(argv: Sequence[str] | None = None) -> int:
    args = parse_args(argv)
    load_api_keys(args.config)

    if args.mode == "test":
        if not args.model:
            print(
                "ERROR: mode=test では model=XXX の指定が必要です。\n"
                "例: python list_litellm_models.py mode=test model=gpt-4o-mini",
                file=sys.stderr,
            )
            return 2
        return 0 if test_model(args.model, args.prompt) else 1

    selected = PROVIDERS if args.api == "all" else (args.api,)
    results = [print_provider_models(provider) for provider in selected]
    return 0 if all(results) else 1


if __name__ == "__main__":
    sys.exit(main())
