156 lines
6.1 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""将 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"]