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
+2 -7
View File
@@ -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%)
+135 -16
View File
@@ -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
}
}
}
+1 -2
View File
@@ -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);
}
+1 -1
View File
@@ -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`)
+31 -4
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,
+14 -14
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 连接并停止心跳检测调度器"""
@@ -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__ = [
+37 -6
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(
+24 -16
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:
+6 -3
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()))
@@ -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
)
+191
View File
@@ -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}