#!/usr/bin/env python3
"""
概要:
    LM Studio API sample (Python standard library only).
詳細説明:
    LM Studioのモデル一覧取得とチャットAPIのサンプルスクリプトです。
    python lmstudio_api_sample.py --list のように実行してモデル一覧を取得したり、
    python lmstudio_api_sample.py --model モデルID --prompt 質問 のようにチャットを行います。
関連リンク:
    lmstudio_api_sample_usage
"""
import argparse
import json
import sys
from urllib.error import HTTPError, URLError
from urllib.request import Request, ProxyHandler, build_opener


def request_json(base_url, endpoint, timeout, payload=None):
    """
    概要:
        JSON APIにHTTPリクエストを送信し、応答をパースして返します。
    詳細説明:
        システムプロキシをバイパスして直接サーバーに接続します。
        payloadが指定された場合はJSONとしてエンコードし、POSTリクエストを送信します。
    引数:
        :param base_url: APIのベースURL
        :type base_url: str
        :param endpoint: リクエスト先のエンドポイント
        :type endpoint: str
        :param timeout: 通信タイムアウト秒
        :type timeout: float
        :param payload: リクエストボディとして送信するデータオブジェクト
        :type payload: dict or None
    戻り値:
        :returns: 応答のJSON文字列をパースしたオブジェクト
        :rtype: dict or list
    """
    data = None if payload is None else json.dumps(payload, ensure_ascii=False).encode("utf-8")
    request = Request(
        base_url.rstrip("/") + endpoint,
        data=data,
        headers={"Content-Type": "application/json", "Accept": "application/json"},
    )
    # Connect directly to the private server, bypassing system HTTP proxies.
    with build_opener(ProxyHandler({})).open(request, timeout=timeout) as response:
        return json.load(response)


def main():
    """
    概要:
        コマンドライン引数を解析し、APIへのリクエストを実行します。
    詳細説明:
        引数に応じてモデル一覧の表示、またはチャットAPIへの問い合わせを行います。
        ネットワークエラーやパースエラーが発生した場合はエラーメッセージを標準エラー出力に書き出します。
    戻り値:
        :returns: プロセスの終了コード
        :rtype: int
    """
    parser = argparse.ArgumentParser(description="LM Studioのモデル一覧取得とチャットAPIのサンプル")
    parser.add_argument("--list", action="store_true", help="利用可能なモデルIDを表示（embeddingモデルも含む）")
    parser.add_argument("--model", help="問い合わせるモデルID")
    parser.add_argument("--prompt", help="モデルへの質問")
    parser.add_argument("--base-url", default="http://192.168.27.18:1234/v1", help="APIのベースURL")
    parser.add_argument("--timeout", type=float, default=300, help="通信タイムアウト秒（既定: 300）")
    parser.add_argument("--max-tokens", type=int, default=2048, help="最大生成トークン数（推論トークンを含む）")
    parser.add_argument("--show-reasoning", action="store_true", help="reasoning_contentも表示")
    args = parser.parse_args()
    if (args.model is None) != (args.prompt is None):
        parser.error("--modelと--promptは両方指定してください")
    if not args.list and args.model is None:
        parser.error("--list、または--modelと--promptを指定してください")
    if args.timeout <= 0 or args.max_tokens <= 0:
        parser.error("--timeoutと--max-tokensは正の値を指定してください")
    try:
        if args.list:
            result = request_json(args.base_url, "/models", args.timeout)
            models = result["data"]
            if not models:
                print("利用可能なモデルはありません。")
            for model in models:
                print(model["id"])
        if args.model is not None:
            result = request_json(args.base_url, "/chat/completions", args.timeout, {
                "model": args.model,
                "messages": [{"role": "user", "content": args.prompt}],
                "max_tokens": args.max_tokens,
                "stream": False,
            })
            choice = result["choices"][0]
            message = choice["message"]
            if args.show_reasoning and message.get("reasoning_content"):
                print("[reasoning]\n" + message["reasoning_content"] + "\n[answer]")
            print(message.get("content") or "")
            if choice.get("finish_reason") == "length":
                print("生成上限に達しました。必要なら--max-tokensを増やしてください。", file=sys.stderr)
        return 0
    except HTTPError as exc:
        print(f"HTTPエラー {exc.code}: {exc.read().decode('utf-8', errors='replace')}", file=sys.stderr)
    except (URLError, TimeoutError, OSError) as exc:
        print(f"接続エラー: {exc}", file=sys.stderr)
    except (ValueError, KeyError, IndexError, TypeError) as exc:
        print(f"API応答を解釈できません: {exc}", file=sys.stderr)
    return 1


if __name__ == "__main__":
    # Preserve Japanese output in Windows terminals and redirected output.
    for stream in (sys.stdout, sys.stderr):
        if hasattr(stream, "reconfigure"):
            stream.reconfigure(encoding="utf-8")
    sys.exit(main())