156 lines
6.1 KiB
Python
156 lines
6.1 KiB
Python
"""将 Pi 的 provider catalog 投影为 DSH pi-ai profile。"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import json
|
||
import os
|
||
import re
|
||
from dataclasses import dataclass
|
||
from pathlib import Path
|
||
from typing import Any, Mapping
|
||
|
||
|
||
_ENV_NAME = re.compile(r"^[A-Z_][A-Z0-9_]*$")
|
||
|
||
|
||
class PiAiRouteConfigError(ValueError):
|
||
"""Pi 配置缺失、格式非法或不能描述一个 DSH 路由。"""
|
||
|
||
|
||
@dataclass(frozen=True)
|
||
class PiAiRouteConfig:
|
||
"""不含密钥的 DSH pi-ai 路由配置。"""
|
||
|
||
provider: str
|
||
model: str
|
||
api: str
|
||
base_url: str
|
||
api_key_env: str = "DEEPSEEK_API_KEY"
|
||
model_name: str | None = None
|
||
context_window: int | None = None
|
||
max_tokens: int | None = None
|
||
reasoning_efforts: Mapping[str, str | None] | None = None
|
||
|
||
def __post_init__(self) -> None:
|
||
for field, value in (
|
||
("provider", self.provider),
|
||
("model", self.model),
|
||
("api", self.api),
|
||
("base_url", self.base_url),
|
||
):
|
||
if not isinstance(value, str) or not value.strip():
|
||
raise PiAiRouteConfigError(f"{field} 必须是非空字符串")
|
||
if not isinstance(self.api_key_env, str) or _ENV_NAME.fullmatch(self.api_key_env) is None:
|
||
raise PiAiRouteConfigError("api_key_env 必须是大写 POSIX 环境变量名")
|
||
for field, value in (("context_window", self.context_window), ("max_tokens", self.max_tokens)):
|
||
if value is not None and (isinstance(value, bool) or not isinstance(value, int) or value <= 0):
|
||
raise PiAiRouteConfigError(f"{field} 必须是正整数")
|
||
|
||
def as_profile(self) -> dict[str, Any]:
|
||
"""生成 DSH `llm-pi-ai.providers.<route>` 配置,不携带 secret。"""
|
||
|
||
model: dict[str, Any] = {"id": self.model}
|
||
if self.model_name:
|
||
model["name"] = self.model_name
|
||
if self.context_window is not None:
|
||
model["contextWindow"] = self.context_window
|
||
if self.max_tokens is not None:
|
||
model["maxTokens"] = self.max_tokens
|
||
if self.reasoning_efforts:
|
||
model["reasoningEfforts"] = dict(self.reasoning_efforts)
|
||
return {
|
||
"apiKeyEnv": self.api_key_env,
|
||
"api": self.api,
|
||
"baseURL": self.base_url,
|
||
"models": [model],
|
||
}
|
||
|
||
|
||
def _default_models_path() -> Path:
|
||
agent_dir = os.environ.get("PI_CODING_AGENT_DIR")
|
||
return Path(agent_dir).expanduser() / "models.json" if agent_dir else Path.home() / ".pi" / "agent" / "models.json"
|
||
|
||
|
||
def _as_mapping(value: Any, field: str) -> Mapping[str, Any]:
|
||
if not isinstance(value, Mapping):
|
||
raise PiAiRouteConfigError(f"{field} 必须是对象")
|
||
return value
|
||
|
||
|
||
def _positive_int(value: Any, field: str) -> int | None:
|
||
if value is None:
|
||
return None
|
||
if isinstance(value, bool) or not isinstance(value, int) or value <= 0:
|
||
raise PiAiRouteConfigError(f"{field} 必须是正整数")
|
||
return value
|
||
|
||
|
||
def load_pi_ai_route(
|
||
models_path: str | Path | None = None,
|
||
*,
|
||
provider: str | None = None,
|
||
model: str | None = None,
|
||
api_key_env: str = "DEEPSEEK_API_KEY",
|
||
) -> PiAiRouteConfig:
|
||
"""读取 Pi 的非敏感 provider/model 元数据,返回 DSH 路由配置。
|
||
|
||
密钥不从 `auth.json` 读取;调用方应在当前进程把它放入 `api_key_env`。
|
||
"""
|
||
|
||
path = Path(models_path).expanduser() if models_path else _default_models_path()
|
||
try:
|
||
raw = json.loads(path.read_text(encoding="utf-8"))
|
||
except (OSError, UnicodeError, json.JSONDecodeError) as exc:
|
||
raise PiAiRouteConfigError(f"Pi models.json 不可读: {path}") from exc
|
||
if not isinstance(raw, Mapping):
|
||
raise PiAiRouteConfigError("Pi models.json 顶层必须是对象")
|
||
providers = _as_mapping(raw.get("providers"), "providers")
|
||
selected_provider = provider or os.environ.get("PI_PROVIDER")
|
||
selected_model = model or os.environ.get("PI_MODEL")
|
||
if not selected_provider:
|
||
raise PiAiRouteConfigError("未指定 provider,且 PI_PROVIDER 未配置")
|
||
if not selected_model:
|
||
raise PiAiRouteConfigError("未指定 model,且 PI_MODEL 未配置")
|
||
provider_config = _as_mapping(providers.get(selected_provider), f"provider {selected_provider}")
|
||
base_url = provider_config.get("baseUrl")
|
||
api = provider_config.get("api")
|
||
if not isinstance(base_url, str) or not base_url.strip():
|
||
raise PiAiRouteConfigError(f"provider {selected_provider} 缺少 baseUrl")
|
||
if not isinstance(api, str) or not api.strip():
|
||
raise PiAiRouteConfigError(f"provider {selected_provider} 缺少 api 协议")
|
||
models = provider_config.get("models")
|
||
if not isinstance(models, list):
|
||
raise PiAiRouteConfigError(f"provider {selected_provider} 的 models 必须是数组")
|
||
model_config = next(
|
||
(item for item in models if isinstance(item, Mapping) and item.get("id") == selected_model),
|
||
None,
|
||
)
|
||
if model_config is None:
|
||
raise PiAiRouteConfigError(f"model {selected_model} 未登记在 provider {selected_provider}")
|
||
thinking_map = model_config.get("thinkingLevelMap")
|
||
reasoning_efforts = None
|
||
if isinstance(thinking_map, Mapping):
|
||
reasoning_efforts = {}
|
||
for level, value in thinking_map.items():
|
||
level_name = str(level)
|
||
# Pi 用 null 表示该档位不支持;DSH 只允许 off 以 null 声明。
|
||
if value is None and level_name != "off":
|
||
continue
|
||
reasoning_efforts[level_name] = (
|
||
value if value is None or isinstance(value, str) else str(value)
|
||
)
|
||
return PiAiRouteConfig(
|
||
provider=str(selected_provider),
|
||
model=str(selected_model),
|
||
api=str(api),
|
||
base_url=base_url.rstrip("/"),
|
||
api_key_env=api_key_env,
|
||
model_name=str(model_config.get("name")) if model_config.get("name") else None,
|
||
context_window=_positive_int(model_config.get("contextWindow"), "contextWindow"),
|
||
max_tokens=_positive_int(model_config.get("maxTokens"), "maxTokens"),
|
||
reasoning_efforts=reasoning_efforts,
|
||
)
|
||
|
||
|
||
__all__ = ["PiAiRouteConfig", "PiAiRouteConfigError", "load_pi_ai_route"]
|