- upstream.py:httpx.AsyncClient 模块级单例(keepalive,limits=100, 超时 connect10/read120/write10/pool30);stream() 始终注入 stream_options.include_usage(计量不依赖客户端)+ filter_usage_chunk (客户端未要求 usage 时剥除该 chunk);D-P4 首 token 前 failover 链、 流中失败抛 UpstreamAborted(不可切换);ttfb_ms 记录 - normalize_usage 三家归一:deepseek prompt_cache_hit_tokens / openai prompt_tokens_details.cached_tokens / anthropic cache_read_input_tokens (命中数>总数时钳制) - model_pool:+provider(枚举校验 deepseek/openai/anthropic)+in_hit_price (缺省 = price_in×1/30,D-P3) - 测试 +6:透传+sink/failover/流中 aborted/全挂 502/三家归一/过滤,全量 345 passed
279 lines
11 KiB
Python
279 lines
11 KiB
Python
"""模型池(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")
|
||
PROVIDERS = ("deepseek", "openai", "anthropic") # 代理层 usage 归一化用(D-P3)
|
||
|
||
# 池条目允许的字段(其余字段拒绝写入)
|
||
ENTRY_FIELDS = {
|
||
"id", "name", "tier", "backend", "base_url", "model", "api_key",
|
||
"price_in", "price_out", "temperature", "max_tokens", "enabled",
|
||
"provider", "in_hit_price", # 代理层扩展(D-P3):usage 归一化 / 按命中价选上游
|
||
}
|
||
|
||
# 单价默认值($/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} 之一")
|
||
if entry.get("provider") and entry["provider"] not in PROVIDERS:
|
||
raise PoolError(f"provider 必须是 {PROVIDERS} 之一")
|
||
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)),
|
||
# 代理层扩展(D-P3):provider 缺省 openai;in_hit_price 缺省 = price_in × 1/30
|
||
"provider": (str(entry["provider"]) if entry.get("provider") else "openai"),
|
||
"in_hit_price": (float(entry["in_hit_price"]) if entry.get("in_hit_price") is not None
|
||
else round(price_in / 30.0, 6)),
|
||
}
|
||
|
||
@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
|