mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-09 04:30:20 +08:00
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:
+2
-7
@@ -36,15 +36,10 @@ ADMIN_PASSWORD=admin123456
|
||||
# APP_IMAGE=ghcr.io/fawney19/aether:latest
|
||||
|
||||
# Gunicorn Worker 数量(默认 2)
|
||||
# Docker 部署下 Tunnel Hub 为容器内部固定服务,可安全使用多 worker。
|
||||
# 非 Docker 运行时若使用 ProxyNode tunnel,建议设置为 1。
|
||||
# Tunnel 请求统一经 Hub 转发,可安全使用多 worker。
|
||||
# 非 Docker 运行时若使用 ProxyNode tunnel,请确保 aether-hub 可达(默认 ws://127.0.0.1:8085)。
|
||||
# GUNICORN_WORKERS=2
|
||||
|
||||
# 本地构建 app 镜像时使用的 Hub 二进制镜像
|
||||
# 默认 aether-hub:local(需先本地构建 aether-hub/Dockerfile)
|
||||
# 也可改为 ghcr.io/fawney19/aether-hub:latest
|
||||
# HUB_BINARY_IMAGE=aether-hub:local
|
||||
|
||||
# Gunicorn Max Requests(默认 4000)
|
||||
# Worker 处理指定数量请求后自动重启,防止内存泄漏
|
||||
# max-requests-jitter 会自动设置为 MAX_REQUESTS/20 (5%)
|
||||
|
||||
@@ -3,6 +3,8 @@
|
||||
use std::sync::atomic::Ordering;
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
use std::time::SystemTime;
|
||||
use std::time::UNIX_EPOCH;
|
||||
|
||||
use bytes::Bytes;
|
||||
use tokio::sync::watch;
|
||||
@@ -16,6 +18,11 @@ use crate::state::ServerContext;
|
||||
use super::protocol::{Frame, MsgType};
|
||||
use super::writer::FrameSender;
|
||||
|
||||
enum AckDecision {
|
||||
Accept(Option<u64>),
|
||||
Ignore,
|
||||
}
|
||||
|
||||
/// Handle for the dispatcher to forward HeartbeatAck frames.
|
||||
#[derive(Clone)]
|
||||
pub struct HeartbeatHandle {
|
||||
@@ -37,6 +44,15 @@ pub fn spawn_noop() -> HeartbeatHandle {
|
||||
HeartbeatHandle { ack_tx }
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, Default)]
|
||||
struct HeartbeatSnapshot {
|
||||
requests: u64,
|
||||
latency_ns: u64,
|
||||
failed: u64,
|
||||
dns_failures: u64,
|
||||
stream_errors: u64,
|
||||
}
|
||||
|
||||
/// Spawn the heartbeat task. Returns a handle for forwarding ACKs.
|
||||
pub fn spawn(
|
||||
_config: Arc<Config>,
|
||||
@@ -50,6 +66,19 @@ pub fn spawn(
|
||||
// Read initial interval from dynamic config (may be updated by remote config).
|
||||
let initial_interval = Duration::from_secs(server.dynamic.load().heartbeat_interval);
|
||||
let mut current_interval = initial_interval;
|
||||
// At most one in-flight heartbeat snapshot is tracked at a time.
|
||||
// Snapshot is only cleared after receiving an ACK, which avoids losing
|
||||
// interval counters when ACK/frame delivery is temporarily unstable.
|
||||
let mut pending: Option<(u64, HeartbeatSnapshot)> = None;
|
||||
let mut next_heartbeat_id: u64 = 1;
|
||||
let heartbeat_session_id = format!(
|
||||
"{}-{}",
|
||||
std::process::id(),
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
.as_nanos()
|
||||
);
|
||||
|
||||
// Skip first immediate tick by sleeping first.
|
||||
tokio::time::sleep(current_interval).await;
|
||||
@@ -57,9 +86,30 @@ pub fn spawn(
|
||||
loop {
|
||||
tokio::select! {
|
||||
_ = tokio::time::sleep(current_interval) => {
|
||||
let payload = build_heartbeat_payload(&server);
|
||||
let (heartbeat_id, snapshot) = if let Some((id, snap)) = pending {
|
||||
(id, snap)
|
||||
} else {
|
||||
let snap = collect_snapshot(&server);
|
||||
let id = next_heartbeat_id;
|
||||
next_heartbeat_id = next_heartbeat_id.wrapping_add(1);
|
||||
if next_heartbeat_id == 0 {
|
||||
next_heartbeat_id = 1;
|
||||
}
|
||||
pending = Some((id, snap));
|
||||
(id, snap)
|
||||
};
|
||||
|
||||
let payload = build_heartbeat_payload(
|
||||
&server,
|
||||
&heartbeat_session_id,
|
||||
heartbeat_id,
|
||||
snapshot
|
||||
);
|
||||
let frame = Frame::control(MsgType::HeartbeatData, payload);
|
||||
if frame_tx.send(frame).await.is_err() {
|
||||
if let Some((_, snap)) = pending.take() {
|
||||
restore_snapshot(&server, snap);
|
||||
}
|
||||
break; // Writer closed
|
||||
}
|
||||
debug!("sent heartbeat data");
|
||||
@@ -79,10 +129,30 @@ pub fn spawn(
|
||||
}
|
||||
}
|
||||
Some(ack_payload) = ack_rx.recv() => {
|
||||
handle_ack(&server, &ack_payload);
|
||||
match handle_ack(&server, &ack_payload) {
|
||||
AckDecision::Accept(ack_id) => {
|
||||
if let Some((pending_id, _)) = pending {
|
||||
match ack_id {
|
||||
Some(id) if id == pending_id => {
|
||||
pending = None;
|
||||
}
|
||||
None => {
|
||||
// Backward-compatible with servers that don't echo
|
||||
// heartbeat_id in ACK payload yet.
|
||||
pending = None;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
AckDecision::Ignore => {}
|
||||
}
|
||||
}
|
||||
_ = shutdown.changed() => {
|
||||
debug!("heartbeat task shutting down");
|
||||
if let Some((_, snap)) = pending.take() {
|
||||
restore_snapshot(&server, snap);
|
||||
}
|
||||
break;
|
||||
}
|
||||
}
|
||||
@@ -92,36 +162,81 @@ pub fn spawn(
|
||||
HeartbeatHandle { ack_tx }
|
||||
}
|
||||
|
||||
fn build_heartbeat_payload(server: &ServerContext) -> Bytes {
|
||||
fn collect_snapshot(server: &ServerContext) -> HeartbeatSnapshot {
|
||||
HeartbeatSnapshot {
|
||||
requests: server.metrics.total_requests.swap(0, Ordering::AcqRel),
|
||||
latency_ns: server.metrics.total_latency_ns.swap(0, Ordering::AcqRel),
|
||||
failed: server.metrics.failed_requests.swap(0, Ordering::AcqRel),
|
||||
dns_failures: server.metrics.dns_failures.swap(0, Ordering::AcqRel),
|
||||
stream_errors: server.metrics.stream_errors.swap(0, Ordering::AcqRel),
|
||||
}
|
||||
}
|
||||
|
||||
fn restore_snapshot(server: &ServerContext, snap: HeartbeatSnapshot) {
|
||||
if snap.requests > 0 {
|
||||
server
|
||||
.metrics
|
||||
.total_requests
|
||||
.fetch_add(snap.requests, Ordering::Release);
|
||||
}
|
||||
if snap.latency_ns > 0 {
|
||||
server
|
||||
.metrics
|
||||
.total_latency_ns
|
||||
.fetch_add(snap.latency_ns, Ordering::Release);
|
||||
}
|
||||
if snap.failed > 0 {
|
||||
server
|
||||
.metrics
|
||||
.failed_requests
|
||||
.fetch_add(snap.failed, Ordering::Release);
|
||||
}
|
||||
if snap.dns_failures > 0 {
|
||||
server
|
||||
.metrics
|
||||
.dns_failures
|
||||
.fetch_add(snap.dns_failures, Ordering::Release);
|
||||
}
|
||||
if snap.stream_errors > 0 {
|
||||
server
|
||||
.metrics
|
||||
.stream_errors
|
||||
.fetch_add(snap.stream_errors, Ordering::Release);
|
||||
}
|
||||
}
|
||||
|
||||
fn build_heartbeat_payload(
|
||||
server: &ServerContext,
|
||||
heartbeat_session_id: &str,
|
||||
heartbeat_id: u64,
|
||||
snapshot: HeartbeatSnapshot,
|
||||
) -> Bytes {
|
||||
let node_id = server.node_id.read().unwrap().clone();
|
||||
|
||||
let interval_requests = server.metrics.total_requests.swap(0, Ordering::AcqRel);
|
||||
let interval_latency_ns = server.metrics.total_latency_ns.swap(0, Ordering::AcqRel);
|
||||
let interval_failed = server.metrics.failed_requests.swap(0, Ordering::AcqRel);
|
||||
let interval_dns_failures = server.metrics.dns_failures.swap(0, Ordering::AcqRel);
|
||||
let interval_stream_errors = server.metrics.stream_errors.swap(0, Ordering::AcqRel);
|
||||
let avg_latency_ms = if interval_requests > 0 {
|
||||
Some(interval_latency_ns as f64 / interval_requests as f64 / 1_000_000.0)
|
||||
let avg_latency_ms = if snapshot.requests > 0 {
|
||||
Some(snapshot.latency_ns as f64 / snapshot.requests as f64 / 1_000_000.0)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
let payload = serde_json::json!({
|
||||
"node_id": node_id,
|
||||
"heartbeat_session_id": heartbeat_session_id,
|
||||
"heartbeat_id": heartbeat_id,
|
||||
"active_connections": server.active_connections.load(Ordering::Acquire),
|
||||
"total_requests": interval_requests,
|
||||
"total_requests": snapshot.requests,
|
||||
"avg_latency_ms": avg_latency_ms,
|
||||
"failed_requests": interval_failed,
|
||||
"dns_failures": interval_dns_failures,
|
||||
"stream_errors": interval_stream_errors,
|
||||
"failed_requests": snapshot.failed,
|
||||
"dns_failures": snapshot.dns_failures,
|
||||
"stream_errors": snapshot.stream_errors,
|
||||
});
|
||||
|
||||
Bytes::from(serde_json::to_vec(&payload).unwrap_or_default())
|
||||
}
|
||||
|
||||
fn handle_ack(server: &ServerContext, payload: &[u8]) {
|
||||
fn handle_ack(server: &ServerContext, payload: &[u8]) -> AckDecision {
|
||||
if payload.is_empty() {
|
||||
return;
|
||||
return AckDecision::Accept(None);
|
||||
}
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
@@ -130,6 +245,8 @@ fn handle_ack(server: &ServerContext, payload: &[u8]) {
|
||||
remote_config: Option<RemoteConfig>,
|
||||
#[serde(default)]
|
||||
config_version: u64,
|
||||
#[serde(default)]
|
||||
heartbeat_id: Option<u64>,
|
||||
}
|
||||
|
||||
match serde_json::from_slice::<AckPayload>(payload) {
|
||||
@@ -137,9 +254,11 @@ fn handle_ack(server: &ServerContext, payload: &[u8]) {
|
||||
if let Some(ref rc) = ack.remote_config {
|
||||
runtime::apply_remote_config(&server.dynamic, rc, ack.config_version);
|
||||
}
|
||||
AckDecision::Accept(ack.heartbeat_id)
|
||||
}
|
||||
Err(e) => {
|
||||
warn!(error = %e, "failed to parse heartbeat ACK");
|
||||
AckDecision::Ignore
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -253,8 +253,7 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn request_meta_accepts_integer_like_float_timeout() {
|
||||
let raw =
|
||||
br#"{"method":"GET","url":"https://example.com","headers":{},"timeout":15.0}"#;
|
||||
let raw = br#"{"method":"GET","url":"https://example.com","headers":{},"timeout":15.0}"#;
|
||||
let meta: RequestMeta = serde_json::from_slice(raw).expect("parse request meta");
|
||||
assert_eq!(meta.timeout, 15);
|
||||
}
|
||||
|
||||
@@ -60,7 +60,7 @@ network_access = "enabled"
|
||||
disable_response_storage = true
|
||||
|
||||
[model_providers.aether]
|
||||
name = "aether"
|
||||
name = "OpenAI"
|
||||
base_url = "${baseUrl.value}/v1"
|
||||
wire_api = "responses"
|
||||
requires_openai_auth = true`)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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 连接并停止心跳检测调度器"""
|
||||
|
||||
@@ -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__ = [
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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()))
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from src.api.admin.provider_oauth import (
|
||||
_PROVIDER_OAUTH_BATCH_IMPORT_PROXY_TIMEOUT_SECONDS,
|
||||
_PROVIDER_OAUTH_DEFAULT_TIMEOUT_SECONDS,
|
||||
_resolve_batch_import_timeout_seconds,
|
||||
)
|
||||
|
||||
|
||||
def test_batch_import_timeout_without_proxy_uses_default() -> None:
|
||||
assert _resolve_batch_import_timeout_seconds(None) == _PROVIDER_OAUTH_DEFAULT_TIMEOUT_SECONDS
|
||||
assert (
|
||||
_resolve_batch_import_timeout_seconds({"enabled": False})
|
||||
== _PROVIDER_OAUTH_DEFAULT_TIMEOUT_SECONDS
|
||||
)
|
||||
|
||||
|
||||
def test_batch_import_timeout_with_proxy_uses_extended_value() -> None:
|
||||
assert (
|
||||
_resolve_batch_import_timeout_seconds({"node_id": "node-1", "enabled": True})
|
||||
== _PROVIDER_OAUTH_BATCH_IMPORT_PROXY_TIMEOUT_SECONDS
|
||||
)
|
||||
assert (
|
||||
_resolve_batch_import_timeout_seconds({"node_id": "node-1"})
|
||||
== _PROVIDER_OAUTH_BATCH_IMPORT_PROXY_TIMEOUT_SECONDS
|
||||
)
|
||||
assert (
|
||||
_resolve_batch_import_timeout_seconds("http://legacy-proxy") # type: ignore[arg-type]
|
||||
== _PROVIDER_OAUTH_BATCH_IMPORT_PROXY_TIMEOUT_SECONDS
|
||||
)
|
||||
@@ -0,0 +1,191 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from src.services.proxy_node.hub_config import HubConfig
|
||||
from src.services.proxy_node.hub_transport import HubConnectionManager
|
||||
from src.services.proxy_node.tunnel_protocol import Frame, MsgType
|
||||
|
||||
|
||||
class _StubRedis:
|
||||
def __init__(self, *, set_result: bool = True, set_exc: Exception | None = None) -> None:
|
||||
self.set_result = set_result
|
||||
self.set_exc = set_exc
|
||||
self.calls: list[tuple[str, str, int | None, bool | None]] = []
|
||||
|
||||
async def set(
|
||||
self,
|
||||
key: str,
|
||||
value: str,
|
||||
ex: int | None = None,
|
||||
nx: bool | None = None,
|
||||
) -> bool:
|
||||
self.calls.append((key, value, ex, nx))
|
||||
if self.set_exc is not None:
|
||||
raise self.set_exc
|
||||
return self.set_result
|
||||
|
||||
|
||||
def _build_manager() -> HubConnectionManager:
|
||||
return HubConnectionManager(
|
||||
HubConfig(
|
||||
enabled=True,
|
||||
url="ws://127.0.0.1:8085",
|
||||
connect_timeout_seconds=1.0,
|
||||
ping_interval_seconds=1.0,
|
||||
send_timeout_seconds=1.0,
|
||||
max_streams=16,
|
||||
max_frame_size=1024 * 1024,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _heartbeat_frame(payload: dict[str, Any]) -> Frame:
|
||||
return Frame(0, MsgType.HEARTBEAT_DATA, 0, json.dumps(payload).encode("utf-8"))
|
||||
|
||||
|
||||
def _decode_ack(frame: Frame) -> dict[str, Any]:
|
||||
assert frame.msg_type == MsgType.HEARTBEAT_ACK
|
||||
if not frame.payload:
|
||||
return {}
|
||||
return json.loads(frame.payload.decode("utf-8"))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_heartbeat_normalizes_id_and_updates_db(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
manager = _build_manager()
|
||||
captured_frames: list[Frame] = []
|
||||
heartbeat_calls: list[dict[str, Any]] = []
|
||||
|
||||
async def _fake_send_frame(frame: Frame) -> None:
|
||||
captured_frames.append(frame)
|
||||
|
||||
class _FakeSession:
|
||||
def close(self) -> None:
|
||||
return
|
||||
|
||||
async def _fake_get_redis_client(*, require_redis: bool = False) -> _StubRedis:
|
||||
_ = require_redis
|
||||
return redis
|
||||
|
||||
def _fake_heartbeat(db: Any, **kwargs: Any) -> Any:
|
||||
heartbeat_calls.append(kwargs)
|
||||
return SimpleNamespace(remote_config={"heartbeat_interval": 8}, config_version=5)
|
||||
|
||||
redis = _StubRedis(set_result=True)
|
||||
monkeypatch.setattr(manager, "_send_frame", _fake_send_frame)
|
||||
monkeypatch.setattr("src.database.create_session", lambda: _FakeSession())
|
||||
monkeypatch.setattr("src.clients.get_redis_client", _fake_get_redis_client)
|
||||
monkeypatch.setattr(
|
||||
"src.services.proxy_node.service.ProxyNodeService.heartbeat", _fake_heartbeat
|
||||
)
|
||||
|
||||
await manager._handle_heartbeat(
|
||||
_heartbeat_frame(
|
||||
{
|
||||
"node_id": "node-1",
|
||||
"heartbeat_session_id": "sess-1",
|
||||
"heartbeat_id": 15.0,
|
||||
"active_connections": 3,
|
||||
"total_requests": 10,
|
||||
"failed_requests": 1,
|
||||
"dns_failures": 2,
|
||||
"stream_errors": 0,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
assert len(redis.calls) == 1
|
||||
assert redis.calls[0][0] == "hub:heartbeat:node-1:sess-1:15"
|
||||
assert heartbeat_calls and heartbeat_calls[0]["node_id"] == "node-1"
|
||||
assert len(captured_frames) == 1
|
||||
|
||||
ack = _decode_ack(captured_frames[0])
|
||||
assert ack["heartbeat_id"] == 15
|
||||
assert isinstance(ack["heartbeat_id"], int)
|
||||
assert ack["remote_config"] == {"heartbeat_interval": 8}
|
||||
assert ack["config_version"] == 5
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_heartbeat_duplicate_skips_db_update(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
manager = _build_manager()
|
||||
captured_frames: list[Frame] = []
|
||||
heartbeat_called = False
|
||||
|
||||
async def _fake_send_frame(frame: Frame) -> None:
|
||||
captured_frames.append(frame)
|
||||
|
||||
async def _fake_get_redis_client(*, require_redis: bool = False) -> _StubRedis:
|
||||
_ = require_redis
|
||||
return redis
|
||||
|
||||
def _fake_heartbeat(db: Any, **kwargs: Any) -> Any:
|
||||
_ = db, kwargs
|
||||
nonlocal heartbeat_called
|
||||
heartbeat_called = True
|
||||
return SimpleNamespace(remote_config={"heartbeat_interval": 8}, config_version=5)
|
||||
|
||||
redis = _StubRedis(set_result=False)
|
||||
monkeypatch.setattr(manager, "_send_frame", _fake_send_frame)
|
||||
monkeypatch.setattr("src.clients.get_redis_client", _fake_get_redis_client)
|
||||
monkeypatch.setattr(
|
||||
"src.services.proxy_node.service.ProxyNodeService.heartbeat", _fake_heartbeat
|
||||
)
|
||||
|
||||
await manager._handle_heartbeat(
|
||||
_heartbeat_frame({"node_id": "node-1", "heartbeat_id": 77, "total_requests": 20})
|
||||
)
|
||||
|
||||
assert len(redis.calls) == 1
|
||||
assert heartbeat_called is False
|
||||
assert len(captured_frames) == 1
|
||||
ack = _decode_ack(captured_frames[0])
|
||||
assert ack == {"heartbeat_id": 77}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_heartbeat_redis_error_falls_back_to_db_update(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
manager = _build_manager()
|
||||
captured_frames: list[Frame] = []
|
||||
heartbeat_calls: list[dict[str, Any]] = []
|
||||
|
||||
async def _fake_send_frame(frame: Frame) -> None:
|
||||
captured_frames.append(frame)
|
||||
|
||||
class _FakeSession:
|
||||
def close(self) -> None:
|
||||
return
|
||||
|
||||
async def _fake_get_redis_client(*, require_redis: bool = False) -> _StubRedis:
|
||||
_ = require_redis
|
||||
return redis
|
||||
|
||||
def _fake_heartbeat(db: Any, **kwargs: Any) -> Any:
|
||||
heartbeat_calls.append(kwargs)
|
||||
return SimpleNamespace(remote_config=None, config_version=0)
|
||||
|
||||
redis = _StubRedis(set_exc=RuntimeError("redis unavailable"))
|
||||
monkeypatch.setattr(manager, "_send_frame", _fake_send_frame)
|
||||
monkeypatch.setattr("src.database.create_session", lambda: _FakeSession())
|
||||
monkeypatch.setattr("src.clients.get_redis_client", _fake_get_redis_client)
|
||||
monkeypatch.setattr(
|
||||
"src.services.proxy_node.service.ProxyNodeService.heartbeat", _fake_heartbeat
|
||||
)
|
||||
|
||||
await manager._handle_heartbeat(
|
||||
_heartbeat_frame({"node_id": "node-1", "heartbeat_id": 99, "total_requests": 1})
|
||||
)
|
||||
|
||||
assert heartbeat_calls and heartbeat_calls[0]["total_requests"] == 1
|
||||
assert len(captured_frames) == 1
|
||||
ack = _decode_ack(captured_frames[0])
|
||||
assert ack == {"heartbeat_id": 99}
|
||||
Reference in New Issue
Block a user