"""模型池(PoolStore)—— 多价位异构模型注册表。 设计(《实现方案_v4_模型池与工具智能体.md》D1): - 叙事从"端云分工"泛化为"按价位分工":local(零边际成本,内置 llama.cpp)、 budget(低价 API)、premium(高价 API)。位置只是价位的属性之一。 - 池条目存"端点 + 凭据 + 模型名 + 价位 + 单价($/1M tokens)",不存模型权重。 - roles 把池条目指派给三个角色:architect(决策/终审)、worker(实现/自验证)、 agent(智能体工具循环)。角色留空 = 沿用经典单模型设置(向后兼容)。 - 持久化到 config/model_pool.json(gitignore,与 settings.json 同级)。 """ import json import re import threading from pathlib import Path from typing import Any, Dict, Optional _POOL_PATH = Path(__file__).resolve().parent.parent / "config" / "model_pool.json" # 合法取值 TIERS = ("local", "budget", "premium") BACKENDS = ("mock", "llama_server", "openai") ROLES = ("architect", "worker", "agent") # 池条目允许的字段(其余字段拒绝写入) ENTRY_FIELDS = { "id", "name", "tier", "backend", "base_url", "model", "api_key", "price_in", "price_out", "temperature", "max_tokens", "enabled", } # 单价默认值($/1M tokens);local 档为 0 PRICE_DEFAULTS = {"local": 0.0, "budget": 0.1, "premium": 1.0} def _empty_pool() -> Dict[str, Any]: return { "roles": {"architect": "", "worker": "", "agent": ""}, "entries": [], } class PoolError(ValueError): """池条目/角色配置非法。""" class PoolStore: """模型池注册表(内存 + model_pool.json 持久化,线程安全)。""" def __init__(self, path: Optional[Path] = None): self._path = Path(path) if path else _POOL_PATH self._lock = threading.Lock() self._data = _empty_pool() self.load() # ---------- 持久化 ---------- def load(self) -> None: if self._path.exists(): try: raw = json.loads(self._path.read_text(encoding="utf-8")) self._data = { "roles": {**_empty_pool()["roles"], **(raw.get("roles") or {})}, "entries": list(raw.get("entries") or []), } except Exception: self._data = _empty_pool() else: self._data = _empty_pool() def save(self) -> None: self._path.parent.mkdir(parents=True, exist_ok=True) self._path.write_text( json.dumps(self._data, ensure_ascii=False, indent=2), encoding="utf-8") # ---------- 条目 CRUD ---------- def list(self) -> Dict[str, Any]: """返回完整池(api_key 打码)。""" with self._lock: return { "roles": dict(self._data["roles"]), "entries": [self._masked(e) for e in self._data["entries"]], } def get(self, entry_id: str) -> Optional[Dict[str, Any]]: with self._lock: for e in self._data["entries"]: if e.get("id") == entry_id: return dict(e) return None def upsert(self, entry: Dict[str, Any]) -> Dict[str, Any]: """新增或更新条目(按 id)。返回打码后的条目。""" clean = self._validate(entry) with self._lock: entries = self._data["entries"] for i, e in enumerate(entries): if e.get("id") == clean["id"]: # 空 api_key 表示保留原值(前端不回传明文) if not clean.get("api_key"): clean["api_key"] = e.get("api_key", "") entries[i] = clean self.save() return self._masked(clean) entries.append(clean) self.save() return self._masked(clean) def delete(self, entry_id: str) -> bool: with self._lock: before = len(self._data["entries"]) self._data["entries"] = [ e for e in self._data["entries"] if e.get("id") != entry_id] changed = len(self._data["entries"]) != before if changed: # 清空指向被删条目的角色指派 for role, rid in self._data["roles"].items(): if rid == entry_id: self._data["roles"][role] = "" self.save() return changed # ---------- 角色指派 ---------- def set_roles(self, roles: Dict[str, str]) -> Dict[str, str]: """指派角色 -> 池条目 id(空串 = 沿用经典设置)。""" with self._lock: ids = {e.get("id") for e in self._data["entries"]} for role, rid in roles.items(): if role not in ROLES: raise PoolError(f"未知角色: {role}") if rid and rid not in ids: raise PoolError(f"角色 {role} 指向不存在的模型条目: {rid}") self._data["roles"][role] = rid or "" self.save() return dict(self._data["roles"]) def resolve(self, role: str) -> Optional[Dict[str, Any]]: """解析角色当前生效的池条目(未指派/条目禁用时返回 None = 用经典设置)。""" if role not in ROLES: return None with self._lock: rid = self._data["roles"].get(role, "") for e in self._data["entries"]: if e.get("id") == rid: return dict(e) if e.get("enabled", True) else None return None def find_by_model(self, model: str) -> Optional[Dict[str, Any]]: """按模型名找条目(用于按模型计价分账)。""" with self._lock: for e in self._data["entries"]: if e.get("model") == model: return dict(e) return None # ---------- 校验与工具 ---------- def _validate(self, entry: Dict[str, Any]) -> Dict[str, Any]: if not isinstance(entry, dict): raise PoolError("条目必须是对象") unknown = set(entry) - ENTRY_FIELDS if unknown: raise PoolError(f"非法字段: {sorted(unknown)}") eid = str(entry.get("id") or "").strip() if not eid: # 未提供 id 时按名称生成 slug base = re.sub(r"[^a-zA-Z0-9_-]+", "-", str(entry.get("name") or entry.get("model") or "model")).strip("-").lower() eid = base or "model" with self._lock: exist = {e.get("id") for e in self._data["entries"]} if eid in exist: n = 2 while f"{eid}-{n}" in exist: n += 1 eid = f"{eid}-{n}" elif not re.fullmatch(r"[a-zA-Z0-9_-]{1,64}", eid): raise PoolError("id 只允许字母/数字/-/_,长度 1-64") backend = entry.get("backend", "openai") if backend not in BACKENDS: raise PoolError(f"backend 必须是 {BACKENDS} 之一") tier = entry.get("tier", "budget") if tier not in TIERS: raise PoolError(f"tier 必须是 {TIERS} 之一") if backend != "mock" and not str(entry.get("base_url") or "").strip(): raise PoolError("非 mock 后端必须填写 base_url") if backend != "mock" and not str(entry.get("model") or "").strip(): raise PoolError("非 mock 后端必须填写 model") try: price_in = float(entry.get("price_in", PRICE_DEFAULTS[tier])) price_out = float(entry.get("price_out", PRICE_DEFAULTS[tier])) except (TypeError, ValueError): raise PoolError("price_in/price_out 必须是数字") if price_in < 0 or price_out < 0: raise PoolError("单价不能为负") try: temperature = float(entry.get("temperature", 0.3)) except (TypeError, ValueError): temperature = 0.3 try: max_tokens = int(entry.get("max_tokens", 4096)) except (TypeError, ValueError): max_tokens = 4096 return { "id": eid, "name": str(entry.get("name") or entry.get("model") or eid), "tier": tier, "backend": backend, "base_url": str(entry.get("base_url") or "").strip(), "model": str(entry.get("model") or "").strip(), "api_key": str(entry.get("api_key") or ""), "price_in": price_in, "price_out": price_out, "temperature": temperature, "max_tokens": max_tokens, "enabled": bool(entry.get("enabled", True)), } @staticmethod def _masked(entry: Dict[str, Any]) -> Dict[str, Any]: out = dict(entry) key = out.get("api_key") or "" out["api_key_set"] = bool(key) out["api_key"] = (key[:6] + "…") if key else "" return out def entry_to_architect_cfg(entry: Dict[str, Any]) -> Dict[str, Any]: """池条目 -> build_architect 配置段。""" cfg: Dict[str, Any] = { "model": entry.get("model") or "local", "base_url": entry.get("base_url") or "http://127.0.0.1:8901/v1", "temperature": float(entry.get("temperature", 0.2)), } if entry.get("api_key"): cfg["api_key"] = entry["api_key"] if entry.get("max_tokens"): cfg["max_tokens"] = int(entry["max_tokens"]) return cfg def entry_to_worker_cfg(entry: Dict[str, Any]) -> Dict[str, Any]: """池条目 -> build_worker 配置段。""" return { "backend": entry.get("backend") or "openai", "base_url": entry.get("base_url") or "", "model": entry.get("model") or "", "temperature": float(entry.get("temperature", 0.3)), } def compute_cost(entry: Dict[str, Any], input_tokens: int, output_tokens: int) -> float: """按条目单价估算成本(USD)。price 单位:$/1M tokens。""" return (input_tokens / 1e6) * float(entry.get("price_in", 0.0)) + \ (output_tokens / 1e6) * float(entry.get("price_out", 0.0)) # ---------- 全局单例 ---------- _store: Optional[PoolStore] = None def get_pool() -> PoolStore: global _store if _store is None: _store = PoolStore() return _store def reset_pool() -> None: """测试用:重置全局池单例。""" global _store _store = None