804 lines
36 KiB
Python
804 lines
36 KiB
Python
#!/usr/bin/env python3
|
||
"""
|
||
ocg_router.py — OpenCode Go 多 key 路由代理
|
||
===========================================
|
||
OpenAI 兼容代理,背后池化 6 个 OpenCode Go API key。
|
||
基于 usage_stats.json 的剩余 quota 智能路由,402/429/5xx 自动故障切换。
|
||
|
||
端点:
|
||
GET /v1/models — 合并所有 key 的模型列表
|
||
POST /v1/chat/completions — 转发(支持 stream)
|
||
GET /health — 健康检查
|
||
GET /api/status — 代理状态(监控用)
|
||
|
||
架构:
|
||
┌─ OMO (Windows) ──────► ocg_router (:19878) ─► key5 (最空闲) ─► opencode.ai
|
||
│ │ key6 (备用)
|
||
│ │ key4 (备用)
|
||
│ ▼ ...
|
||
│ usage_stats.json
|
||
└─ dashboard (:5803) ◄── /api/status
|
||
"""
|
||
|
||
import os, sys, json, time, logging, threading
|
||
from http.server import ThreadingHTTPServer, BaseHTTPRequestHandler
|
||
from urllib.request import Request, urlopen, HTTPError
|
||
from urllib.error import URLError
|
||
from datetime import datetime, timezone
|
||
from pathlib import Path
|
||
|
||
# ── 路径(平台自适应)──────────────────────────────────────────
|
||
_SCRIPT_DIR = Path(__file__).resolve().parent # gateway/scripts/
|
||
_GATEWAY_DIR = _SCRIPT_DIR.parent # gateway/
|
||
_PROJECT_DIR = _GATEWAY_DIR.parent # AgentsMeeting/
|
||
sys.path.insert(0, str(_SCRIPT_DIR))
|
||
from proc_guard import guard
|
||
|
||
# ── 配置 ────────────────────────────────────────────────────
|
||
LISTEN_HOST = "0.0.0.0"
|
||
LISTEN_PORT = 19878
|
||
OPENCODE_BASE = "https://opencode.ai/zen/go/v1"
|
||
USAGE_STATS_FILE = _GATEWAY_DIR / "temp" / "usage_stats.json"
|
||
ACCOUNTS_FILE = _SCRIPT_DIR / "usage_monitor" / "accounts.json"
|
||
|
||
# 故障切换:同一个 key 失败后冷却时间(秒)
|
||
FAIL_COOLDOWN_SEC = 300 # 5 分钟(key 自身故障)
|
||
UPSTREAM_COOLDOWN_SEC = 30 # 上游模型级故障(Router.Unavailable)短冷却:换 key 无意义
|
||
# 用量数据刷新间隔
|
||
USAGE_REFRESH_INTERVAL = 120 # 2 分钟
|
||
# 最大重试次数(所有 key 耗尽)
|
||
MAX_RETRIES = 6
|
||
|
||
# ── 日志 ────────────────────────────────────────────────────
|
||
LOG_DIR = _GATEWAY_DIR / "logs"
|
||
LOG_DIR.mkdir(parents=True, exist_ok=True)
|
||
logging.basicConfig(
|
||
level=logging.INFO,
|
||
format="%(asctime)s [ocg_router] %(levelname)s: %(message)s",
|
||
handlers=[
|
||
logging.FileHandler(str(LOG_DIR / "ocg_router.log"), encoding="utf-8"),
|
||
logging.StreamHandler(),
|
||
],
|
||
)
|
||
log = logging.getLogger("ocg_router")
|
||
|
||
# ── 全局状态 ─────────────────────────────────────────────────
|
||
_keys: list[dict] = [] # [{key_id, label, api_key, workspace_id}]
|
||
_key_usage: dict[str, dict] = {} # key_id → {rolling_pct, weekly_pct, monthly_pct}
|
||
_key_failures: dict[str, float] = {} # key_id → fail_until_timestamp
|
||
_key_failure_reason: dict[str, str] = {} # key_id → 最近失败原因(冷却跳过时用于报错)
|
||
_key_fail_count: dict[str, int] = {} # key_id → 连续失败次数(指数退避基数,2026-08-14 老莫)
|
||
_route_stats: dict[str, int] = {} # key_id → hit_count
|
||
_total_hits = 0
|
||
_current_key_id = ""
|
||
# 2026-08-13 并发模式:round-robin 轮询索引(锁 key 轮着来)
|
||
_rr_index = 0
|
||
_state_lock = threading.Lock()
|
||
|
||
|
||
# ═══════════════════════════════════════════════════════════════
|
||
# 用量数据加载
|
||
# ═══════════════════════════════════════════════════════════════
|
||
|
||
def load_accounts():
|
||
"""从 accounts.json 读取所有 OCG key。enabled=false 的 key 直接跳过(不在路由池中)。"""
|
||
keys = []
|
||
if not ACCOUNTS_FILE.exists():
|
||
log.error("accounts.json not found: %s", ACCOUNTS_FILE)
|
||
return keys
|
||
try:
|
||
with open(str(ACCOUNTS_FILE), "r", encoding="utf-8") as f:
|
||
data = json.load(f)
|
||
for a in data.get("accounts", []):
|
||
if a.get("provider", "").startswith("kimi"):
|
||
continue # 跳过 Kimi
|
||
kid = a.get("key_id", "")
|
||
if not kid:
|
||
continue
|
||
# 显式停用的 key 不加入路由池
|
||
if a.get("enabled", True) is False:
|
||
log.info("key %s disabled in accounts.json, skipped", kid)
|
||
continue
|
||
keys.append({
|
||
"key_id": kid,
|
||
"label": a.get("label", kid),
|
||
"api_key": a.get("key", ""),
|
||
"workspace_id": a.get("workspace_id", ""),
|
||
# 续订取消/显式标记 → 视为不可用(采集失败时也能正确跳过)
|
||
"subscribed": a.get("subscribed", a.get("renewal") != "cancelled"),
|
||
})
|
||
log.info("loaded %d OCG keys from accounts.json", len(keys))
|
||
except Exception as e:
|
||
log.error("failed to load accounts.json: %s", e)
|
||
return keys
|
||
|
||
|
||
def load_usage():
|
||
"""从 usage_stats.json 读取各 key 用量。"""
|
||
usage = {}
|
||
if not USAGE_STATS_FILE.exists():
|
||
log.warning("usage_stats.json not found: %s", USAGE_STATS_FILE)
|
||
return usage
|
||
try:
|
||
with open(str(USAGE_STATS_FILE), "r", encoding="utf-8") as f:
|
||
data = json.load(f)
|
||
for a in data.get("accounts", []):
|
||
kid = a.get("key_id", "")
|
||
rolling = a.get("rolling") or {}
|
||
weekly = a.get("weekly") or {}
|
||
monthly = a.get("monthly") or {}
|
||
has_error = bool(a.get("error"))
|
||
sub_active = a.get("subscribed", True)
|
||
usage[kid] = {
|
||
"rolling_pct": rolling.get("usage_percent", 0) if rolling else 0,
|
||
"weekly_pct": weekly.get("usage_percent", 0) if weekly else 0,
|
||
"monthly_pct": monthly.get("usage_percent", 0) if monthly else 0,
|
||
"error": has_error,
|
||
"subscribed": sub_active,
|
||
"last_update": a.get("last_update_iso", ""),
|
||
}
|
||
log.info("loaded usage for %d keys", len(usage))
|
||
except Exception as e:
|
||
log.error("failed to load usage_stats.json: %s", e)
|
||
return usage
|
||
|
||
|
||
def refresh_state():
|
||
"""刷新 key 池和用量数据。调用此函数前必须先获取 _state_lock。"""
|
||
global _keys, _key_usage
|
||
_keys = load_accounts()
|
||
_key_usage = load_usage()
|
||
# 合并 accounts.json 的订阅标记:采集失败时 usage 缺失该信息,
|
||
# 以 accounts.json 的 renewal/subscribed 为准
|
||
for k in _keys:
|
||
kid = k["key_id"]
|
||
if not k.get("subscribed", True):
|
||
_key_usage.setdefault(kid, {})["subscribed"] = False
|
||
log.info("state refreshed: %d keys, %d with usage data", len(_keys), len(_key_usage))
|
||
|
||
|
||
# ═══════════════════════════════════════════════════════════════
|
||
# Key 选择(核心路由逻辑)
|
||
# ═══════════════════════════════════════════════════════════════
|
||
|
||
def _key_health_score(kid):
|
||
"""
|
||
计算 key 健康度(越低越好):
|
||
- 已取消订阅的 key 排最末(score = 999)
|
||
- 在冷却期内的 key 排倒数第二(score = 888)
|
||
- 用量采集失败(error)仅降权(score = 500),key 仍可用
|
||
- 否则按 rolling% → weekly% → monthly% 加权
|
||
"""
|
||
usage = _key_usage.get(kid, {})
|
||
|
||
# 订阅取消 → 不可用
|
||
if not usage.get("subscribed", True):
|
||
return 999
|
||
|
||
# 用量采集失败 → 降权但可用(不因采集问题误杀 key)
|
||
if usage.get("error"):
|
||
return 500
|
||
|
||
# 故障冷却期内 → 低优先级
|
||
now = time.time()
|
||
fail_until = _key_failures.get(kid, 0)
|
||
if now < fail_until:
|
||
return 888
|
||
|
||
# 正常 key:按用量打分(用量越低分越低=越优先)
|
||
r = usage.get("rolling_pct", 0) or 0
|
||
w = usage.get("weekly_pct", 0) or 0
|
||
m = usage.get("monthly_pct", 0) or 0
|
||
return r * 0.6 + w * 0.3 + m * 0.1 # 滚动窗口占比最大
|
||
|
||
|
||
def pick_key():
|
||
"""选择当前最空闲的 key。返回 key dict 或 None。"""
|
||
if not _keys:
|
||
return None
|
||
# 按健康度排序
|
||
ranked = sorted(_keys, key=lambda k: _key_health_score(k["key_id"]))
|
||
best = ranked[0]
|
||
score = _key_health_score(best["key_id"])
|
||
if score >= 888: # 所有 key 都不可用
|
||
log.warning("ALL keys unavailable — best score=%.1f for %s", score, best["key_id"])
|
||
return best
|
||
|
||
|
||
def pick_key_round_robin():
|
||
"""并发模式:round-robin 轮询选 key(锁 key 轮着来)。
|
||
每次请求用下一个 key,避免并发压同一 key(2026-08-13 老莫设计)。
|
||
跳过不可用 key(失败冷却中/未订阅)。"""
|
||
global _rr_index
|
||
if not _keys:
|
||
return None
|
||
with _state_lock:
|
||
n = len(_keys)
|
||
if n == 0:
|
||
return None
|
||
# 轮询:从 _rr_index 开始找下一个可用 key
|
||
for offset in range(n):
|
||
idx = (_rr_index + offset) % n
|
||
k = _keys[idx]
|
||
kid = k["key_id"]
|
||
# 跳过不可用(失败冷却中/未订阅)
|
||
if time.time() < _key_failures.get(kid, 0):
|
||
continue
|
||
usage = _key_usage.get(kid, {})
|
||
if not usage.get("subscribed", True):
|
||
continue
|
||
# 选中:更新索引到下一个
|
||
_rr_index = (idx + 1) % n
|
||
return k
|
||
# 全部不可用 → 返回第一个(兜底)
|
||
return _keys[0]
|
||
|
||
|
||
def mark_key_failed(kid, reason="", cooldown=None):
|
||
"""标记 key 故障,进入指数退避冷却期(按 key 独立统计失败次数)。
|
||
2026-08-14 老莫设计:失败冷却时长指数增加(base × 2^count),上限 600s(10分钟)。
|
||
空输出/报错 → 立即换下一个 key,同时该 key 进入指数退避冷却。"""
|
||
now = time.time()
|
||
with _state_lock:
|
||
# 失败计数 +1(独立统计)
|
||
_key_fail_count[kid] = _key_fail_count.get(kid, 0) + 1
|
||
count = _key_fail_count[kid]
|
||
# 指数退避:base=2s, 每次失败 ×2,上限 600s(10分钟)
|
||
base = 2
|
||
cool = min(base * (2 ** (count - 1)), 600)
|
||
until = now + cool
|
||
_key_failures[kid] = until
|
||
if reason:
|
||
_key_failure_reason[kid] = reason
|
||
log.warning("key %s marked failed (count=%d, cooldown=%ds): %s", kid, count, cool, reason)
|
||
|
||
|
||
def reset_key_fail(kid):
|
||
"""key 成功 → 重置失败计数(清除指数退避)"""
|
||
with _state_lock:
|
||
if kid in _key_fail_count:
|
||
del _key_fail_count[kid]
|
||
if kid in _key_failures:
|
||
del _key_failures[kid]
|
||
if kid in _key_failure_reason:
|
||
del _key_failure_reason[kid]
|
||
|
||
|
||
def record_route(kid):
|
||
"""记录一次成功路由 + 重置失败计数(指数退避清除)。"""
|
||
global _total_hits, _current_key_id
|
||
with _state_lock:
|
||
_route_stats[kid] = _route_stats.get(kid, 0) + 1
|
||
_total_hits += 1
|
||
_current_key_id = kid
|
||
reset_key_fail(kid) # 成功重置失败计数
|
||
|
||
|
||
# ═══════════════════════════════════════════════════════════════
|
||
# HTTP 代理核心
|
||
# ═══════════════════════════════════════════════════════════════
|
||
|
||
class RouterHandler(BaseHTTPRequestHandler):
|
||
"""HTTP 请求路由处理器。"""
|
||
|
||
def log_message(self, format, *args):
|
||
pass # 禁用标准库日志
|
||
|
||
def do_GET(self):
|
||
self._handle()
|
||
|
||
def do_POST(self):
|
||
self._handle()
|
||
|
||
def do_OPTIONS(self):
|
||
self.send_response(200)
|
||
self.send_header("Access-Control-Allow-Origin", "*")
|
||
self.send_header("Access-Control-Allow-Methods", "GET,POST,OPTIONS")
|
||
self.send_header("Access-Control-Allow-Headers", "*")
|
||
self.end_headers()
|
||
|
||
def _handle(self):
|
||
path = self.path.rstrip("/") or "/"
|
||
|
||
# ── 健康检查 ──
|
||
if path == "/health":
|
||
self._json_response({"ok": True, "time": datetime.now(timezone.utc).isoformat()})
|
||
return
|
||
|
||
# ── 代理状态 ──
|
||
if path == "/api/status":
|
||
self._serve_status()
|
||
return
|
||
|
||
# ── Key 启用/停用 ──
|
||
if path == "/api/keys/toggle" and self.command == "POST":
|
||
self._toggle_key()
|
||
return
|
||
|
||
# ── 立即刷新状态(accounts.json 被外部修改后调用)──
|
||
if path == "/api/reload" and self.command == "POST":
|
||
with _state_lock:
|
||
refresh_state()
|
||
self._json_response({"ok": True, "keys": len(_keys)})
|
||
return
|
||
|
||
# ── 模型列表 ──
|
||
if path == "/v1/models":
|
||
self._serve_models()
|
||
return
|
||
|
||
# ── Chat Completions ──
|
||
if path.endswith("/v1/chat/completions") or path == "/v1/chat/completions":
|
||
self._proxy_chat()
|
||
return
|
||
|
||
# ── 其他请求尝试透传 ──
|
||
self._proxy_any()
|
||
|
||
# ── Key 启用/停用 ──
|
||
def _toggle_key(self):
|
||
"""POST /api/keys/toggle {key_id, enabled} → 写 accounts.json 并立即 reload。"""
|
||
try:
|
||
content_length = int(self.headers.get("Content-Length", 0))
|
||
body = self.rfile.read(content_length) if content_length > 0 else b""
|
||
req_data = json.loads(body) if body else {}
|
||
kid = req_data.get("key_id", "")
|
||
enabled = bool(req_data.get("enabled", True))
|
||
if not kid:
|
||
self._error_response(400, "key_id required")
|
||
return
|
||
if not ACCOUNTS_FILE.exists():
|
||
self._error_response(500, "accounts.json not found")
|
||
return
|
||
with open(str(ACCOUNTS_FILE), "r", encoding="utf-8") as f:
|
||
data = json.load(f)
|
||
found = False
|
||
for a in data.get("accounts", []):
|
||
if a.get("key_id") == kid:
|
||
a["enabled"] = enabled
|
||
found = True
|
||
break
|
||
if not found:
|
||
self._error_response(404, f"key {kid} not found")
|
||
return
|
||
with open(str(ACCOUNTS_FILE), "w", encoding="utf-8") as f:
|
||
json.dump(data, f, ensure_ascii=False, indent=2)
|
||
# 立即 reload,使配置生效
|
||
with _state_lock:
|
||
refresh_state()
|
||
log.info("key %s enabled=%s (via API), pool now %d keys", kid, enabled, len(_keys))
|
||
self._json_response({"ok": True, "key_id": kid, "enabled": enabled, "pool_size": len(_keys)})
|
||
except Exception as e:
|
||
log.error("toggle key failed: %s", e)
|
||
self._error_response(500, str(e))
|
||
|
||
# ── Status API ──
|
||
def _serve_status(self):
|
||
with _state_lock:
|
||
pool = []
|
||
for k in _keys:
|
||
kid = k["key_id"]
|
||
usage = _key_usage.get(kid, {})
|
||
fail_until = _key_failures.get(kid, 0)
|
||
pool.append({
|
||
"key_id": kid,
|
||
"label": k.get("label", kid),
|
||
"rolling_pct": usage.get("rolling_pct", None),
|
||
"weekly_pct": usage.get("weekly_pct", None),
|
||
"monthly_pct": usage.get("monthly_pct", None),
|
||
"healthy": usage.get("subscribed", True) and (time.time() >= fail_until),
|
||
"in_cooldown": time.time() < fail_until,
|
||
"fail_reason": _key_failure_reason.get(kid, ""),
|
||
"error": usage.get("error", False),
|
||
"hits": _route_stats.get(kid, 0),
|
||
"last_update": usage.get("last_update", ""),
|
||
})
|
||
status = {
|
||
"ok": True,
|
||
"uptime_seconds": time.time() - _start_time,
|
||
"current_key": _current_key_id,
|
||
"total_requests": _total_hits,
|
||
"pool_size": len(_keys),
|
||
"pool_healthy": sum(1 for p in pool if p["healthy"]),
|
||
"pool": pool,
|
||
"last_refresh": _last_refresh.isoformat() if _last_refresh else None,
|
||
}
|
||
self._json_response(status)
|
||
|
||
# ── Models ──
|
||
def _serve_models(self):
|
||
"""合并所有 key 的模型列表。"""
|
||
seen = set()
|
||
models = []
|
||
with _state_lock:
|
||
keys_snapshot = list(_keys)
|
||
|
||
for k in keys_snapshot:
|
||
kid = k["key_id"]
|
||
api_key = k.get("api_key", "")
|
||
if not api_key:
|
||
continue
|
||
# 跳过不可用 key(采集 error 不跳过,仅降权)
|
||
usage = _key_usage.get(kid, {})
|
||
if not usage.get("subscribed", True):
|
||
continue
|
||
if time.time() < _key_failures.get(kid, 0):
|
||
continue
|
||
|
||
try:
|
||
req = Request(
|
||
f"{OPENCODE_BASE}/models",
|
||
headers={
|
||
"Authorization": f"Bearer {api_key}",
|
||
"User-Agent": "Mozilla/5.0",
|
||
},
|
||
method="GET",
|
||
)
|
||
with urlopen(req, timeout=10) as resp:
|
||
data = json.loads(resp.read().decode("utf-8"))
|
||
for m in data.get("data", []):
|
||
mid = m.get("id", "")
|
||
if mid not in seen:
|
||
seen.add(mid)
|
||
models.append(m)
|
||
break # 只问第一个健康的 key 就够了(模型列表相同)
|
||
except Exception:
|
||
continue
|
||
|
||
if not models:
|
||
# fallback:硬编码已知模型列表
|
||
models = [
|
||
{"id": "deepseek-v4-pro", "object": "model"},
|
||
{"id": "deepseek-v4-flash", "object": "model"},
|
||
{"id": "kimi-k2.7-code", "object": "model"},
|
||
{"id": "glm-5.2", "object": "model"},
|
||
{"id": "qwen3.7-plus", "object": "model"},
|
||
{"id": "qwen3.7-max", "object": "model"},
|
||
]
|
||
|
||
self._json_response({"object": "list", "data": models})
|
||
|
||
# ── Chat Completions Proxy ──
|
||
def _proxy_chat(self):
|
||
start = time.time()
|
||
content_length = int(self.headers.get("Content-Length", 0))
|
||
body = self.rfile.read(content_length) if content_length > 0 else b""
|
||
|
||
is_stream = False
|
||
model_name = "unknown"
|
||
if body:
|
||
try:
|
||
req_data = json.loads(body)
|
||
is_stream = req_data.get("stream", False)
|
||
model_name = req_data.get("model", "unknown")
|
||
log.info("chat request: model=%s stream=%s body_len=%d", model_name, is_stream, len(body))
|
||
except (json.JSONDecodeError, UnicodeDecodeError) as e:
|
||
log.warning("failed to parse request body (len=%d): %s", len(body), e)
|
||
|
||
# 2026-08-13 并发模式:header X-OCG-Mode: concurrent → round-robin 锁 key 轮询
|
||
concurrent_mode = self.headers.get("X-OCG-Mode", "").lower() == "concurrent"
|
||
|
||
# 尝试所有 key(按用量排序),直到成功
|
||
with _state_lock:
|
||
usage_snapshot = dict(_key_usage)
|
||
failures_snapshot = dict(_key_failures)
|
||
if concurrent_mode:
|
||
# 并发模式:round-robin 轮询(锁 key 轮着来)
|
||
ranked_keys = []
|
||
n = len(_keys)
|
||
global _rr_index
|
||
for offset in range(n):
|
||
idx = (_rr_index + offset) % n
|
||
k = _keys[idx]
|
||
kid = k["key_id"]
|
||
if time.time() < failures_snapshot.get(kid, 0):
|
||
continue
|
||
usage = usage_snapshot.get(kid, {})
|
||
if not usage.get("subscribed", True):
|
||
continue
|
||
ranked_keys.append(k)
|
||
# 更新轮询索引
|
||
if ranked_keys:
|
||
first_kid = ranked_keys[0]["key_id"]
|
||
for i, k in enumerate(_keys):
|
||
if k["key_id"] == first_kid:
|
||
_rr_index = (i + 1) % len(_keys)
|
||
break
|
||
log.info("concurrent mode: round-robin, next_idx=%d, pool=%d", _rr_index, len(ranked_keys))
|
||
else:
|
||
ranked_keys = sorted(
|
||
list(_keys),
|
||
key=lambda k: _key_health_score(k["key_id"])
|
||
)
|
||
|
||
last_error = ""
|
||
success = False
|
||
client_error = False # True=客户端请求错误(400/ModelError),应返回 4xx 而非 502
|
||
|
||
for attempt, k in enumerate(ranked_keys):
|
||
kid = k["key_id"]
|
||
api_key = k.get("api_key", "")
|
||
if not api_key:
|
||
continue
|
||
|
||
# 检查 key 是否可用(采集 error 不跳过,仅降权)
|
||
usage = usage_snapshot.get(kid, {})
|
||
if not usage.get("subscribed", True):
|
||
continue
|
||
if time.time() < failures_snapshot.get(kid, 0):
|
||
continue
|
||
|
||
log.info("→ %s via %s (attempt %d/%d)", model_name, kid, attempt + 1, len(ranked_keys))
|
||
|
||
try:
|
||
result = self._forward_request(api_key, body, is_stream)
|
||
if result["ok"]:
|
||
record_route(kid)
|
||
elapsed = time.time() - start
|
||
log.info("✓ %s via %s OK (%.2fs)", model_name, kid, elapsed)
|
||
success = True
|
||
break
|
||
else:
|
||
last_error = result.get("error", "unknown")
|
||
log.warning("⚠ %s via %s FAILED: %s", model_name, kid, last_error)
|
||
# 错误分类:只有真正的 key 级错误才冷却 key
|
||
# - 400 客户端错误:不重试、不冷却(换 key 无意义),直接返回给调用方
|
||
# - 401/403 认证失败:key 级故障 → 冷却 300s(但 ModelError=客户端传了不存在的模型,不冷却)
|
||
# - 402/429 配额/限流:key 级 → 冷却 60s
|
||
# - 5xx/Router.Unavailable:上游服务问题 → 不冷却 key(误杀元凶),立即失败
|
||
status = result.get("status", 0)
|
||
if status == -1:
|
||
# 客户端已断开:直接终止,不重试不冷却
|
||
log.info("client disconnected during %s via %s, aborting", model_name, kid)
|
||
success = False
|
||
break
|
||
if status == 400:
|
||
log.warning("client error 400, returning to caller without retry")
|
||
client_error = True
|
||
success = False
|
||
break
|
||
if status in (401, 403):
|
||
# 区分:ModelError(模型不存在/不支持)是客户端请求错误,换 key 无用且会
|
||
# 把所有 key 打进冷却(雪崩);认证失败(Invalid key)才是 key 级故障。
|
||
if "ModelError" in last_error or "model is not" in last_error.lower() \
|
||
or "not supported" in last_error.lower() or "model not found" in last_error.lower():
|
||
log.warning("client ModelError (bad model name '%s'), returning 400 without cooling keys", model_name)
|
||
client_error = True
|
||
success = False
|
||
break
|
||
mark_key_failed(kid, last_error, cooldown=FAIL_COOLDOWN_SEC)
|
||
elif status in (402, 429):
|
||
mark_key_failed(kid, last_error, cooldown=60)
|
||
elif status >= 500:
|
||
# 上游服务端错误(Internal server error / Router.Unavailable):
|
||
# 不杀 key。立即失败返回,避免"全部 key 被误杀进冷却"的雪崩。
|
||
log.warning("upstream %d (service-level) via %s — NOT cooling key, failing request", status, kid)
|
||
success = False
|
||
break
|
||
except Exception as e:
|
||
last_error = str(e)
|
||
log.warning("⚠ %s via %s exception: %s", model_name, kid, e)
|
||
mark_key_failed(kid, last_error)
|
||
|
||
if not success:
|
||
elapsed = time.time() - start
|
||
if not last_error:
|
||
# 所有 key 在冷却中被跳过 → 组装原因+最早到期时间
|
||
now = time.time()
|
||
cooling = [(kid, until) for kid, until in failures_snapshot.items() if now < until]
|
||
if cooling:
|
||
earliest = min(until for _, until in cooling)
|
||
reasons = {_key_failure_reason.get(kid, "unknown") for kid, _ in cooling}
|
||
reason_str = "; ".join(sorted(r for r in reasons if r and r != "unknown"))
|
||
last_error = (f"all keys in cooldown until "
|
||
f"{datetime.fromtimestamp(earliest).strftime('%H:%M:%S')}"
|
||
f"{' (' + reason_str + ')' if reason_str else ''}")
|
||
if client_error:
|
||
# 客户端请求错误(400/ModelError):返回 400,不是服务端故障
|
||
log.warning("✗ %s client error (%.2fs): %s", model_name, elapsed, last_error)
|
||
self._error_response(400, last_error)
|
||
else:
|
||
log.error("✗ %s ALL KEYS EXHAUSTED (%.2fs): %s", model_name, elapsed, last_error)
|
||
self._error_response(502, f"All OCG keys exhausted: {last_error}")
|
||
|
||
def _forward_request(self, api_key, body, is_stream):
|
||
"""转发单次请求到指定 key。返回 {ok, status, error}。"""
|
||
try:
|
||
req = Request(
|
||
f"{OPENCODE_BASE}/chat/completions",
|
||
data=body,
|
||
headers={
|
||
"Authorization": f"Bearer {api_key}",
|
||
"Content-Type": "application/json",
|
||
"User-Agent": "Mozilla/5.0",
|
||
},
|
||
method="POST",
|
||
)
|
||
with urlopen(req, timeout=180) as resp:
|
||
status = resp.status
|
||
if status >= 400:
|
||
err_body = resp.read().decode("utf-8", errors="replace")[:500]
|
||
return {"ok": False, "status": status, "error": f"HTTP {status}: {err_body}"}
|
||
|
||
# 透传响应 + 空输出检测(2026-08-14 老莫:空输出自动换 key)
|
||
resp_body = resp.read()
|
||
if not is_stream:
|
||
# 非流式:解析 JSON 检测空输出
|
||
try:
|
||
import json as _json
|
||
d = _json.loads(resp_body.decode("utf-8", errors="replace"))
|
||
msg = d.get("choices", [{}])[0].get("message", {}) or {}
|
||
content = msg.get("content", "")
|
||
# 推理模型(deepseek-v4-flash/pro 等)回复在 reasoning_content,
|
||
# content 可能为空(尤其 max_tokens 小时 token 全用于推理)——
|
||
# 有 reasoning_content 或 finish_reason=length 都视为正常,不算空输出
|
||
reasoning = msg.get("reasoning_content", "")
|
||
finish_reason = d.get("choices", [{}])[0].get("finish_reason", "")
|
||
has_reasoning = bool(reasoning and str(reasoning).strip())
|
||
if (not content or not str(content).strip()) and not has_reasoning:
|
||
if finish_reason == "length":
|
||
# max_tokens 截断,content 可能为空但属于正常响应
|
||
log.warning("content empty but finish_reason=length (truncated), passing through")
|
||
else:
|
||
log.warning("empty content detected via key, retrying next key")
|
||
return {"ok": False, "status": 0, "error": "empty_content"}
|
||
except Exception:
|
||
pass # 解析失败不阻断(正常透传)
|
||
self._send_response(status, dict(resp.headers), resp_body)
|
||
return {"ok": True}
|
||
|
||
except HTTPError as e:
|
||
err_body = e.read().decode("utf-8", errors="replace")[:500]
|
||
return {"ok": False, "status": e.code, "error": f"HTTP {e.code}: {err_body}"}
|
||
except (BrokenPipeError, ConnectionResetError, ConnectionAbortedError) as e:
|
||
# 客户端断开:不是 key 问题,标记特殊错误码让调用方不重试不冷却
|
||
return {"ok": False, "status": -1, "error": f"client_disconnected: {e}"}
|
||
except URLError as e:
|
||
return {"ok": False, "status": 0, "error": f"URLError: {e.reason}"}
|
||
except Exception as e:
|
||
return {"ok": False, "status": 0, "error": str(e)}
|
||
|
||
# ── 通用代理(其他路径) ──
|
||
def _proxy_any(self):
|
||
with _state_lock:
|
||
key = pick_key()
|
||
if not key:
|
||
self._error_response(503, "No available OCG keys")
|
||
return
|
||
api_key = key["api_key"]
|
||
self._forward_to_upstream(api_key)
|
||
|
||
def _forward_to_upstream(self, api_key):
|
||
content_length = int(self.headers.get("Content-Length", 0))
|
||
body = self.rfile.read(content_length) if content_length > 0 else b""
|
||
method = self.command
|
||
url = f"{OPENCODE_BASE}{self.path}"
|
||
|
||
excluded = {"host", "connection", "keep-alive", "transfer-encoding", "content-length"}
|
||
headers = {
|
||
k: v for k, v in self.headers.items()
|
||
if k.lower() not in excluded
|
||
}
|
||
headers["Authorization"] = f"Bearer {api_key}"
|
||
headers["User-Agent"] = "Mozilla/5.0"
|
||
|
||
try:
|
||
req = Request(url, data=body or None, headers=headers, method=method)
|
||
with urlopen(req, timeout=120) as resp:
|
||
resp_body = resp.read()
|
||
self._send_response(resp.status, dict(resp.headers), resp_body)
|
||
except HTTPError as e:
|
||
err_body = e.read().decode("utf-8", errors="replace")[:500]
|
||
self._error_response(e.code, err_body)
|
||
except Exception as e:
|
||
self._error_response(502, str(e))
|
||
|
||
# ── 工具方法 ──
|
||
def _json_response(self, data):
|
||
body = json.dumps(data, ensure_ascii=False).encode("utf-8")
|
||
self.send_response(200)
|
||
self.send_header("Content-Type", "application/json")
|
||
self.send_header("Access-Control-Allow-Origin", "*")
|
||
self.send_header("Content-Length", str(len(body)))
|
||
self.end_headers()
|
||
self.wfile.write(body)
|
||
|
||
def _error_response(self, code, message):
|
||
data = json.dumps({"error": str(message)}, ensure_ascii=False).encode("utf-8")
|
||
self.send_response(code)
|
||
self.send_header("Content-Type", "application/json")
|
||
self.send_header("Access-Control-Allow-Origin", "*")
|
||
self.send_header("Content-Length", str(len(data)))
|
||
self.end_headers()
|
||
self.wfile.write(data)
|
||
|
||
def _send_response(self, status, headers, body):
|
||
self.send_response(status)
|
||
allowed = {
|
||
"content-type", "content-encoding", "cache-control",
|
||
"x-request-id", "x-ratelimit-remaining", "x-ratelimit-reset",
|
||
"access-control-allow-origin",
|
||
}
|
||
# 注意:不转发 transfer-encoding: chunked — 代理已缓存整个响应,不是 chunked
|
||
for k, v in headers.items():
|
||
if k.lower() in allowed:
|
||
self.send_header(k, v)
|
||
self.send_header("Content-Length", str(len(body)))
|
||
self.send_header("Access-Control-Allow-Origin", "*")
|
||
self.end_headers()
|
||
self.wfile.write(body)
|
||
self.wfile.flush()
|
||
|
||
|
||
# ═══════════════════════════════════════════════════════════════
|
||
# 后台线程:周期性刷新状态
|
||
# ═══════════════════════════════════════════════════════════════
|
||
|
||
_start_time = time.time()
|
||
_last_refresh: datetime | None = None
|
||
|
||
|
||
def _bg_refresh():
|
||
"""后台线程:每 USAGE_REFRESH_INTERVAL 秒刷新用量数据。"""
|
||
global _last_refresh
|
||
while True:
|
||
time.sleep(USAGE_REFRESH_INTERVAL)
|
||
try:
|
||
with _state_lock:
|
||
refresh_state()
|
||
_last_refresh = datetime.now(timezone.utc)
|
||
# 清空已过冷却期的 key
|
||
now = time.time()
|
||
with _state_lock:
|
||
expired = [k for k, v in _key_failures.items() if now >= v]
|
||
for k in expired:
|
||
del _key_failures[k]
|
||
_key_failure_reason.pop(k, None)
|
||
log.info("key %s cooldown expired, restored", k)
|
||
log.info("bg refresh: %d keys, healthy pool", len(_keys))
|
||
except Exception as e:
|
||
log.error("bg refresh failed: %s", e)
|
||
|
||
|
||
# ═══════════════════════════════════════════════════════════════
|
||
# 启动
|
||
# ═══════════════════════════════════════════════════════════════
|
||
|
||
def main():
|
||
# PID 锁
|
||
g = guard("ocg_router")
|
||
if not g.ok:
|
||
log.error("ocg_router 已有实例在运行 (PID %s),退出", g.message)
|
||
sys.exit(1)
|
||
|
||
# 初始加载
|
||
with _state_lock:
|
||
refresh_state()
|
||
global _last_refresh
|
||
_last_refresh = datetime.now(timezone.utc)
|
||
|
||
# 启动后台刷新线程
|
||
refresh_thread = threading.Thread(target=_bg_refresh, name="ocg_router_refresh", daemon=True)
|
||
refresh_thread.start()
|
||
|
||
# 启动 HTTP 服务(ThreadingHTTPServer — 多线程,避免大请求阻塞 /api/status 等监控端点)
|
||
# daemon_threads=True:客户端断开/线程卡住不会阻塞服务退出
|
||
server = ThreadingHTTPServer((LISTEN_HOST, LISTEN_PORT), RouterHandler)
|
||
server.daemon_threads = True
|
||
# 限制最大并发连接数(每个连接一个线程;信号量在 handler 层控制)
|
||
MAX_CONCURRENT = 16
|
||
log.info("ocg_router 启动 → http://%s:%d (max_parallel=%d)", LISTEN_HOST, LISTEN_PORT, MAX_CONCURRENT)
|
||
log.info("keys loaded: %d (%d healthy)", len(_keys),
|
||
sum(1 for k in _keys if _key_health_score(k["key_id"]) < 888))
|
||
log.info("fail cooldown: %ds, usage refresh: %ds", FAIL_COOLDOWN_SEC, USAGE_REFRESH_INTERVAL)
|
||
|
||
try:
|
||
server.serve_forever()
|
||
except KeyboardInterrupt:
|
||
log.info("收到中断信号,关闭...")
|
||
server.shutdown()
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|