list_models.py ダウンロード/コピー

list_models.py をダウンロード

list_models.py
list_models.py
  1#!/usr/bin/env python3
  2"""LiteLLMを使ってOpenAI/Geminiのモデル一覧取得と応答テストを行うモジュール。
  3
  4このモジュールはLiteLLMライブラリを利用し、OpenAIおよびGoogle GeminiのAIモデルに対して
  5以下の操作を実行します。
  6
  7- 利用可能なモデルの一覧を取得し、コンソールに表示する。
  8- 指定されたモデルに対して接続テストを行い、応答を表示する。
  9
 10APIキーの管理には、`tkai_lib` または `tkai_lib_litellm` ライブラリ(もし利用可能であれば)
 11または環境変数を使用します。特に、Google AI StudioのAPIキーは `GOOGLE_API_KEY` または
 12`GEMINI_API_KEY` のどちらかに設定することで利用可能です。
 13
 14使用例::
 15
 16    python list_litellm_models.py mode=list
 17    python list_litellm_models.py mode=list api=openai
 18    python list_litellm_models.py mode=test model=gpt-4o-mini
 19    python list_litellm_models.py mode=test model=gemini-2.0-flash
 20    python list_litellm_models.py --mode test --model openai/gpt-4o-mini
 21
 22.. seealso::
 23   :doc:`list_models_usage`
 24
 25"""
 26
 27from __future__ import annotations
 28
 29import argparse
 30import os
 31import sys
 32from collections.abc import Sequence
 33from typing import Any
 34
 35from litellm import completion, get_valid_models
 36
 37try:
 38    from tkai_lib_litellm import read_ai_config
 39except ImportError:
 40    try:
 41        from tkai_lib import read_ai_config
 42    except ImportError:
 43        read_ai_config = None  # type: ignore[assignment]
 44
 45
 46PROVIDERS = ("openai", "gemini")
 47DEFAULT_TEST_PROMPT = (
 48    "これはAPI接続テストです。日本語で、あなたのモデル名を名乗らずに、"
 49    "『接続テストに成功しました』と短く回答してください。"
 50)
 51
 52
 53def parse_key_value_args(argv: Sequence[str]) -> list[str]:
 54    """mode=list のような形式のコマンドライン引数をargparse互換形式へ変換する。
 55
 56    特定のキーと値のペアを `--key value` の形式に変換します。
 57    これにより、ユーザーは `mode=list` のような簡潔な記法と、
 58    標準的な `--mode list` の両方を使用できます。
 59
 60    :param argv: コマンドライン引数のリスト。
 61    :type argv: Sequence[str]
 62    :returns: 変換されたコマンドライン引数のリスト。
 63    :rtype: list[str]
 64    """
 65    supported = {"mode", "api", "model", "prompt", "config"}
 66    converted: list[str] = []
 67
 68    for arg in argv:
 69        if "=" in arg and not arg.startswith("--"):
 70            key, value = arg.split("=", 1)
 71            if key in supported:
 72                converted.extend((f"--{key}", value))
 73                continue
 74        converted.append(arg)
 75
 76    return converted
 77
 78
 79def parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace:
 80    """コマンドライン引数をパースする。
 81
 82    スクリプトの実行モード(モデル一覧表示またはテスト)と、
 83    それに伴うオプション(プロバイダー、モデル名、プロンプト、設定ファイル)を
 84    定義し、パースします。
 85
 86    :param argv: コマンドライン引数のリスト。None の場合は sys.argv[1:] を使用。
 87    :type argv: Sequence[str] | None
 88    :returns: パースされた引数を格納するNamespaceオブジェクト。
 89    :rtype: argparse.Namespace
 90    """
 91    parser = argparse.ArgumentParser(
 92        description="LiteLLMでOpenAI/Geminiのモデル一覧取得と応答テストを行います。"
 93    )
 94    parser.add_argument(
 95        "--mode",
 96        choices=("list", "test"),
 97        default="list",
 98        help="list: モデル一覧、test: 指定モデルへテスト送信",
 99    )
100    parser.add_argument(
101        "--api",
102        choices=("all", *PROVIDERS),
103        default="all",
104        help="listモードで対象にするプロバイダー。既定値: all",
105    )
106    parser.add_argument(
107        "--model",
108        help="testモードで使用するモデル名。例: gpt-4o-mini, gemini-2.0-flash",
109    )
110    parser.add_argument(
111        "--prompt",
112        default=DEFAULT_TEST_PROMPT,
113        help="testモードで送信するプロンプト",
114    )
115    parser.add_argument(
116        "--config",
117        default="ai.env",
118        help="read_ai_config() に渡す設定ファイル。既定値: ai.env",
119    )
120    return parser.parse_args(parse_key_value_args(argv or sys.argv[1:]))
121
122
123def load_api_keys(config_path: str) -> None:
124    """添付ライブラリを使ってAPIキー設定を読み込む。
125
126    `tkai_lib_litellm` または `tkai_lib` が利用可能な場合、指定された設定ファイルから
127    APIキーを環境変数に読み込みます。これらのライブラリが見つからない場合は、
128    既存の環境変数がそのまま使用されます。
129
130    また、`GOOGLE_API_KEY` が設定されていれば、LiteLLMがGeminiモデルで使用できるよう
131    `GEMINI_API_KEY` にも設定をコピーします。
132
133    :param config_path: APIキー設定ファイルのパス。
134    :type config_path: str
135    :returns: なし
136    :rtype: None
137    """
138    if read_ai_config is None:
139        print(
140            "Warning: tkai_lib_litellm または tkai_lib が見つからないため、"
141            "現在の環境変数だけを使用します。",
142            file=sys.stderr,
143        )
144    else:
145        read_ai_config(config_path)
146
147    # 既存環境の GOOGLE_API_KEY をLiteLLM用にも利用する。
148    if not os.getenv("GEMINI_API_KEY") and os.getenv("GOOGLE_API_KEY"):
149        os.environ["GEMINI_API_KEY"] = os.environ["GOOGLE_API_KEY"]
150
151
152def required_key_exists(provider: str) -> bool:
153    """指定されたプロバイダーに必要なAPIキーが環境変数に存在するかを確認する。
154
155    - 'openai' の場合: `OPENAI_API_KEY` の有無をチェックします。
156    - 'gemini' の場合: `GEMINI_API_KEY` または `GOOGLE_API_KEY` のいずれかの有無をチェックします。
157
158    :param provider: チェックするプロバイダー名 ('openai' または 'gemini')。
159    :type provider: str
160    :returns: 必要なAPIキーが存在すれば True、そうでなければ False。
161    :rtype: bool
162    :raises ValueError: 未知のプロバイダーが指定された場合。
163    """
164    if provider == "openai":
165        return bool(os.getenv("OPENAI_API_KEY"))
166    if provider == "gemini":
167        return bool(os.getenv("GEMINI_API_KEY") or os.getenv("GOOGLE_API_KEY"))
168    raise ValueError(f"Unknown provider: {provider}")
169
170
171def list_provider_models(provider: str) -> list[str]:
172    """プロバイダーAPIへ問い合わせ、利用可能モデル名を返す。
173
174    LiteLLMの `get_valid_models` 関数を使用して、指定されたプロバイダーが提供する
175    利用可能なモデルのリストを取得します。エンドポイントの有効性も確認されます。
176
177    :param provider: モデル一覧を取得するプロバイダー名。
178    :type provider: str
179    :returns: 利用可能なモデル名のソート済みリスト。
180    :rtype: list[str]
181    """
182    models = get_valid_models(
183        check_provider_endpoint=True,
184        custom_llm_provider=provider,
185    )
186    return sorted({str(model) for model in models}, key=str.casefold)
187
188
189def print_provider_models(provider: str) -> bool:
190    """指定されたプロバイダーの利用可能モデル一覧をコンソールに表示する。
191
192    まず、必要なAPIキーが設定されているかを確認します。
193    キーが設定されていれば、`list_provider_models` を呼び出してモデル一覧を取得し、表示します。
194    エラーが発生した場合やキーが未設定の場合は、エラーメッセージを表示します。
195
196    :param provider: モデル一覧を表示するプロバイダー名。
197    :type provider: str
198    :returns: モデル一覧の取得と表示が成功した場合は True、失敗した場合は False。
199    :rtype: bool
200    """
201    title = "OpenAI" if provider == "openai" else "Gemini (Google AI)"
202    print(f"\n--- {title}: 利用可能モデル一覧 ---")
203
204    if not required_key_exists(provider):
205        key_name = (
206            "OPENAI_API_KEY"
207            if provider == "openai"
208            else "GOOGLE_API_KEY / GEMINI_API_KEY"
209        )
210        print(f"SKIP: {key_name} が設定されていません", file=sys.stderr)
211        return False
212
213    try:
214        models = list_provider_models(provider)
215    except Exception as exc:
216        print(f"ERROR: {title} のモデル一覧取得に失敗しました", file=sys.stderr)
217        print(f"       [{type(exc).__name__}] {exc}", file=sys.stderr)
218        return False
219
220    for model in models:
221        print(f"ID: {model}")
222
223    print(f"--- 完了: {len(models)} 件 ---")
224    return True
225
226
227def normalize_model_name(model: str) -> tuple[str, str]:
228    """モデル名からプロバイダーを判定し、LiteLLM形式へ正規化する。
229
230    入力されたモデル名に基づいてプロバイダーを識別し、LiteLLMが認識する形式
231    (例: 'openai/gpt-4o-mini', 'gemini/gemini-2.0-flash') に変換します。
232    モデル名に `/` が含まれる場合はそれを使用してプロバイダーを判断し、
233    そうでなければ一般的なプレフィックス('gemini' など)で判断します。
234    Gemini以外のモデルはデフォルトでOpenAIとして扱われます。
235
236    :param model: ユーザーが指定したモデル名。
237    :type model: str
238    :returns: (プロバイダー名, LiteLLM形式のモデル名) のタプル。
239    :rtype: tuple[str, str]
240    :raises ValueError: モデル名が空の場合、または対応していないプロバイダーが指定された場合。
241    """
242    normalized = model.strip()
243    if not normalized:
244        raise ValueError("model が空です")
245
246    if "/" in normalized:
247        provider = normalized.split("/", 1)[0].lower()
248        if provider == "google":  # 'google/gemini-pro' のようなケースを 'gemini/gemini-pro' に変換
249            provider = "gemini"
250            normalized = "gemini/" + normalized.split("/", 1)[1]
251        if provider not in PROVIDERS:
252            raise ValueError(
253                f"対応していないプロバイダーです: {provider} "
254                f"(対応: {', '.join(PROVIDERS)})"
255            )
256        return provider, normalized
257
258    lower_name = normalized.lower()
259    if lower_name.startswith("gemini"):
260        return "gemini", f"gemini/{normalized}"
261
262    # OpenAIモデルは gpt-, o1, o3, o4 など複数の接頭辞があるため、
263    # Gemini以外はOpenAIとして扱う。
264    return "openai", f"openai/{normalized}"
265
266
267def extract_response_text(response: Any) -> str:
268    """LiteLLMの応答オブジェクトからメッセージ本文を取り出す。
269
270    LiteLLMの `completion` 関数が返す応答オブジェクトから、
271    AIモデルの生成したメッセージのテキストコンテンツを抽出します。
272    応答の構造が異なる可能性に対応するため、複数のアクセス方法を試行します。
273
274    :param response: LiteLLMの `completion` 関数が返した応答オブジェクト。
275    :type response: Any
276    :returns: 抽出されたメッセージ本文の文字列。本文がNoneの場合は空文字列。
277    :rtype: str
278    :raises ValueError: 応答オブジェクトから本文を抽出できなかった場合。
279    """
280    try:
281        content = response.choices[0].message.content
282    except (AttributeError, IndexError, TypeError):
283        try:
284            content = response["choices"][0]["message"]["content"]
285        except (KeyError, IndexError, TypeError) as exc:
286            raise ValueError("応答本文を取得できませんでした") from exc
287
288    if content is None:
289        return ""
290    if isinstance(content, str):
291        return content
292    return str(content)
293
294
295def test_model(model: str, prompt: str) -> bool:
296    """指定モデルへプロンプトを送信し、回答を表示する。
297
298    まず、モデル名を正規化し、必要なAPIキーが設定されているかを確認します。
299    その後、LiteLLMの `completion` 関数を使用して指定されたプロンプトをモデルに送信し、
300    得られた応答のテキストコンテンツをコンソールに表示します。
301    APIキーの不足、モデルの認識失敗、またはAPI呼び出し中のエラーが発生した場合は、
302    適切なエラーメッセージを表示します。
303
304    :param model: テスト対象のモデル名。
305    :type model: str
306    :param prompt: モデルに送信するプロンプト。
307    :type prompt: str
308    :returns: テストが成功した場合は True、失敗した場合は False。
309    :rtype: bool
310    """
311    try:
312        provider, litellm_model = normalize_model_name(model)
313    except ValueError as exc:
314        print(f"ERROR: {exc}", file=sys.stderr)
315        return False
316
317    if not required_key_exists(provider):
318        key_name = (
319            "OPENAI_API_KEY"
320            if provider == "openai"
321            else "GOOGLE_API_KEY / GEMINI_API_KEY"
322        )
323        print(f"ERROR: {key_name} が設定されていません", file=sys.stderr)
324        return False
325
326    print("--- モデル応答テスト ---")
327    print(f"Provider : {provider}")
328    print(f"Model    : {litellm_model}")
329    print(f"Prompt   : {prompt}")
330    print("--- 応答 ---")
331
332    try:
333        response = completion(
334            model=litellm_model,
335            messages=[{"role": "user", "content": prompt}],
336        )
337        text = extract_response_text(response)
338    except Exception as exc:
339        print("ERROR: モデルからの回答取得に失敗しました", file=sys.stderr)
340        print(f"       [{type(exc).__name__}] {exc}", file=sys.stderr)
341        print(
342            "       APIキー、モデル名、利用権限、クォータを確認してください。",
343            file=sys.stderr,
344        )
345        return False
346
347    print(text if text else "(空の応答)")
348    print("--- テスト完了 ---")
349    return True
350
351
352def main(argv: Sequence[str] | None = None) -> int:
353    """スクリプトのメイン処理を実行する。
354
355    コマンドライン引数をパースし、APIキーを読み込み、指定されたモードに応じて
356    モデル一覧表示またはモデル応答テストを実行します。
357
358    - `mode=test` の場合、`--model` オプションが必須です。
359    - `mode=list` の場合、`--api` オプションで対象プロバイダーを選択できます。
360
361    :param argv: コマンドライン引数のリスト。None の場合は `sys.argv[1:]` を使用。
362    :type argv: Sequence[str] | None
363    :returns: 終了コード (0: 成功、1: 失敗、2: 不適切な引数)。
364    :rtype: int
365    """
366    args = parse_args(argv)
367    load_api_keys(args.config)
368
369    if args.mode == "test":
370        if not args.model:
371            print(
372                "ERROR: mode=test では model=XXX の指定が必要です。\n"
373                "例: python list_litellm_models.py mode=test model=gpt-4o-mini",
374                file=sys.stderr,
375            )
376            return 2
377        return 0 if test_model(args.model, args.prompt) else 1
378
379    selected = PROVIDERS if args.api == "all" else (args.api,)
380    results = [print_provider_models(provider) for provider in selected]
381    return 0 if all(results) else 1
382
383
384if __name__ == "__main__":
385    sys.exit(main())