Files
projectAIpopular/gateway/model_pool.py
T
tzt e00e2cd2b1 feat(proxy): 架构与算法优化二轮——sense 运行时单例激活 T-X3、打码密钥泄漏修复、账本热路径 2.4x
- gateway/proxy/routes.py:sense 运行时(SenseStore+Grader+晋升表)按配置单例化,
  原每请求重建导致:每请求跑全量 DDL、重读工件文件,且 T-X3 60s 决策缓存随
  Grader 丢弃、命中率恒为 0(本轮最大收益,激活既有已测组件);
  T1 审计抽样修复原 NameError×3 被 except 吞掉(load_config/get_review/request_id
  均未定义、§9.4 从未入队)——request_id 上提、review 队列经 review_getter 注入
  (api.py 反向依赖解除)、sample_rate 走 load_config
- gateway/model_pool.py:新增 usable_entries() 未打码内通道;routes 的降级链/
  预算降档/能力位重定向/档位映射四处改走该通道——修复打码 api_key 流入上游
  Bearer 头导致带密钥条目重定向必 401 的隐性缺陷(HTTP 管理面仍用打码 list())
- gateway/proxy/ledger.py + billing.py:共享持久连接(原每操作新建,全链路
  每请求 3-5 次 connect)+ 日重置 UPDATE 每实例每日一次短路(原每请求全表扫描
  抢写锁);check_and_count A/B 0.75→0.31 ms/op(2.4x);新增 usage_stats()
  SQL 聚合,/admin/stats 从 list_usage(limit=50 万) Python 四遍扫描改为下推聚合
  (微基准 ~28x,随流水线性扩大)
- gateway/sense/grader.py:决策留痕 insert_decision 移入 asyncio.to_thread
  (原同步 sqlite 写直接跑在事件循环线程,高并发阻塞网关;与观察写批量队列同等保护)
- 全量 474 项两轮复核全绿(基线 474)
2026-09-18 23:59:44 +08:00

332 lines
13 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.
"""模型池(PoolStore)—— 多价位异构模型注册表。
设计(《实现方案_v4_模型池与工具智能体.md》D1):
- 叙事从"端云分工"泛化为"按价位分工"local(零边际成本,内置 llama.cpp)、
budget(低价 API)、premium(高价 API)。位置只是价位的属性之一。
- 池条目存"端点 + 凭据 + 模型名 + 价位 + 单价($/1M tokens",不存模型权重。
- roles 把池条目指派给三个角色:architect(决策/终审)、worker(实现/自验证)、
agent(智能体工具循环)。角色留空 = 沿用经典单模型设置(向后兼容)。
- 持久化到 config/model_pool.jsongitignore,与 settings.json 同级)。
"""
import json
import re
import threading
from pathlib import Path
from typing import Any, Dict, List, 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 归一化 / 按命中价选上游
"capabilities", # 能力位(T-X6):vision/tools/context_window
}
# 单价默认值($/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": [],
}
def _normalize_capabilities(raw: Any) -> Dict[str, Any]:
"""能力位归一(T-X6):缺省全兼容。
- vision/tools:缺省 True(未声明 = 不设限,老条目行为不变);
- context_window:缺省 0 = 不按上下文长度过滤。
"""
if not isinstance(raw, dict):
return {"vision": True, "tools": True, "context_window": 0}
try:
ctx = int(raw.get("context_window", 0) or 0)
except (TypeError, ValueError):
ctx = 0
return {"vision": bool(raw.get("vision", True)),
"tools": bool(raw.get("tools", True)),
"context_window": max(0, ctx)}
def filter_by_capabilities(entries: List[Dict[str, Any]], *,
need_vision: bool = False, need_tools: bool = False,
min_context_tokens: int = 0) -> List[Dict[str, Any]]:
"""能力位硬过滤(T-X6,采纳 cortiq capabilities 过滤)。
依次校验 vision/tools 声明位与 context_window >= 提示+输出预估
context_window=0 视为不限)。返回保持原顺序的合格条目子集。
"""
out: List[Dict[str, Any]] = []
for e in entries:
cap = e.get("capabilities") or _normalize_capabilities(None)
if need_vision and not cap.get("vision", True):
continue
if need_tools and not cap.get("tools", True):
continue
ctx = int(cap.get("context_window") or 0)
if min_context_tokens > 0 and 0 < ctx < min_context_tokens:
continue
out.append(e)
return out
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 打码,供 HTTP 管理面)。"""
with self._lock:
return {
"roles": dict(self._data["roles"]),
"entries": [self._masked(e) for e in self._data["entries"]],
}
def usable_entries(self) -> List[Dict[str, Any]]:
"""未打码条目快照——仅供代理转发链内部使用(failover/降档/能力重定向),
绝不进入任何 HTTP 响应。
背景:此前降级链/能力重定向从 list() 取打码条目,带 api_key 的条目
被重定向时向上游发送打码密钥(必 401),T-X1/T-X2/T-X6 链路静默劣化。
"""
with self._lock:
return [dict(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
capabilities = _normalize_capabilities(entry.get("capabilities"))
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 缺省 openaiin_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)),
# 能力位(T-X6):缺省全兼容(不破坏既有条目/选型行为)
"capabilities": capabilities,
}
@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