feat(heartbeat): 心跳可靠性增强,原子计数与去重优化

- Rust proxy: 引入 snapshot+ACK 确认机制,心跳未确认时保留快照重发,
  避免指标丢失;添加 heartbeat_session_id 防跨进程去重误判
- Hub transport: Redis SETNX 心跳去重,避免多 worker 重复写库;
  ACK 回显 heartbeat_id 供 Rust 端匹配
- ProxyNodeService.heartbeat: 改用 SQLAlchemy atomic update 原子累加
  指标,避免 ORM read-modify-write 的并发覆盖问题
- 启动顺序修正: tunnel 状态重置移到 Hub 连接建立之前,避免竞态
- OAuth 批量导入: 动态超时(默认30s,走代理60s),Kiro 适配器透传
- 提取 normalize_heartbeat_id 到 tunnel_protocol 共享模块,消除重复
This commit is contained in:
fawney19
2026-03-02 12:27:51 +08:00
parent f3b9f42202
commit f978888759
13 changed files with 507 additions and 71 deletions

View File

@@ -125,6 +125,8 @@ async def _consume_state(redis: Redis, nonce: str) -> ProviderOAuthStateData | N
_PROVIDER_OAUTH_BATCH_TASK_PREFIX = "provider_oauth_batch_task:"
_PROVIDER_OAUTH_BATCH_TASK_TTL_SECONDS = 24 * 3600
_PROVIDER_OAUTH_BATCH_TASK_MAX_ERROR_SAMPLES = 20
_PROVIDER_OAUTH_DEFAULT_TIMEOUT_SECONDS = 30.0
_PROVIDER_OAUTH_BATCH_IMPORT_PROXY_TIMEOUT_SECONDS = 60.0
_PROVIDER_OAUTH_BATCH_TASK_ALLOWED_STATUSES = {
"submitted",
"processing",
@@ -175,7 +177,9 @@ async def _save_batch_task_state(
)
async def _load_batch_task_state(task_id: str, *, redis: Redis | None = None) -> dict[str, Any] | None:
async def _load_batch_task_state(
task_id: str, *, redis: Redis | None = None
) -> dict[str, Any] | None:
now_ts = int(time.time())
_cleanup_in_memory_batch_tasks(now_ts)
redis_client = redis if redis is not None else await get_redis_client(require_redis=False)
@@ -311,6 +315,21 @@ def _resolve_proxy_for_oauth(
return provider_proxy, None
def _resolve_batch_import_timeout_seconds(proxy_config: dict[str, Any] | None) -> float:
"""返回 OAuth 批量导入 token 刷新超时。
走代理链路时延更高,适当放宽超时,减少批量导入误超时。
"""
if not proxy_config:
return _PROVIDER_OAUTH_DEFAULT_TIMEOUT_SECONDS
if isinstance(proxy_config, dict):
if not proxy_config.get("enabled", True):
return _PROVIDER_OAUTH_DEFAULT_TIMEOUT_SECONDS
return _PROVIDER_OAUTH_BATCH_IMPORT_PROXY_TIMEOUT_SECONDS
# 兼容历史数据:非 dict 但存在代理配置时同样使用放宽超时
return _PROVIDER_OAUTH_BATCH_IMPORT_PROXY_TIMEOUT_SECONDS
def _pkce_s256(verifier: str) -> str:
digest = hashlib.sha256(verifier.encode("utf-8")).digest()
return base64.urlsafe_b64encode(digest).decode("utf-8").rstrip("=")
@@ -1750,6 +1769,7 @@ async def _batch_import_standard_oauth_internal(
) -> BatchImportResponse:
"""标准 OAuth Provider 批量导入(不含 Kiro"""
template = _require_oauth_template(provider_type)
timeout_seconds = _resolve_batch_import_timeout_seconds(proxy_config)
tokens = _parse_tokens_input(raw_credentials)
if not tokens:
@@ -1811,7 +1831,7 @@ async def _batch_import_standard_oauth_internal(
data=data,
json_body=json_body,
proxy_config=proxy_config,
timeout_seconds=30.0,
timeout_seconds=timeout_seconds,
)
except Exception as exc:
result_item = BatchImportResultItem(
@@ -2109,7 +2129,9 @@ async def _run_batch_import_task(
try:
await _save_batch_task_state(task_id, state, redis=redis)
except Exception as exc:
logger.debug("[BATCH_IMPORT_TASK] save progress failed (task_id={}): {}", task_id, exc)
logger.debug(
"[BATCH_IMPORT_TASK] save progress failed (task_id={}): {}", task_id, exc
)
if provider_type == ProviderType.KIRO.value:
result = await _batch_import_kiro_internal(
@@ -2264,6 +2286,7 @@ async def _batch_import_kiro_internal(
credentials = _parse_kiro_import_input(raw_credentials)
if not credentials:
raise InvalidRequestException("未找到有效的凭据数据")
timeout_seconds = _resolve_batch_import_timeout_seconds(proxy_config)
api_formats = _get_provider_api_formats(provider)
@@ -2297,7 +2320,11 @@ async def _batch_import_kiro_internal(
cfg.provider_type = ProviderType.KIRO.value
try:
access_token, new_cfg = await refresh_access_token(cfg, proxy_config=proxy_config)
access_token, new_cfg = await refresh_access_token(
cfg,
proxy_config=proxy_config,
timeout_seconds=timeout_seconds,
)
except Exception as exc:
result_item = BatchImportResultItem(
index=idx,

View File

@@ -87,18 +87,6 @@ async def _on_startup() -> None:
config.worker_processes,
)
# 启动时主动建立 /worker 长连接:
# - 立即接收 NODE_STATUS 广播避免“proxy 已连但 UI 仍显示离线”的窗口期
# - 确保后续 tunnel 请求不需要首请求触发懒连接
from src.services.proxy_node.hub_transport import get_hub_connection_manager
try:
await get_hub_connection_manager().ensure_connected()
logger.info("Hub worker channel initialized on startup")
except Exception as e:
# ensure_connected 失败时内部会启动重连循环,这里仅记录告警不阻塞启动
logger.warning("Hub worker channel init failed, reconnecting in background: {}", e)
from src.clients import get_redis_client
redis_client = await get_redis_client()
@@ -110,11 +98,23 @@ async def _on_startup() -> None:
# 仅 leader worker 执行启动重置,避免多 worker 并发启动/重启时
# 把其他 worker 已建立的 tunnel 状态错误重置为 OFFLINE。
_reset_tunnel_connected_on_startup()
logger.info("启动 ProxyNode 心跳检测调度器...")
await proxy_node_health_scheduler.start()
else:
logger.info("检测到其他 worker 已运行 ProxyNode 心跳检测,本实例跳过")
# 在可能的状态重置之后再建立 /worker 长连接,避免“先同步在线,再被重置离线”的竞态。
from src.services.proxy_node.hub_transport import get_hub_connection_manager
try:
await get_hub_connection_manager().ensure_connected()
logger.info("Hub worker channel initialized on startup")
except Exception as e:
# ensure_connected 失败时内部会启动重连循环,这里仅记录告警不阻塞启动
logger.warning("Hub worker channel init failed, reconnecting in background: {}", e)
if active:
logger.info("启动 ProxyNode 心跳检测调度器...")
await proxy_node_health_scheduler.start()
async def _on_shutdown() -> None:
"""优雅关闭 tunnel 连接并停止心跳检测调度器"""

View File

@@ -295,11 +295,20 @@ async def refresh_access_token(
cfg: KiroAuthConfig,
*,
proxy_config: dict[str, Any] | None,
timeout_seconds: float = 30.0,
) -> tuple[str, KiroAuthConfig]:
method = (cfg.auth_method or "social").strip().lower()
if method == "idc":
return await refresh_idc_token(cfg, proxy_config=proxy_config)
return await refresh_social_token(cfg, proxy_config=proxy_config)
return await refresh_idc_token(
cfg,
proxy_config=proxy_config,
timeout_seconds=timeout_seconds,
)
return await refresh_social_token(
cfg,
proxy_config=proxy_config,
timeout_seconds=timeout_seconds,
)
__all__ = [

View File

@@ -20,7 +20,7 @@ from src.core.logger import logger
from .hub_config import HubConfig, get_hub_config
from .tunnel_manager import TunnelStreamError, _StreamState
from .tunnel_protocol import Frame, FrameFlags, MsgType
from .tunnel_protocol import Frame, FrameFlags, MsgType, normalize_heartbeat_id
if TYPE_CHECKING:
from collections.abc import AsyncGenerator, Coroutine
@@ -28,6 +28,7 @@ if TYPE_CHECKING:
_TUNNEL_COMPRESS_MIN_SIZE = 512
_RECONNECT_DELAYS_SECONDS: tuple[float, ...] = (0.0, 0.05, 0.1, 0.2, 0.5, 1.0, 2.0, 5.0)
_HEARTBEAT_DEDUP_TTL_SECONDS = 600
_HOP_BY_HOP_HEADERS = frozenset(
{
@@ -317,6 +318,36 @@ class HubConnectionManager:
data = {}
node_id = str(data.get("node_id") or "").strip()
heartbeat_session_id = str(data.get("heartbeat_session_id") or "").strip()
if len(heartbeat_session_id) > 128:
heartbeat_session_id = heartbeat_session_id[:128]
heartbeat_id = normalize_heartbeat_id(data.get("heartbeat_id"))
ack: dict[str, object] = {}
if heartbeat_id is not None:
ack["heartbeat_id"] = heartbeat_id
should_process = True
if node_id and heartbeat_id is not None:
if heartbeat_session_id:
dedup_key = f"hub:heartbeat:{node_id}:{heartbeat_session_id}:{heartbeat_id}"
else:
dedup_key = f"hub:heartbeat:{node_id}:{heartbeat_id}"
try:
from src.clients import get_redis_client
redis = await get_redis_client()
if redis:
acquired = await redis.set(
dedup_key,
"1",
ex=_HEARTBEAT_DEDUP_TTL_SECONDS,
nx=True,
)
if not acquired:
should_process = False
except Exception:
# Redis 不可用时降级为不去重,避免心跳链路阻塞
pass
def _sync_heartbeat() -> dict[str, object]:
from src.database import create_session
@@ -345,11 +376,11 @@ class HubConnectionManager:
finally:
db.close()
try:
ack = await asyncio.to_thread(_sync_heartbeat)
except Exception as e:
logger.warning("hub heartbeat DB update failed: {}", e)
ack = {}
if should_process:
try:
ack.update(await asyncio.to_thread(_sync_heartbeat))
except Exception as e:
logger.warning("hub heartbeat DB update failed: {}", e)
try:
await self._send_frame(

View File

@@ -14,6 +14,7 @@ from typing import Any
from urllib.parse import urlparse
import httpx
from sqlalchemy import update
from sqlalchemy.orm import Session
from src.core.exceptions import InvalidRequestException, NotFoundException
@@ -338,36 +339,43 @@ class ProxyNodeService:
)
now = datetime.now(timezone.utc)
values: dict[str, Any] = {"last_heartbeat_at": now}
# 心跳通过 tunnel 连接传输,能收到心跳说明 tunnel 一定连通。
# 如果状态不是 ONLINE 或 tunnel_connected 不一致(例如并发写入覆盖),修正状态
# 状态不一致(例如并发写入覆盖),修正为 ONLINE
if node.status != ProxyNodeStatus.ONLINE or not node.tunnel_connected:
node.status = ProxyNodeStatus.ONLINE
node.tunnel_connected = True
node.tunnel_connected_at = now
node.updated_at = now
node.last_heartbeat_at = now
values["status"] = ProxyNodeStatus.ONLINE
values["tunnel_connected"] = True
values["tunnel_connected_at"] = now
values["updated_at"] = now
if heartbeat_interval is not None:
node.heartbeat_interval = heartbeat_interval
values["heartbeat_interval"] = heartbeat_interval
# 实时快照指标 -- 直接覆盖
if active_connections is not None:
node.active_connections = active_connections
values["active_connections"] = active_connections
if avg_latency_ms is not None:
node.avg_latency_ms = avg_latency_ms
values["avg_latency_ms"] = avg_latency_ms
# 区间增量指标 -- 累加到累计值
# 区间增量指标 -- 使用数据库原子自增,避免并发心跳读改写丢增量
if total_requests is not None and total_requests > 0:
node.total_requests = (node.total_requests or 0) + total_requests
values["total_requests"] = ProxyNode.total_requests + int(total_requests)
if failed_requests is not None and failed_requests > 0:
node.failed_requests = (node.failed_requests or 0) + failed_requests
values["failed_requests"] = ProxyNode.failed_requests + int(failed_requests)
if dns_failures is not None and dns_failures > 0:
node.dns_failures = (node.dns_failures or 0) + dns_failures
values["dns_failures"] = ProxyNode.dns_failures + int(dns_failures)
if stream_errors is not None and stream_errors > 0:
node.stream_errors = (node.stream_errors or 0) + stream_errors
values["stream_errors"] = ProxyNode.stream_errors + int(stream_errors)
db.execute(update(ProxyNode).where(ProxyNode.id == node_id).values(**values))
db.commit()
db.refresh(node)
return node
db.expire_all()
refreshed = db.query(ProxyNode).filter(ProxyNode.id == node_id).first()
if not refreshed:
raise NotFoundException(f"ProxyNode {node_id} 不存在", "proxy_node")
return refreshed
@staticmethod
def unregister_node(db: Session, *, node_id: str) -> ProxyNode:

View File

@@ -20,7 +20,7 @@ from starlette.websockets import WebSocket, WebSocketState
from src.core.logger import logger
from .tunnel_protocol import Frame, FrameFlags, MsgType
from .tunnel_protocol import Frame, FrameFlags, MsgType, normalize_heartbeat_id
# 隧道帧压缩的最小 payload 大小(字节)
# 小于此值的帧压缩收益不大,反而增加 CPU 开销
@@ -463,6 +463,10 @@ class TunnelManager:
data = json.loads(frame.payload) if frame.payload else {}
except Exception:
data = {}
heartbeat_id = normalize_heartbeat_id(data.get("heartbeat_id"))
ack: dict[str, Any] = {}
if heartbeat_id is not None:
ack["heartbeat_id"] = heartbeat_id
def _sync_heartbeat() -> dict[str, Any]:
from src.database import create_session
@@ -489,10 +493,9 @@ class TunnelManager:
db.close()
try:
ack = await asyncio.to_thread(_sync_heartbeat)
ack.update(await asyncio.to_thread(_sync_heartbeat))
except Exception as e:
logger.warning("tunnel heartbeat DB update failed: {}", e)
ack = {}
try:
await conn.send_frame(Frame(0, MsgType.HEARTBEAT_ACK, 0, json.dumps(ack).encode()))

View File

@@ -11,6 +11,8 @@ import struct
from enum import IntEnum
from typing import Self
_U64_MAX = (1 << 64) - 1
HEADER_SIZE = 10 # 4 + 1 + 1 + 4 bytes
@@ -98,3 +100,25 @@ class Frame:
f"Frame(stream={self.stream_id}, type={self.msg_type.name}, "
f"flags=0x{self.flags:02x}, payload_len={len(self.payload)})"
)
def normalize_heartbeat_id(value: object) -> int | None:
"""Normalize heartbeat_id to u64-compatible int for Rust ACK parsing."""
if isinstance(value, bool) or value is None:
return None
parsed: int | None = None
if isinstance(value, int):
parsed = value
elif isinstance(value, float):
if value.is_integer():
parsed = int(value)
elif isinstance(value, str):
stripped = value.strip()
if stripped.isdigit():
try:
parsed = int(stripped)
except ValueError:
parsed = None
if parsed is None or parsed < 0 or parsed > _U64_MAX:
return None
return parsed