noke / core /model_config.py
i3t2y
feat: noke v0.3.0 — 5-batch architecture upgrade
5039d41
Raw
History Blame Contribute Delete
4.54 kB
"""
动态模型列表:从上游获取最新模型配置,本地硬编码作 fallback。
"""
import time
import logging
import re
import httpx
from core.noke_client import MODEL_MAP as _FALLBACK_MAP
logger = logging.getLogger("noke")
# 上游模型配置接口
from core.config import DEFAULT_UPSTREAM_URL
_MODELS_URL = f"{DEFAULT_UPSTREAM_URL}/proxy/v1/model_config/models?a=0"
# 缓存
_cache: dict[str, str] | None = None
_cache_at: float = 0
_cache_ttl: float = 600 # 10 分钟
def _normalize_display_name(raw: str) -> str:
"""把前端显示名转成 API model ID。
例: 'GPT-5.4-Chat' → 'gpt-5.4-chat', 'Claude-Opus-4.7' → 'claude-opus-4.7'
"""
return raw.lower().replace(" ", "-")
def _extract_id_from_model(m: dict) -> str | None:
"""从 model_config 条目中提取可用于 API 的 model ID。
优先用 modelId(前端字段名可能随版本变化),fallback 到 name/displayName。
"""
for key in ("modelId", "model_id", "id"):
v = m.get(key)
if v and isinstance(v, str):
return _normalize_display_name(v)
# displayName 如 "GPT-5.4-Chat"
dn = m.get("displayName") or m.get("display_name") or m.get("name")
if dn and isinstance(dn, str):
return _normalize_display_name(dn)
return None
def _extract_display_name(m: dict) -> str | None:
"""提取上游内部发送用的模型显示名(如 'GPT-5.4-Chat')。
这个名字会被 send_message 的 selected_model 字段使用。
"""
for key in ("displayName", "display_name", "name"):
v = m.get(key)
if v and isinstance(v, str):
return v
return None
async def fetch_models(base_url: str = DEFAULT_UPSTREAM_URL) -> dict[str, str]:
"""从上游获取 model_config,返回 {api_model_id: display_name}。
失败时 fallback 到硬编码 MAP。
"""
global _cache, _cache_at
now = time.time()
if _cache is not None and (now - _cache_at) < _cache_ttl:
return _cache
url = f"{base_url.rstrip('/')}/proxy/v1/model_config/models?a=0"
try:
async with httpx.AsyncClient(timeout=10) as client:
resp = await client.get(url, headers={
"User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36",
"Referer": f"{base_url}/newtab",
})
resp.raise_for_status()
data = resp.json()
models_list = data if isinstance(data, list) else data.get("data", data.get("models", []))
if not isinstance(models_list, list):
logger.warning("model_config: unexpected shape, falling back")
return dict(_FALLBACK_MAP)
result: dict[str, str] = {}
# 始终保留 best
result["best"] = "最佳"
for m in models_list:
if not isinstance(m, dict):
continue
api_id = _extract_id_from_model(m)
display_name = _extract_display_name(m)
if api_id and display_name:
result[api_id] = display_name
if len(result) > 1: # 至少 best + 1个真实模型
_cache = result
_cache_at = now
logger.info("Fetched %d models from model_config", len(result))
return result
logger.warning("model_config returned too few models (%d), falling back", len(result))
return dict(_FALLBACK_MAP)
except Exception as e:
logger.warning("Failed to fetch model_config: %s, falling back", e)
return dict(_FALLBACK_MAP)
def get_fallback_map() -> dict[str, str]:
"""返回硬编码 fallback map(不走网络)。"""
return dict(_FALLBACK_MAP)
def invalidate_cache():
"""手动清缓存,下次请求重新拉取。"""
global _cache, _cache_at
_cache = None
_cache_at = 0
# Priority model routing (for fallback on 429)
PRIORITY_ROUTE = [
{"display": "最佳", "tier": "primary"},
{"display": "Claude-Sonnet-4.6", "tier": "primary"},
{"display": "GPT-5.2-Chat", "tier": "backup"},
{"display": "Gemini-3.1-Pro", "tier": "backup"},
]
def resolve_priority_chain(upstream_model: str, models: dict) -> list[str]:
"""返回优先尝试的模型列表(按顺序)。"""
if upstream_model != "最佳":
return [upstream_model]
chain = []
for entry in PRIORITY_ROUTE:
for api_id, display in models.items():
if display == entry["display"]:
chain.append(api_id)
break
return chain