Files
projectAIpopular/gateway/model_pool.py
tzt f24f016e95 feat(proxy): T-P2 上游客户端(流式派发/usage 注入过滤/三家归一化/首 token 前 failover)
- 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
2026-09-05 09:14:31 +08:00

279 lines
11 KiB
Python
Raw Permalink 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, 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 缺省 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)),
}
@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