feat(hub,stability): bounded outbound queue、worker liveness 检测、事件循环 watchdog 及 DB 操作异步化

- aether-hub: unbounded channel 改为 bounded channel (BoundedOutbound),队列满时标记拥塞并主动关闭连接,防止内存无限增长
- aether-hub: worker idle timeout 从命令行参数改为基于心跳的 liveness 检测,默认 60 秒
- aether-hub: 新增 ConnConfig 统一管理连接配置,新增 outbound_queue_capacity 参数
- hub_transport: 新增事件循环 watchdog,检测 lag 超过阈值时临时降级暂停新流
- gunicorn_conf: 启用 faulthandler,worker abort 时自动 dump 全部线程栈用于诊断
- health/endpoint_checker/recording: 同步 DB 操作移至 asyncio.to_thread,避免阻塞事件循环
- Dockerfile: 移除 --worker-idle-timeout 0 命令行参数,改由环境变量和默认值控制
This commit is contained in:
fawney19
2026-03-12 12:19:04 +08:00
parent 4d338ebd3d
commit 71ae1a2307
14 changed files with 676 additions and 280 deletions

View File

@@ -202,7 +202,7 @@ RUN printf '%s\n' \
'environment=PYTHONUNBUFFERED=1,PYTHONIOENCODING=utf-8,LANG=C.UTF-8,LC_ALL=C.UTF-8,DOCKER_CONTAINER=true' \ 'environment=PYTHONUNBUFFERED=1,PYTHONIOENCODING=utf-8,LANG=C.UTF-8,LC_ALL=C.UTF-8,DOCKER_CONTAINER=true' \
'' \ '' \
'[program:tunnel-hub]' \ '[program:tunnel-hub]' \
'command=/usr/local/bin/aether-hub --bind 0.0.0.0:8085 --worker-idle-timeout 0' \ 'command=/usr/local/bin/aether-hub --bind 0.0.0.0:8085' \
'autostart=true' \ 'autostart=true' \
'autorestart=true' \ 'autorestart=true' \
'stdout_logfile=/dev/stdout' \ 'stdout_logfile=/dev/stdout' \

View File

@@ -213,7 +213,7 @@ RUN printf '%s\n' \
'environment=PYTHONUNBUFFERED=1,PYTHONIOENCODING=utf-8,LANG=C.UTF-8,LC_ALL=C.UTF-8,DOCKER_CONTAINER=true' \ 'environment=PYTHONUNBUFFERED=1,PYTHONIOENCODING=utf-8,LANG=C.UTF-8,LC_ALL=C.UTF-8,DOCKER_CONTAINER=true' \
'' \ '' \
'[program:tunnel-hub]' \ '[program:tunnel-hub]' \
'command=/usr/local/bin/aether-hub --bind 0.0.0.0:8085 --worker-idle-timeout 0' \ 'command=/usr/local/bin/aether-hub --bind 0.0.0.0:8085' \
'autostart=true' \ 'autostart=true' \
'autorestart=true' \ 'autorestart=true' \
'stdout_logfile=/dev/stdout' \ 'stdout_logfile=/dev/stdout' \

2
aether-hub/Cargo.lock generated
View File

@@ -10,7 +10,7 @@ checksum = "320119579fcad9c21884f5c4861d16174d0e06250625266f50fe6898340abefa"
[[package]] [[package]]
name = "aether-hub" name = "aether-hub"
version = "0.1.4" version = "0.1.5"
dependencies = [ dependencies = [
"axum", "axum",
"clap", "clap",

View File

@@ -28,6 +28,11 @@ cd /path/to/Aether
- `--load`: 加载到本地 Docker单平台 - `--load`: 加载到本地 Docker单平台
- `--latest`: 额外打 `latest` tag - `--latest`: 额外打 `latest` tag
## 运行时参数
- `TUNNEL_HUB_WORKER_IDLE_TIMEOUT`worker 心跳空闲超时,默认 `60`
- `TUNNEL_HUB_OUTBOUND_QUEUE_CAPACITY`:单连接出站队列容量,默认 `128`;队列打满时会把连接视为拥塞并主动关闭,避免 Hub 内存无限增长
## 与部署脚本关系 ## 与部署脚本关系
- `./deploy.sh`: 本地构建部署(会本地构建 app/base并在构建 app 时从 GitHub Release 下载 Hub可用 `--hub-tag` 固定版本)。 - `./deploy.sh`: 本地构建部署(会本地构建 app/base并在构建 app 时从 GitHub Release 下载 Hub可用 `--hub-tag` 固定版本)。

View File

@@ -2,17 +2,89 @@
/// ///
/// Manages proxy connections (node_id -> [ProxyConn]) and worker connections (conn_id -> WorkerConn). /// Manages proxy connections (node_id -> [ProxyConn]) and worker connections (conn_id -> WorkerConn).
/// Routes frames between workers and proxies with stream_id remapping. /// Routes frames between workers and proxies with stream_id remapping.
use std::sync::atomic::{AtomicU32, AtomicU64, AtomicUsize, Ordering}; use std::sync::atomic::{AtomicBool, AtomicU32, AtomicU64, AtomicUsize, Ordering};
use std::sync::Arc; use std::sync::Arc;
use std::time::Duration;
use axum::extract::ws::Message; use axum::extract::ws::Message;
use dashmap::DashMap; use dashmap::DashMap;
use parking_lot::RwLock; use parking_lot::RwLock;
use tokio::sync::mpsc; use tokio::sync::mpsc;
use tokio::sync::mpsc::error::TrySendError;
use tokio::sync::watch;
use tracing::{debug, info, warn}; use tracing::{debug, info, warn};
use crate::protocol; use crate::protocol;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SendStatus {
Queued,
Closed,
Congested,
}
// ---------------------------------------------------------------------------
// Connection configuration (shared by proxy and worker handlers)
// ---------------------------------------------------------------------------
#[derive(Debug, Clone, Copy)]
pub struct ConnConfig {
pub ping_interval: Duration,
pub idle_timeout: Duration,
pub outbound_queue_capacity: usize,
}
// ---------------------------------------------------------------------------
// Bounded outbound channel with congestion-aware close
// ---------------------------------------------------------------------------
pub struct BoundedOutbound {
tx: mpsc::Sender<Message>,
close_tx: watch::Sender<bool>,
closing: AtomicBool,
}
impl BoundedOutbound {
pub fn new(tx: mpsc::Sender<Message>, close_tx: watch::Sender<bool>) -> Self {
Self {
tx,
close_tx,
closing: AtomicBool::new(false),
}
}
pub fn send(&self, msg: Message) -> SendStatus {
if self.is_closing() {
return SendStatus::Closed;
}
match self.tx.try_send(msg) {
Ok(()) => SendStatus::Queued,
Err(TrySendError::Closed(_)) => {
self.mark_closing();
SendStatus::Closed
}
Err(TrySendError::Full(_)) => {
self.mark_closing();
SendStatus::Congested
}
}
}
pub fn is_closing(&self) -> bool {
self.closing.load(Ordering::Acquire)
}
/// Mark as closing. Returns `true` if this call was the first to flip the flag.
pub fn mark_closing(&self) -> bool {
if self.closing.swap(true, Ordering::AcqRel) {
return false;
}
let _ = self.close_tx.send(true);
true
}
}
// --------------------------------------------------------------------------- // ---------------------------------------------------------------------------
// Proxy connection // Proxy connection
// --------------------------------------------------------------------------- // ---------------------------------------------------------------------------
@@ -21,7 +93,7 @@ pub struct ProxyConn {
pub id: u64, pub id: u64,
pub node_id: String, pub node_id: String,
pub node_name: String, pub node_name: String,
pub tx: mpsc::UnboundedSender<Message>, pub outbound: BoundedOutbound,
next_stream_id: AtomicU32, next_stream_id: AtomicU32,
pub stream_count: AtomicUsize, pub stream_count: AtomicUsize,
pub max_streams: usize, pub max_streams: usize,
@@ -32,14 +104,15 @@ impl ProxyConn {
id: u64, id: u64,
node_id: String, node_id: String,
node_name: String, node_name: String,
tx: mpsc::UnboundedSender<Message>, tx: mpsc::Sender<Message>,
close_tx: watch::Sender<bool>,
max_streams: usize, max_streams: usize,
) -> Self { ) -> Self {
Self { Self {
id, id,
node_id, node_id,
node_name, node_name,
tx, outbound: BoundedOutbound::new(tx, close_tx),
next_stream_id: AtomicU32::new(2), // even IDs, start at 2 next_stream_id: AtomicU32::new(2), // even IDs, start at 2
stream_count: AtomicUsize::new(0), stream_count: AtomicUsize::new(0),
max_streams, max_streams,
@@ -51,7 +124,7 @@ impl ProxyConn {
// Reserve one stream slot first (CAS to honor max_streams under contention). // Reserve one stream slot first (CAS to honor max_streams under contention).
let mut current = self.stream_count.load(Ordering::Relaxed); let mut current = self.stream_count.load(Ordering::Relaxed);
loop { loop {
if current >= self.max_streams { if current >= self.max_streams || !self.is_available() {
return None; return None;
} }
match self.stream_count.compare_exchange_weak( match self.stream_count.compare_exchange_weak(
@@ -99,8 +172,27 @@ impl ProxyConn {
} }
} }
pub fn send(&self, msg: Message) -> bool { pub fn is_available(&self) -> bool {
self.tx.send(msg).is_ok() !self.outbound.is_closing()
}
pub fn request_close(&self) {
self.outbound.mark_closing();
}
pub fn send(&self, msg: Message) -> SendStatus {
let was_closing = self.outbound.is_closing();
let status = self.outbound.send(msg);
if status == SendStatus::Congested && !was_closing {
warn!(
conn_id = self.id,
node_id = %self.node_id,
node_name = %self.node_name,
queued_streams = self.stream_count.load(Ordering::Relaxed),
"proxy outbound queue full, closing congested connection"
);
}
status
} }
} }
@@ -110,16 +202,35 @@ impl ProxyConn {
pub struct WorkerConn { pub struct WorkerConn {
pub id: u64, pub id: u64,
pub tx: mpsc::UnboundedSender<Message>, pub outbound: BoundedOutbound,
} }
impl WorkerConn { impl WorkerConn {
pub fn new(id: u64, tx: mpsc::UnboundedSender<Message>) -> Self { pub fn new(id: u64, tx: mpsc::Sender<Message>, close_tx: watch::Sender<bool>) -> Self {
Self { id, tx } Self {
id,
outbound: BoundedOutbound::new(tx, close_tx),
}
} }
pub fn send(&self, msg: Message) -> bool { pub fn is_available(&self) -> bool {
self.tx.send(msg).is_ok() !self.outbound.is_closing()
}
pub fn request_close(&self) {
self.outbound.mark_closing();
}
pub fn send(&self, msg: Message) -> SendStatus {
let was_closing = self.outbound.is_closing();
let status = self.outbound.send(msg);
if status == SendStatus::Congested && !was_closing {
warn!(
worker_id = self.id,
"worker outbound queue full, closing congested connection"
);
}
status
} }
} }
@@ -242,6 +353,7 @@ impl HubRouter {
let conns = map.get(node_id)?; let conns = map.get(node_id)?;
conns conns
.iter() .iter()
.filter(|c| c.is_available())
.min_by_key(|c| c.stream_count.load(Ordering::Relaxed)) .min_by_key(|c| c.stream_count.load(Ordering::Relaxed))
.cloned() .cloned()
} }
@@ -404,14 +516,17 @@ impl HubRouter {
}, },
); );
if !proxy_conn.send(Message::Binary(rebuilt_frame.into())) { match proxy_conn.send(Message::Binary(rebuilt_frame.into())) {
// Send failed, clean up mapping SendStatus::Queued => {}
self.worker_to_proxy SendStatus::Closed | SendStatus::Congested => {
.remove(&(worker_conn_id, worker_stream_id)); // Send failed, clean up mapping
self.proxy_to_worker self.worker_to_proxy
.remove(&(proxy_conn.id, proxy_stream_id)); .remove(&(worker_conn_id, worker_stream_id));
proxy_conn.release_stream(); self.proxy_to_worker
return Some("proxy connection send failed".to_string()); .remove(&(proxy_conn.id, proxy_stream_id));
proxy_conn.release_stream();
return Some("proxy connection congested".to_string());
}
} }
None None
@@ -558,7 +673,10 @@ impl HubRouter {
let workers: Vec<Arc<WorkerConn>> = self let workers: Vec<Arc<WorkerConn>> = self
.worker_conns .worker_conns
.iter() .iter()
.map(|e| e.value().clone()) .filter_map(|e| {
let worker = e.value().clone();
worker.is_available().then_some(worker)
})
.collect(); .collect();
if workers.is_empty() { if workers.is_empty() {
debug!("no workers to forward heartbeat to"); debug!("no workers to forward heartbeat to");
@@ -654,7 +772,7 @@ impl HubRouter {
let mut sent = 0usize; let mut sent = 0usize;
for entry in self.worker_conns.iter() { for entry in self.worker_conns.iter() {
if entry.value().send(msg.clone()) { if matches!(entry.value().send(msg.clone()), SendStatus::Queued) {
sent += 1; sent += 1;
} }
} }

View File

@@ -14,7 +14,7 @@ use axum::Router;
use clap::Parser; use clap::Parser;
use tracing::{info, warn}; use tracing::{info, warn};
use crate::hub::HubRouter; use crate::hub::{ConnConfig, HubRouter};
#[derive(Parser, Debug)] #[derive(Parser, Debug)]
#[command(name = "aether-hub", about = "Tunnel Hub for Aether")] #[command(name = "aether-hub", about = "Tunnel Hub for Aether")]
@@ -28,7 +28,7 @@ struct Args {
proxy_idle_timeout: u64, proxy_idle_timeout: u64,
/// Worker-side idle timeout in seconds (0 to disable) /// Worker-side idle timeout in seconds (0 to disable)
#[arg(long, default_value_t = 0, env = "TUNNEL_HUB_WORKER_IDLE_TIMEOUT")] #[arg(long, default_value_t = 60, env = "TUNNEL_HUB_WORKER_IDLE_TIMEOUT")]
worker_idle_timeout: u64, worker_idle_timeout: u64,
/// Ping interval in seconds (for both sides) /// Ping interval in seconds (for both sides)
@@ -38,14 +38,21 @@ struct Args {
/// Max concurrent streams per proxy connection /// Max concurrent streams per proxy connection
#[arg(long, default_value_t = 2048, env = "TUNNEL_HUB_MAX_STREAMS")] #[arg(long, default_value_t = 2048, env = "TUNNEL_HUB_MAX_STREAMS")]
max_streams: usize, max_streams: usize,
/// Per-connection outbound queue capacity before treating the socket as congested
#[arg(
long,
default_value_t = 128,
env = "TUNNEL_HUB_OUTBOUND_QUEUE_CAPACITY"
)]
outbound_queue_capacity: usize,
} }
#[derive(Clone)] #[derive(Clone)]
struct AppState { struct AppState {
hub: Arc<HubRouter>, hub: Arc<HubRouter>,
proxy_idle_timeout: Duration, proxy_conn_cfg: ConnConfig,
worker_idle_timeout: Duration, worker_conn_cfg: ConnConfig,
ping_interval: Duration,
max_streams: usize, max_streams: usize,
} }
@@ -62,11 +69,20 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
let args = Args::parse(); let args = Args::parse();
let hub = HubRouter::new(); let hub = HubRouter::new();
let outbound_queue_capacity = args.outbound_queue_capacity.clamp(8, 4096);
let ping_interval = Duration::from_secs(args.ping_interval);
let state = AppState { let state = AppState {
hub, hub,
proxy_idle_timeout: Duration::from_secs(args.proxy_idle_timeout), proxy_conn_cfg: ConnConfig {
worker_idle_timeout: Duration::from_secs(args.worker_idle_timeout), ping_interval,
ping_interval: Duration::from_secs(args.ping_interval), idle_timeout: Duration::from_secs(args.proxy_idle_timeout),
outbound_queue_capacity,
},
worker_conn_cfg: ConnConfig {
ping_interval,
idle_timeout: Duration::from_secs(args.worker_idle_timeout),
outbound_queue_capacity,
},
max_streams: args.max_streams, max_streams: args.max_streams,
}; };
@@ -139,8 +155,7 @@ async fn ws_proxy(
node_id, node_id,
node_name, node_name,
max_streams, max_streams,
state.ping_interval, state.proxy_conn_cfg,
state.proxy_idle_timeout,
) )
}) })
.into_response() .into_response()
@@ -149,11 +164,6 @@ async fn ws_proxy(
async fn ws_worker(ws: WebSocketUpgrade, State(state): State<AppState>) -> impl IntoResponse { async fn ws_worker(ws: WebSocketUpgrade, State(state): State<AppState>) -> impl IntoResponse {
ws.max_frame_size(64 * 1024 * 1024) ws.max_frame_size(64 * 1024 * 1024)
.on_upgrade(move |socket| { .on_upgrade(move |socket| {
worker_conn::handle_worker_connection( worker_conn::handle_worker_connection(socket, state.hub, state.worker_conn_cfg)
socket,
state.hub,
state.ping_interval,
state.worker_idle_timeout,
)
}) })
} }

View File

@@ -7,10 +7,10 @@ use std::time::Duration;
use axum::extract::ws::{Message, WebSocket}; use axum::extract::ws::{Message, WebSocket};
use futures_util::{SinkExt, StreamExt}; use futures_util::{SinkExt, StreamExt};
use tokio::sync::mpsc; use tokio::sync::{mpsc, watch};
use tracing::{debug, info, warn}; use tracing::{debug, info, warn};
use crate::hub::{HubRouter, ProxyConn}; use crate::hub::{ConnConfig, HubRouter, ProxyConn, SendStatus};
use crate::protocol; use crate::protocol;
/// Maximum single frame size: 64 MB /// Maximum single frame size: 64 MB
@@ -22,69 +22,71 @@ pub async fn handle_proxy_connection(
node_id: String, node_id: String,
node_name: String, node_name: String,
max_streams: usize, max_streams: usize,
ping_interval: Duration, cfg: ConnConfig,
idle_timeout: Duration,
) { ) {
let conn_id = hub.alloc_conn_id(); let conn_id = hub.alloc_conn_id();
let (mut ws_tx, ws_rx) = ws.split(); let (mut ws_tx, ws_rx) = ws.split();
// Create channel for outbound messages let (tx, mut rx) = mpsc::channel::<Message>(cfg.outbound_queue_capacity);
let (tx, mut rx) = mpsc::unbounded_channel::<Message>(); let (close_tx, mut close_rx) = watch::channel(false);
let conn = Arc::new(ProxyConn::new( let conn = Arc::new(ProxyConn::new(
conn_id, conn_id,
node_id.clone(), node_id.clone(),
node_name.clone(), node_name.clone(),
tx, tx,
close_tx,
max_streams, max_streams,
)); ));
hub.register_proxy(conn.clone()); hub.register_proxy(conn.clone());
// Spawn writer task: drains channel -> WebSocket
let writer = tokio::spawn(async move { let writer = tokio::spawn(async move {
while let Some(msg) = rx.recv().await { loop {
if ws_tx.send(msg).await.is_err() { tokio::select! {
break; msg = rx.recv() => match msg {
Some(msg) => {
if ws_tx.send(msg).await.is_err() {
break;
}
}
None => break,
},
changed = close_rx.changed() => {
if changed.is_err() || *close_rx.borrow() {
break;
}
}
} }
} }
let _ = ws_tx.close().await; let _ = ws_tx.close().await;
}); });
// Spawn ping task let ping_conn = conn.clone();
let ping_tx = conn.tx.clone(); let ping_interval = cfg.ping_interval;
let ping_task = tokio::spawn(async move { let ping_task = tokio::spawn(async move {
loop { loop {
tokio::time::sleep(ping_interval).await; tokio::time::sleep(ping_interval).await;
let ping = protocol::encode_ping(); let ping = protocol::encode_ping();
if ping_tx.send(Message::Binary(ping.into())).is_err() { if !matches!(
ping_conn.send(Message::Binary(ping.into())),
SendStatus::Queued
) {
break; break;
} }
} }
}); });
// Spawn reader task
let reader_hub = hub.clone(); let reader_hub = hub.clone();
let reader_node_id = node_id.clone(); let reader_conn = conn.clone();
let reader_tx = conn.tx.clone();
let reader = tokio::spawn(async move { let reader = tokio::spawn(async move {
run_proxy_reader( run_proxy_reader(ws_rx, reader_hub, reader_conn, cfg.idle_timeout).await;
ws_rx,
reader_hub,
conn_id,
reader_node_id,
reader_tx,
idle_timeout,
)
.await;
}); });
// Wait for reader to end, then cleanup.
let _ = reader.await; let _ = reader.await;
ping_task.abort(); ping_task.abort();
conn.request_close();
hub.unregister_proxy(conn_id, &node_id); hub.unregister_proxy(conn_id, &node_id);
// conn still holds an Arc<ProxyConn> with a channel sender clone.
// Drop it so the writer can drain and exit.
drop(conn); drop(conn);
tokio::time::sleep(Duration::from_millis(100)).await; tokio::time::sleep(Duration::from_millis(100)).await;
writer.abort(); writer.abort();
@@ -94,9 +96,7 @@ pub async fn handle_proxy_connection(
async fn run_proxy_reader( async fn run_proxy_reader(
mut ws_rx: futures_util::stream::SplitStream<WebSocket>, mut ws_rx: futures_util::stream::SplitStream<WebSocket>,
hub: Arc<HubRouter>, hub: Arc<HubRouter>,
conn_id: u64, conn: Arc<ProxyConn>,
node_id: String,
tx: mpsc::UnboundedSender<Message>,
idle_timeout: Duration, idle_timeout: Duration,
) { ) {
let idle_enabled = !idle_timeout.is_zero(); let idle_enabled = !idle_timeout.is_zero();
@@ -106,8 +106,9 @@ async fn run_proxy_reader(
tokio::select! { tokio::select! {
msg = ws_rx.next() => msg, msg = ws_rx.next() => msg,
_ = tokio::time::sleep(idle_timeout) => { _ = tokio::time::sleep(idle_timeout) => {
warn!(conn_id = conn_id, node_id = %node_id, "proxy idle timeout"); warn!(conn_id = conn.id, node_id = %conn.node_id, "proxy idle timeout");
let _ = tx.send(Message::Binary(protocol::encode_goaway().into())); let _ = conn.send(Message::Binary(protocol::encode_goaway().into()));
conn.request_close();
break; break;
} }
} }
@@ -121,12 +122,13 @@ async fn run_proxy_reader(
if data.len() > MAX_FRAME_SIZE { if data.len() > MAX_FRAME_SIZE {
oversized_count += 1; oversized_count += 1;
warn!( warn!(
conn_id = conn_id, conn_id = conn.id,
size = data.len(), size = data.len(),
"oversized frame from proxy" "oversized frame from proxy"
); );
if oversized_count >= 5 { if oversized_count >= 5 {
warn!(conn_id = conn_id, "too many oversized frames, closing"); warn!(conn_id = conn.id, "too many oversized frames, closing");
conn.request_close();
break; break;
} }
continue; continue;
@@ -134,21 +136,21 @@ async fn run_proxy_reader(
oversized_count = 0; oversized_count = 0;
if data.len() < protocol::HEADER_SIZE { if data.len() < protocol::HEADER_SIZE {
debug!(conn_id = conn_id, "frame too small, skipping"); debug!(conn_id = conn.id, "frame too small, skipping");
continue; continue;
} }
hub.handle_proxy_frame(conn_id, &mut data); hub.handle_proxy_frame(conn.id, &mut data);
} }
Some(Ok(Message::Close(_))) | None => { Some(Ok(Message::Close(_))) | None => {
info!(conn_id = conn_id, node_id = %node_id, "proxy WebSocket closed"); info!(conn_id = conn.id, node_id = %conn.node_id, "proxy WebSocket closed");
break; break;
} }
Some(Err(e)) => { Some(Err(e)) => {
warn!(conn_id = conn_id, error = %e, "proxy WebSocket error"); warn!(conn_id = conn.id, error = %e, "proxy WebSocket error");
break; break;
} }
_ => {} // Ignore text/ping/pong at WS level _ => {}
} }
} }
} }

View File

@@ -2,76 +2,92 @@
/// ///
/// Handles the lifecycle of a single Gunicorn worker connection: /// Handles the lifecycle of a single Gunicorn worker connection:
/// accept -> read loop (route frames via Hub) -> cleanup /// accept -> read loop (route frames via Hub) -> cleanup
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc; use std::sync::Arc;
use std::time::Duration; use std::time::{Duration, Instant};
use axum::extract::ws::{Message, WebSocket}; use axum::extract::ws::{Message, WebSocket};
use futures_util::{SinkExt, StreamExt}; use futures_util::{SinkExt, StreamExt};
use tokio::sync::mpsc; use tokio::sync::{mpsc, watch};
use tracing::{debug, info, warn}; use tracing::{debug, info, warn};
use crate::hub::{HubRouter, WorkerConn}; use crate::hub::{ConnConfig, HubRouter, SendStatus, WorkerConn};
use crate::protocol; use crate::protocol;
pub async fn handle_worker_connection( pub async fn handle_worker_connection(ws: WebSocket, hub: Arc<HubRouter>, cfg: ConnConfig) {
ws: WebSocket,
hub: Arc<HubRouter>,
ping_interval: Duration,
idle_timeout: Duration,
) {
let conn_id = hub.alloc_conn_id(); let conn_id = hub.alloc_conn_id();
let (mut ws_tx, ws_rx) = ws.split(); let (mut ws_tx, ws_rx) = ws.split();
// Create channel for outbound messages let (tx, mut rx) = mpsc::channel::<Message>(cfg.outbound_queue_capacity);
let (tx, mut rx) = mpsc::unbounded_channel::<Message>(); let (close_tx, mut close_rx) = watch::channel(false);
let conn = Arc::new(WorkerConn::new(conn_id, tx)); let conn = Arc::new(WorkerConn::new(conn_id, tx, close_tx));
hub.register_worker(conn.clone()); hub.register_worker(conn.clone());
// Spawn writer task
let writer = tokio::spawn(async move { let writer = tokio::spawn(async move {
while let Some(msg) = rx.recv().await { loop {
if ws_tx.send(msg).await.is_err() { tokio::select! {
break; msg = rx.recv() => match msg {
Some(msg) => {
if ws_tx.send(msg).await.is_err() {
break;
}
}
None => break,
},
changed = close_rx.changed() => {
if changed.is_err() || *close_rx.borrow() {
break;
}
}
} }
} }
let _ = ws_tx.close().await; let _ = ws_tx.close().await;
}); });
// Spawn ping task let liveness_clock = Instant::now();
let ping_tx = conn.tx.clone(); let last_seen_ms = Arc::new(AtomicU64::new(0));
let ping_task = tokio::spawn(async move {
loop {
tokio::time::sleep(ping_interval).await;
let ping = protocol::encode_ping();
if ping_tx.send(Message::Binary(ping.into())).is_err() {
break;
}
}
});
// Spawn reader task
let reader_hub = hub.clone(); let reader_hub = hub.clone();
let reader_tx = conn.tx.clone(); let reader_conn = conn.clone();
let reader = tokio::spawn(async move { let reader_last_seen_ms = last_seen_ms.clone();
let liveness_conn = conn.clone();
let mut reader = tokio::spawn(async move {
run_worker_reader( run_worker_reader(
ws_rx, ws_rx,
reader_hub, reader_hub,
conn_id, conn_id,
conn.clone(), reader_conn,
reader_tx, reader_last_seen_ms,
idle_timeout, liveness_clock,
) )
.await; .await;
}); });
// Wait for reader to end, then cleanup. let liveness_last_seen_ms = last_seen_ms.clone();
let mut liveness = tokio::spawn(async move {
run_worker_liveness(
conn_id,
liveness_conn,
cfg.ping_interval,
cfg.idle_timeout,
liveness_last_seen_ms,
liveness_clock,
)
.await;
});
tokio::select! {
_ = &mut reader => {}
_ = &mut liveness => {}
}
conn.request_close();
reader.abort();
liveness.abort();
let _ = reader.await; let _ = reader.await;
ping_task.abort(); let _ = liveness.await;
// Unregister first so all channel senders are dropped (reader_tx dropped
// when reader completes, ping_tx dropped by abort, conn.tx dropped when
// the last Arc<WorkerConn> is removed from hub). This lets the writer
// drain buffered messages (e.g. GOAWAY) before we force-abort it.
hub.unregister_worker(conn_id); hub.unregister_worker(conn_id);
tokio::time::sleep(Duration::from_millis(100)).await; tokio::time::sleep(Duration::from_millis(100)).await;
writer.abort(); writer.abort();
@@ -83,26 +99,14 @@ async fn run_worker_reader(
hub: Arc<HubRouter>, hub: Arc<HubRouter>,
conn_id: u64, conn_id: u64,
conn: Arc<WorkerConn>, conn: Arc<WorkerConn>,
tx: mpsc::UnboundedSender<Message>, last_seen_ms: Arc<AtomicU64>,
idle_timeout: Duration, liveness_clock: Instant,
) { ) {
let idle_enabled = !idle_timeout.is_zero();
loop { loop {
let msg = if idle_enabled { match ws_rx.next().await {
tokio::select! {
msg = ws_rx.next() => msg,
_ = tokio::time::sleep(idle_timeout) => {
warn!(worker_id = conn_id, "worker idle timeout");
let _ = tx.send(Message::Binary(protocol::encode_goaway().into()));
break;
}
}
} else {
ws_rx.next().await
};
match msg {
Some(Ok(Message::Binary(data))) => { Some(Ok(Message::Binary(data))) => {
last_seen_ms.store(elapsed_millis(liveness_clock), Ordering::Relaxed);
let mut data = data.to_vec(); let mut data = data.to_vec();
if data.len() < protocol::HEADER_SIZE { if data.len() < protocol::HEADER_SIZE {
debug!(worker_id = conn_id, "frame too small, skipping"); debug!(worker_id = conn_id, "frame too small, skipping");
@@ -114,15 +118,12 @@ async fn run_worker_reader(
None => continue, None => continue,
}; };
// HEARTBEAT_ACK from worker -> route back to proxy
if header.msg_type == protocol::HEARTBEAT_ACK { if header.msg_type == protocol::HEARTBEAT_ACK {
hub.handle_worker_heartbeat_ack(&mut data); hub.handle_worker_heartbeat_ack(&mut data);
continue; continue;
} }
// Regular frames: route via hub
if let Some(err_msg) = hub.handle_worker_frame(conn_id, &mut data) { if let Some(err_msg) = hub.handle_worker_frame(conn_id, &mut data) {
// Send STREAM_ERROR back to worker
let err_frame = protocol::encode_stream_error(header.stream_id, &err_msg); let err_frame = protocol::encode_stream_error(header.stream_id, &err_msg);
let _ = conn.send(Message::Binary(err_frame.into())); let _ = conn.send(Message::Binary(err_frame.into()));
} }
@@ -135,7 +136,55 @@ async fn run_worker_reader(
warn!(worker_id = conn_id, error = %e, "worker WebSocket error"); warn!(worker_id = conn_id, error = %e, "worker WebSocket error");
break; break;
} }
_ => {} // Ignore text/ping/pong at WS level _ => {}
} }
} }
} }
async fn run_worker_liveness(
conn_id: u64,
conn: Arc<WorkerConn>,
ping_interval: Duration,
idle_timeout: Duration,
last_seen_ms: Arc<AtomicU64>,
liveness_clock: Instant,
) {
let ping_interval_ms = ping_interval.as_millis().max(1) as u64;
let idle_timeout_ms = idle_timeout.as_millis() as u64;
loop {
tokio::time::sleep(ping_interval).await;
let ping = protocol::encode_ping();
if !matches!(conn.send(Message::Binary(ping.into())), SendStatus::Queued) {
break;
}
if idle_timeout.is_zero() {
continue;
}
let now_ms = elapsed_millis(liveness_clock);
let last_seen = last_seen_ms.load(Ordering::Relaxed);
let silent_for_ms = now_ms.saturating_sub(last_seen);
if silent_for_ms < idle_timeout_ms {
continue;
}
let missed_heartbeats = (silent_for_ms / ping_interval_ms).max(1);
warn!(
worker_id = conn_id,
idle_timeout_secs = idle_timeout.as_secs(),
silent_for_ms = silent_for_ms,
missed_heartbeats = missed_heartbeats,
"worker heartbeat timeout"
);
let _ = conn.send(Message::Binary(protocol::encode_goaway().into()));
conn.request_close();
break;
}
}
fn elapsed_millis(started_at: Instant) -> u64 {
started_at.elapsed().as_millis().min(u64::MAX as u128) as u64
}

View File

@@ -1,8 +1,12 @@
# Gunicorn configuration file # Gunicorn configuration file
from __future__ import annotations from __future__ import annotations
import faulthandler
import gc import gc
import os import os
import signal
import sys
import traceback
from typing import Any from typing import Any
# worker 心跳超时(秒):异步 worker 在此时间内必须向 arbiter 发送心跳 # worker 心跳超时(秒):异步 worker 在此时间内必须向 arbiter 发送心跳
@@ -25,6 +29,52 @@ def _log_current_rss(log: Any, message: str) -> None:
pass pass
def _enable_fault_handler(log: Any) -> None:
try:
faulthandler.enable(file=sys.stderr, all_threads=True)
except Exception as exc:
log.warning(f"Failed to enable faulthandler: {exc}")
return
sigusr2 = getattr(signal, "SIGUSR2", None)
if sigusr2 is None:
return
try:
faulthandler.unregister(sigusr2)
except Exception:
pass
try:
faulthandler.register(sigusr2, file=sys.stderr, all_threads=True, chain=False)
log.info("Registered faulthandler stack dump on SIGUSR2")
except Exception as exc:
log.warning(f"Failed to register faulthandler SIGUSR2 hook: {exc}")
def _dump_all_thread_traces(log: Any, reason: str) -> None:
pid = os.getpid()
log.critical(f"===== Python stack dump begin: reason={reason}, pid={pid} =====")
_log_current_rss(log, f"Worker {pid} RSS before traceback dump")
try:
faulthandler.dump_traceback(file=sys.stderr, all_threads=True)
except Exception as exc:
log.warning(f"faulthandler.dump_traceback failed: {exc}")
try:
current_frames = sys._current_frames()
for thread_id, frame in current_frames.items():
stack = "".join(traceback.format_stack(frame))
log.critical(
f"--- thread_id={thread_id} stack begin ---\n{stack}--- thread_id={thread_id} stack end ---"
)
except Exception as exc:
log.warning(f"Failed to dump Python frames via sys._current_frames(): {exc}")
log.critical(f"===== Python stack dump end: reason={reason}, pid={pid} =====")
def when_ready(server: Any) -> None: def when_ready(server: Any) -> None:
""" """
Called just after the server is started. Called just after the server is started.
@@ -43,3 +93,8 @@ def post_fork(server: Any, worker: Any) -> None:
def post_worker_init(worker: Any) -> None: def post_worker_init(worker: Any) -> None:
_log_current_rss(worker.log, f"Worker {worker.pid} RSS after app init") _log_current_rss(worker.log, f"Worker {worker.pid} RSS after app init")
_enable_fault_handler(worker.log)
def worker_abort(worker: Any) -> None:
_dump_all_thread_traces(worker.log, "gunicorn worker timeout / SIGABRT")

View File

@@ -4,6 +4,7 @@ Endpoint 健康监控 API
from __future__ import annotations from __future__ import annotations
import asyncio
from collections import defaultdict from collections import defaultdict
from dataclasses import dataclass from dataclasses import dataclass
from datetime import datetime, timedelta, timezone from datetime import datetime, timedelta, timezone
@@ -33,6 +34,60 @@ from src.services.health.monitor import HealthMonitor, health_monitor
router = APIRouter(tags=["Endpoint Health"]) router = APIRouter(tags=["Endpoint Health"])
def _recover_key_health_sync(db: Session, key_id: str, api_format: str | None) -> dict[str, Any]:
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
if not key:
raise NotFoundException(f"Key {key_id} 不存在")
success = health_monitor.reset_health(db, key_id=key_id, api_format=api_format)
if not success:
raise Exception("重置健康度失败")
if not key.is_active:
key.is_active = True # type: ignore[assignment]
db.commit()
return {
"is_active": bool(key.is_active),
"api_format": api_format,
}
def _recover_all_keys_health_sync(db: Session) -> list[dict[str, Any]]:
candidates = (
db.query(ProviderAPIKey)
.filter(
ProviderAPIKey.circuit_breaker_by_format.isnot(None),
ProviderAPIKey.circuit_breaker_by_format != "{}",
)
.all()
)
circuit_open_keys = [
key
for key in candidates
if any(cb.get("open") for cb in (key.circuit_breaker_by_format or {}).values())
]
recovered_keys: list[dict[str, Any]] = []
for key in circuit_open_keys:
key.health_by_format = {} # type: ignore[assignment]
key.circuit_breaker_by_format = {} # type: ignore[assignment]
recovered_keys.append(
{
"key_id": key.id,
"key_name": key.name,
"provider_id": key.provider_id,
"api_formats": key.api_formats,
}
)
if recovered_keys:
db.commit()
return recovered_keys
def _format_str(api_format_enum: Any) -> str: def _format_str(api_format_enum: Any) -> str:
"""将 DB 查询返回的 api_format可能是 enum 或 str统一转为 str。""" """将 DB 查询返回的 api_format可能是 enum 或 str统一转为 str。"""
return api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum) return api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
@@ -491,20 +546,7 @@ class AdminRecoverKeyHealthAdapter(AdminApiAdapter):
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override] async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db db = context.db
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == self.key_id).first() await asyncio.to_thread(_recover_key_health_sync, db, self.key_id, self.api_format)
if not key:
raise NotFoundException(f"Key {self.key_id} 不存在")
# 使用 health_monitor.reset_health 重置健康度
success = health_monitor.reset_health(db, key_id=self.key_id, api_format=self.api_format)
if not success:
raise Exception("重置健康度失败")
# 如果 Key 被禁用,重新启用
if not key.is_active:
key.is_active = True # type: ignore[assignment]
db.commit()
if self.api_format: if self.api_format:
logger.info(f"管理员恢复Key健康状态: {self.key_id}/{self.api_format}") logger.info(f"管理员恢复Key健康状态: {self.key_id}/{self.api_format}")
@@ -534,47 +576,15 @@ class AdminRecoverAllKeysHealthAdapter(AdminApiAdapter):
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override] async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db db = context.db
recovered_keys = await asyncio.to_thread(_recover_all_keys_health_sync, db)
# 粗过滤:仅加载 circuit_breaker_by_format 非空的 Key避免全表扫描 if not recovered_keys:
candidates = (
db.query(ProviderAPIKey)
.filter(
ProviderAPIKey.circuit_breaker_by_format.isnot(None),
ProviderAPIKey.circuit_breaker_by_format != "{}",
)
.all()
)
# 精确筛选有任何格式熔断的 Key
circuit_open_keys = [
key
for key in candidates
if any(cb.get("open") for cb in (key.circuit_breaker_by_format or {}).values())
]
if not circuit_open_keys:
return { return {
"message": "没有需要恢复的 Key", "message": "没有需要恢复的 Key",
"recovered_count": 0, "recovered_count": 0,
"recovered_keys": [], "recovered_keys": [],
} }
recovered_keys = []
for key in circuit_open_keys:
# 重置所有格式的健康度
key.health_by_format = {} # type: ignore[assignment]
key.circuit_breaker_by_format = {} # type: ignore[assignment]
recovered_keys.append(
{
"key_id": key.id,
"key_name": key.name,
"provider_id": key.provider_id,
"api_formats": key.api_formats,
}
)
db.commit()
# 重置健康监控器的熔断计数 # 重置健康监控器的熔断计数
HealthMonitor.reset_open_circuit_count() HealthMonitor.reset_open_circuit_count()

View File

@@ -161,29 +161,43 @@ async def _calculate_and_record_usage(
from src.services.request.candidate import RequestCandidateService from src.services.request.candidate import RequestCandidateService
from src.services.usage.service import UsageService from src.services.usage.service import UsageService
# 获取Provider API Key对象不是用户API Key def _load_usage_context() -> tuple[Any, Any, Any]:
provider_api_key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == api_key_id).first() # 获取Provider API Key对象不是用户API Key
provider_api_key_local = (
db.query(ProviderAPIKey).filter(ProviderAPIKey.id == api_key_id).first()
)
if not provider_api_key_local:
return None, None, None
provider_endpoint_local = None
if api_format and provider_api_key_local.provider_id:
from src.models.database import Provider
provider = (
db.query(Provider).filter(Provider.id == provider_api_key_local.provider_id).first()
)
if provider:
for ep in provider.endpoints:
if ep.api_format == api_format:
provider_endpoint_local = ep
break
user_api_key_local = None
if user:
try:
user_api_key_local = db.query(ApiKey).filter(ApiKey.user_id == user.id).first()
except Exception:
user_api_key_local = None
return provider_api_key_local, provider_endpoint_local, user_api_key_local
provider_api_key, provider_endpoint, user_api_key = await asyncio.to_thread(_load_usage_context)
if not provider_api_key: if not provider_api_key:
logger.warning(f"Provider API Key not found for usage calculation: {api_key_id}") logger.warning(f"Provider API Key not found for usage calculation: {api_key_id}")
return {"error": "Provider API Key not found"} return {"error": "Provider API Key not found"}
# 获取Provider Endpoint信息通过 api_format 查找)
provider_endpoint = None
if api_format and provider_api_key.provider_id:
from src.models.database import Provider
provider = db.query(Provider).filter(Provider.id == provider_api_key.provider_id).first()
if provider:
for ep in provider.endpoints:
if ep.api_format == api_format:
provider_endpoint = ep
break
# 获取用户的API Key用于记录关联即使实际使用的是Provider API Key
user_api_key = None
if user: if user:
try: try:
user_api_key = db.query(ApiKey).filter(ApiKey.user_id == user.id).first()
logger.info( logger.info(
f"[endpoint_check] User API Key found: {user_api_key.id if user_api_key else None}" f"[endpoint_check] User API Key found: {user_api_key.id if user_api_key else None}"
) )
@@ -317,45 +331,46 @@ async def _calculate_and_record_usage(
# 创建RequestCandidate记录用于监控追踪API # 创建RequestCandidate记录用于监控追踪API
try: try:
# 首先创建候选记录
candidate = RequestCandidateService.create_candidate(
db=db,
request_id=f"test_{request_id}",
candidate_index=0, # 测试请求只有一个候选
user_id=user.id if user else None,
api_key_id=user_api_key.id if user_api_key else None,
provider_id=provider_id,
endpoint_id=provider_endpoint.id if provider_endpoint else None,
key_id=api_key_id,
status="available",
extra_data={"model_name": model_name, "request_type": "endpoint_test"},
)
# 立即标记为开始执行 def _record_candidate_sync() -> str:
RequestCandidateService.mark_candidate_started(db, candidate.id) candidate = RequestCandidateService.create_candidate(
# 根据结果标记为成功或失败
if status_code == 200:
RequestCandidateService.mark_candidate_success(
db=db, db=db,
candidate_id=candidate.id, request_id=f"test_{request_id}",
status_code=status_code, candidate_index=0, # 测试请求只有一个候选
latency_ms=response_time_ms, user_id=user.id if user else None,
extra_data={"model_name": model_name, "api_format": api_format}, api_key_id=user_api_key.id if user_api_key else None,
) provider_id=provider_id,
else: endpoint_id=provider_endpoint.id if provider_endpoint else None,
RequestCandidateService.mark_candidate_failed( key_id=api_key_id,
db=db, status="available",
candidate_id=candidate.id, extra_data={"model_name": model_name, "request_type": "endpoint_test"},
error_type="http_error" if status_code > 0 else "network_error",
error_message=error_message or "Unknown error",
status_code=status_code,
latency_ms=response_time_ms,
extra_data={"model_name": model_name, "api_format": api_format},
) )
RequestCandidateService.mark_candidate_started(db, candidate.id)
if status_code == 200:
RequestCandidateService.mark_candidate_success(
db=db,
candidate_id=candidate.id,
status_code=status_code,
latency_ms=response_time_ms,
extra_data={"model_name": model_name, "api_format": api_format},
)
else:
RequestCandidateService.mark_candidate_failed(
db=db,
candidate_id=candidate.id,
error_type="http_error" if status_code > 0 else "network_error",
error_message=error_message or "Unknown error",
status_code=status_code,
latency_ms=response_time_ms,
extra_data={"model_name": model_name, "api_format": api_format},
)
return str(candidate.id)
candidate_id = await asyncio.to_thread(_record_candidate_sync)
logger.info( logger.info(
f"[endpoint_check] RequestCandidate created | request_id=test_{request_id}, candidate_id={candidate.id}" f"[endpoint_check] RequestCandidate created | request_id=test_{request_id}, candidate_id={candidate_id}"
) )
except Exception as e: except Exception as e:
logger.warning(f"[endpoint_check] Failed to create RequestCandidate: {e}") logger.warning(f"[endpoint_check] Failed to create RequestCandidate: {e}")

View File

@@ -30,6 +30,12 @@ if TYPE_CHECKING:
_TUNNEL_COMPRESS_MIN_SIZE = 512 _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) _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 _HEARTBEAT_DEDUP_TTL_SECONDS = 600
_LOOP_WATCHDOG_INTERVAL_SECONDS = 1.0
_LOOP_LAG_WARNING_SECONDS = 1.0
_LOOP_LAG_DEGRADE_SECONDS = 3.0
_LOOP_LAG_DEGRADE_MIN_COOLDOWN_SECONDS = 10.0
_LOOP_LAG_DEGRADE_MAX_COOLDOWN_SECONDS = 30.0
_LOOP_LAG_WARNING_LOG_INTERVAL_SECONDS = 10.0
_HOP_BY_HOP_HEADERS = frozenset( _HOP_BY_HOP_HEADERS = frozenset(
{ {
@@ -65,9 +71,12 @@ class HubConnectionManager:
self._reader_task: asyncio.Task[None] | None = None self._reader_task: asyncio.Task[None] | None = None
self._ping_task: asyncio.Task[None] | None = None self._ping_task: asyncio.Task[None] | None = None
self._reconnect_task: asyncio.Task[None] | None = None self._reconnect_task: asyncio.Task[None] | None = None
self._watchdog_task: asyncio.Task[None] | None = None
self._background_tasks: set[asyncio.Task[None]] = set() self._background_tasks: set[asyncio.Task[None]] = set()
self._closing = False self._closing = False
self._degraded_until: float = 0.0
self._last_loop_lag_warning_ts: float = 0.0
self._disconnect_count = 0 # 连续断开计数,用于抑制重复日志 self._disconnect_count = 0 # 连续断开计数,用于抑制重复日志
# 连续快速断开退避:防止 Hub 端持续发送 GOAWAY 时产生重连风暴 # 连续快速断开退避:防止 Hub 端持续发送 GOAWAY 时产生重连风暴
@@ -85,6 +94,7 @@ class HubConnectionManager:
task.add_done_callback(self._background_tasks.discard) task.add_done_callback(self._background_tasks.discard)
async def ensure_connected(self) -> None: async def ensure_connected(self) -> None:
self._ensure_watchdog_running()
if self._closing: if self._closing:
raise TunnelStreamError("hub connection manager is shutting down") raise TunnelStreamError("hub connection manager is shutting down")
if self.is_connected: if self.is_connected:
@@ -142,6 +152,55 @@ class HubConnectionManager:
) )
self._disconnect_count = 0 self._disconnect_count = 0
def _ensure_watchdog_running(self) -> None:
if self._closing:
return
if self._watchdog_task is not None and not self._watchdog_task.done():
return
self._watchdog_task = asyncio.create_task(self._loop_watchdog())
def _record_loop_lag(self, lag_seconds: float) -> None:
if lag_seconds < _LOOP_LAG_WARNING_SECONDS:
return
now = _time.monotonic()
if lag_seconds >= _LOOP_LAG_DEGRADE_SECONDS:
cooldown = min(
_LOOP_LAG_DEGRADE_MAX_COOLDOWN_SECONDS,
max(_LOOP_LAG_DEGRADE_MIN_COOLDOWN_SECONDS, lag_seconds * 3.0),
)
degraded_until = now + cooldown
self._degraded_until = max(self._degraded_until, degraded_until)
logger.warning(
"Hub worker event loop lag detected: lag={:.2f}s, pausing new streams for {:.1f}s",
lag_seconds,
cooldown,
)
return
if now - self._last_loop_lag_warning_ts >= _LOOP_LAG_WARNING_LOG_INTERVAL_SECONDS:
self._last_loop_lag_warning_ts = now
logger.warning("Hub worker event loop lag observed: lag={:.2f}s", lag_seconds)
def _raise_if_degraded(self) -> None:
remaining = self._degraded_until - _time.monotonic()
if remaining <= 0:
return
raise TunnelStreamError(f"hub worker event loop degraded, retry in {remaining:.1f}s")
async def _loop_watchdog(self) -> None:
interval = _LOOP_WATCHDOG_INTERVAL_SECONDS
expected_at = _time.monotonic() + interval
try:
while not self._closing:
await asyncio.sleep(interval)
now = _time.monotonic()
lag_seconds = max(0.0, now - expected_at)
expected_at = now + interval
self._record_loop_lag(lag_seconds)
except asyncio.CancelledError:
return
def _start_reconnect_loop(self) -> None: def _start_reconnect_loop(self) -> None:
if self._closing: if self._closing:
return return
@@ -530,6 +589,7 @@ class HubConnectionManager:
timeout: float = 60.0, timeout: float = 60.0,
) -> _StreamState: ) -> _StreamState:
await self.ensure_connected() await self.ensure_connected()
self._raise_if_degraded()
if len(self._pending_streams) >= self._config.max_streams: if len(self._pending_streams) >= self._config.max_streams:
raise TunnelStreamError( raise TunnelStreamError(
@@ -587,6 +647,8 @@ class HubConnectionManager:
self._reader_task.cancel() self._reader_task.cancel()
if self._ping_task is not None: if self._ping_task is not None:
self._ping_task.cancel() self._ping_task.cancel()
if self._watchdog_task is not None:
self._watchdog_task.cancel()
tasks = list(self._background_tasks) tasks = list(self._background_tasks)
for task in tasks: for task in tasks:

View File

@@ -102,6 +102,7 @@ def _increment_provider_api_key_totals(
return return
from sqlalchemy import func as sql_func from sqlalchemy import func as sql_func
token_increment = int(total_tokens or 0) token_increment = int(total_tokens or 0)
cost_increment = to_money_decimal(total_cost) cost_increment = to_money_decimal(total_cost)
@@ -329,57 +330,60 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
usage_params, total_cost = await cls._prepare_usage_record(params) usage_params, total_cost = await cls._prepare_usage_record(params)
total_cost = to_money_decimal(total_cost) total_cost = to_money_decimal(total_cost)
# 创建 Usage 记录 def _sync_record() -> Usage:
usage = Usage(**usage_params) # 创建 Usage 记录与相关统计;同步 SQLAlchemy 操作统一移到线程池,避免阻塞事件循环。
db.add(usage) usage = Usage(**usage_params)
db.add(usage)
# 更新 GlobalModel 使用计数(原子操作) # 更新 GlobalModel 使用计数(原子操作)
from sqlalchemy import update from sqlalchemy import update
from src.models.database import GlobalModel from src.models.database import GlobalModel
db.execute(
update(GlobalModel)
.where(GlobalModel.name == model)
.values(usage_count=GlobalModel.usage_count + 1)
)
# 更新用户-模型调用次数计数器
cls._increment_user_model_usage(db, user, model)
# 更新 Provider 月度使用量(原子操作)
if provider_id:
actual_total_cost = Decimal(str(usage_params["actual_total_cost_usd"]))
db.execute( db.execute(
update(Provider) update(GlobalModel)
.where(Provider.id == provider_id) .where(GlobalModel.name == model)
.values(monthly_used_usd=Provider.monthly_used_usd + actual_total_cost) .values(usage_count=GlobalModel.usage_count + 1)
) )
accounted, _charge_applied = cls._finalize_usage_billing( # 更新用户-模型调用次数计数器
db, cls._increment_user_model_usage(db, user, model)
usage=usage,
total_cost=total_cost,
status=status,
finalized_at=finalized_at,
)
if accounted: # 更新 Provider 月度使用量(原子操作)
_increment_provider_api_key_totals( if provider_id:
actual_total_cost = Decimal(str(usage_params["actual_total_cost_usd"]))
db.execute(
update(Provider)
.where(Provider.id == provider_id)
.values(monthly_used_usd=Provider.monthly_used_usd + actual_total_cost)
)
accounted, _charge_applied = cls._finalize_usage_billing(
db, db,
provider_api_key_id, usage=usage,
total_tokens=int(usage_params.get("total_tokens") or 0), total_cost=total_cost,
total_cost=_get_actual_total_cost_usd(usage_params), status=status,
finalized_at=finalized_at,
) )
dispatch_codex_quota_sync_from_response_headers( if accounted:
provider_api_key_id=provider_api_key_id, _increment_provider_api_key_totals(
response_headers=response_headers, db,
db=db, provider_api_key_id,
) total_tokens=int(usage_params.get("total_tokens") or 0),
total_cost=_get_actual_total_cost_usd(usage_params),
)
db.commit() # 立即提交事务,释放数据库锁 dispatch_codex_quota_sync_from_response_headers(
return usage provider_api_key_id=provider_api_key_id,
response_headers=response_headers,
db=db,
)
db.commit() # 立即提交事务,释放数据库锁
return usage
return await asyncio.to_thread(_sync_record)
@classmethod @classmethod
async def record_usage( async def record_usage(

View File

@@ -0,0 +1,66 @@
from __future__ import annotations
import time
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_manager import TunnelStreamError
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 test_record_loop_lag_warning_does_not_degrade() -> None:
manager = _build_manager()
manager._record_loop_lag(1.5)
assert manager._degraded_until == 0.0
def test_record_loop_lag_degrades_manager(monkeypatch: pytest.MonkeyPatch) -> None:
manager = _build_manager()
now = 1234.0
monkeypatch.setattr("src.services.proxy_node.hub_transport._time.monotonic", lambda: now)
manager._record_loop_lag(4.0)
assert manager._degraded_until == pytest.approx(now + 12.0)
@pytest.mark.asyncio
async def test_send_request_rejects_while_manager_degraded(
monkeypatch: pytest.MonkeyPatch,
) -> None:
manager = _build_manager()
manager._degraded_until = time.monotonic() + 5.0
async def _fake_ensure_connected() -> None:
return None
monkeypatch.setattr(manager, "ensure_connected", _fake_ensure_connected)
with pytest.raises(TunnelStreamError, match="event loop degraded"):
await manager.send_request(
"node-1",
method="POST",
url="https://example.com/v1/chat/completions",
headers={"content-type": "application/json"},
body=b"{}",
timeout=5.0,
)
assert manager._pending_streams == {}