feat(failover): 支持 Provider 级别故障转移规则,默认全部转移策略

- 新增 failover_rules 配置:支持 success_failover_patterns(成功响应匹配时转移)
  和 error_stop_patterns(错误响应匹配时终止),支持按 status_code 过滤
- 修改默认转移策略:ErrorClassifier 不再返回 RAISE,所有错误默认继续转移
- TaskService 中客户端错误不再直接抛出,改为 break 继续尝试下一个候选
- 修复 proxy tunnel 连接/断连竞态:引入 per-node 锁和事件时间戳排序
- 优化 ProxyNode 状态判定:OFFLINE 统一由心跳超时判定,兼容多 worker 场景
- has_tunnel 改为纯检查方法,避免在 finally 块中误清理新注册连接
- Redis stream NOGROUP 异常自愈处理
- OAuthAccountDialog 输入框焦点样式补全
This commit is contained in:
fawney19
2026-03-01 00:21:16 +08:00
parent fbcb54a8a5
commit 005cc3e388
18 changed files with 941 additions and 116 deletions

View File

@@ -82,6 +82,32 @@ def _merge_pool_advanced_config(
return merged_config or None, config_changed
def _merge_failover_rules_config(
*,
provider_config: dict[str, Any] | None,
failover_rules: dict[str, Any] | None,
failover_rules_in_payload: bool,
) -> tuple[dict[str, Any] | None, bool]:
"""合并 failover_rules 到 provider.config。"""
merged_config = dict(provider_config or {})
config_changed = False
if not failover_rules_in_payload:
return merged_config or None, config_changed
if failover_rules is None:
if "failover_rules" in merged_config:
merged_config.pop("failover_rules", None)
config_changed = True
else:
next_value = dict(failover_rules)
if merged_config.get("failover_rules") != next_value:
merged_config["failover_rules"] = next_value
config_changed = True
return merged_config or None, config_changed
def _merge_claude_code_advanced_config(
*,
provider_type: str | None,
@@ -394,6 +420,15 @@ class AdminCreateProviderAdapter(AdminApiAdapter):
),
pool_advanced_in_payload=validated_data.pool_advanced is not None,
)
provider_config, _ = _merge_failover_rules_config(
provider_config=provider_config,
failover_rules=(
validated_data.failover_rules.model_dump()
if validated_data.failover_rules is not None
else None
),
failover_rules_in_payload=validated_data.failover_rules is not None,
)
# 创建 Provider 对象
provider = Provider(
@@ -509,6 +544,7 @@ class AdminUpdateProviderAdapter(AdminApiAdapter):
config_in_payload = "config" in update_data
claude_advanced_in_payload = "claude_code_advanced" in update_data
pool_advanced_in_payload = "pool_advanced" in update_data
failover_rules_in_payload = "failover_rules" in update_data
provider_config = (
dict(update_data.pop("config") or {})
if config_in_payload
@@ -518,6 +554,9 @@ class AdminUpdateProviderAdapter(AdminApiAdapter):
update_data.pop("claude_code_advanced") if claude_advanced_in_payload else None
)
pool_advanced = update_data.pop("pool_advanced") if pool_advanced_in_payload else None
failover_rules = (
update_data.pop("failover_rules") if failover_rules_in_payload else None
)
target_provider_type = (
update_data.get("provider_type")
or getattr(provider, "provider_type", None)
@@ -535,6 +574,11 @@ class AdminUpdateProviderAdapter(AdminApiAdapter):
pool_advanced=pool_advanced,
pool_advanced_in_payload=pool_advanced_in_payload,
)
provider_config, config_changed_by_failover = _merge_failover_rules_config(
provider_config=provider_config,
failover_rules=failover_rules,
failover_rules_in_payload=failover_rules_in_payload,
)
config_touched = (
config_in_payload
@@ -542,6 +586,8 @@ class AdminUpdateProviderAdapter(AdminApiAdapter):
or config_changed_by_claude
or pool_advanced_in_payload
or config_changed_by_pool
or failover_rules_in_payload
or config_changed_by_failover
)
if config_touched:
update_data["config"] = provider_config

View File

@@ -18,7 +18,11 @@ from src.core.enums import ProviderBillingType
from src.core.exceptions import InvalidRequestException, NotFoundException
from src.core.logger import logger
from src.database import get_db
from src.models.admin_requests import ClaudeCodeAdvancedConfig, PoolAdvancedConfig
from src.models.admin_requests import (
ClaudeCodeAdvancedConfig,
FailoverRulesConfig,
PoolAdvancedConfig,
)
from src.models.database import (
Model,
Provider,
@@ -282,6 +286,38 @@ def _extract_claude_code_advanced_from_config(
return None
def _extract_failover_rules_from_config(
provider_config: dict[str, Any] | None,
*,
provider_id: str,
) -> FailoverRulesConfig | None:
"""从 Provider.config 中安全提取故障转移规则配置。"""
raw = (provider_config or {}).get("failover_rules")
if raw is None:
return None
if isinstance(raw, FailoverRulesConfig):
return raw
if not isinstance(raw, dict):
logger.warning(
"Provider {} 的 failover_rules 类型无效: {},已忽略",
provider_id,
type(raw).__name__,
)
return None
try:
return FailoverRulesConfig.model_validate(raw)
except Exception as exc:
logger.warning(
"Provider {} 的 failover_rules 配置无效,已忽略: {}",
provider_id,
str(exc),
)
return None
def _build_provider_summary(db: Session, provider: Provider) -> ProviderWithEndpointsSummary:
endpoints = db.query(ProviderEndpoint).filter(ProviderEndpoint.provider_id == provider.id).all()
@@ -402,6 +438,10 @@ def _build_provider_summary(db: Session, provider: Provider) -> ProviderWithEndp
provider_config,
provider_id=str(provider.id),
)
failover_rules = _extract_failover_rules_from_config(
provider_config,
provider_id=str(provider.id),
)
return ProviderWithEndpointsSummary(
id=provider.id,
@@ -425,6 +465,7 @@ def _build_provider_summary(db: Session, provider: Provider) -> ProviderWithEndp
request_timeout=provider.request_timeout,
claude_code_advanced=claude_code_advanced,
pool_advanced=pool_advanced,
failover_rules=failover_rules,
total_endpoints=total_endpoints,
active_endpoints=active_endpoints,
total_keys=total_keys,

View File

@@ -8,10 +8,12 @@ aether-proxy 通过此端点建立 tunnel 连接。
from __future__ import annotations
import asyncio
from datetime import datetime, timezone
from fastapi import APIRouter, WebSocket, WebSocketDisconnect
from src.core.logger import logger
from src.services.proxy_node.health_scheduler import heartbeat_is_stale
from src.services.proxy_node.tunnel_manager import (
TunnelConnection,
get_tunnel_manager,
@@ -20,6 +22,18 @@ from src.services.proxy_node.tunnel_protocol import Frame, MsgType
router = APIRouter()
# Per-node 锁: 防止并发的 connect/disconnect 写入 DB 时出现竞态(后断连覆盖先连接)
_node_status_locks: dict[str, asyncio.Lock] = {}
def _get_node_lock(node_id: str) -> asyncio.Lock:
lock = _node_status_locks.get(node_id)
if lock is None:
lock = asyncio.Lock()
_node_status_locks[node_id] = lock
return lock
# 单帧最大 64 MB -- AI API 请求体可能包含多张 base64 图片,需要足够余量
_MAX_FRAME_SIZE = 64 * 1024 * 1024
@@ -111,10 +125,17 @@ async def proxy_tunnel_ws(ws: WebSocket) -> None:
manager = get_tunnel_manager()
conn = TunnelConnection(node_id, node_name, ws, max_streams=max_streams)
node_lock = _get_node_lock(node_id)
manager.register(conn)
# 更新 DB: tunnel_connected = True
await _update_tunnel_status(node_id, connected=True)
# 在 per-node 锁保护下更新 DB防止并发的 connect/disconnect 写入竞态
async with node_lock:
await _update_tunnel_status(
node_id,
connected=True,
observed_at=datetime.now(timezone.utc),
)
# 启动服务端 ping 任务,防止中间代理因空闲超时关闭连接
ping_task = asyncio.create_task(_ping_loop(conn))
@@ -156,11 +177,20 @@ async def proxy_tunnel_ws(ws: WebSocket) -> None:
logger.error("tunnel WebSocket error for node_id={}: {}", node_id, e)
finally:
ping_task.cancel()
manager.unregister(conn)
if not manager.has_tunnel(node_id):
await _update_tunnel_status(node_id, connected=False, detail=disconnect_reason)
else:
logger.info("tunnel connection closed but pool still active: node_id={}", node_id)
# 在 per-node 锁保护下执行 unregister + 连接池计数检查 + DB 更新,
# 确保整个序列是原子的,避免"断连写 OFFLINE 覆盖新连接写 ONLINE"的竞态
async with node_lock:
manager.unregister(conn)
if manager.connection_count(node_id) == 0:
await _update_tunnel_status(
node_id,
connected=False,
detail=disconnect_reason,
observed_at=datetime.now(timezone.utc),
)
# 不清理锁: asyncio.Lock 极轻量,清理可能导致并发新连接拿到不同锁实例
else:
logger.info("tunnel connection closed but pool still active: node_id={}", node_id)
async def _ping_loop(conn: TunnelConnection) -> None:
@@ -180,13 +210,15 @@ async def _ping_loop(conn: TunnelConnection) -> None:
async def _update_tunnel_status(
node_id: str, *, connected: bool, detail: str | None = None
node_id: str,
*,
connected: bool,
detail: str | None = None,
observed_at: datetime | None = None,
) -> None:
"""更新 ProxyNode 的 tunnel 连接状态并记录事件(在线程池中执行)"""
def _sync_update() -> None:
from datetime import datetime, timezone
from src.database import create_session
from src.models.database import ProxyNode, ProxyNodeEvent, ProxyNodeStatus
@@ -194,20 +226,47 @@ async def _update_tunnel_status(
try:
node = db.query(ProxyNode).filter(ProxyNode.id == node_id).first()
if node:
node.tunnel_connected = connected
now = datetime.now(timezone.utc)
event_time = observed_at or datetime.now(timezone.utc)
last_transition = node.tunnel_connected_at
if last_transition and last_transition.tzinfo is None:
last_transition = last_transition.replace(tzinfo=timezone.utc)
# 忽略乱序的旧事件,避免快速重连时旧状态覆盖新状态
stale_event = bool(last_transition and event_time < last_transition)
if stale_event:
detail_text = f"[stale_ignored] {detail}" if detail else "[stale_ignored]"
db.add(
ProxyNodeEvent(
node_id=node_id,
event_type="connected" if connected else "disconnected",
detail=detail_text,
)
)
db.commit()
return
event_detail = detail
if connected:
node.tunnel_connected_at = now
node.tunnel_connected = True
node.tunnel_connected_at = event_time
node.status = ProxyNodeStatus.ONLINE
else:
node.tunnel_connected_at = now
node.status = ProxyNodeStatus.OFFLINE
# 断连不立即强制 OFFLINE。若心跳仍新鲜可能仍有其他连接存活
# (连接池或跨 worker避免误判写回 OFFLINE
if heartbeat_is_stale(node, event_time):
node.tunnel_connected = False
node.tunnel_connected_at = event_time
node.status = ProxyNodeStatus.OFFLINE
else:
event_detail = (
f"[heartbeat_fresh] {detail}" if detail else "[heartbeat_fresh]"
)
# 记录连接事件
event = ProxyNodeEvent(
node_id=node_id,
event_type="connected" if connected else "disconnected",
detail=detail,
detail=event_detail,
)
db.add(event)
db.commit()