mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
refactor(hub): 用本地 HTTP relay 替代 Worker WebSocket 长连接
Hub 数据面改为 /local/relay/{node_id} HTTP 端点,Worker 通过本机
HTTP 请求转发 tunnel 帧,不再维护 /worker WebSocket 长连接。
Hub 侧:
- 新增 control_plane.rs: Hub 通过 HTTP 回调 Aether app 处理心跳 ACK 和节点状态变更
- 新增 local_relay.rs: 接收本地 HTTP 请求,在 Hub 内部打开 LocalStream 并透传到 proxy
- 移除 worker_conn.rs 及 Worker WebSocket 处理逻辑
- 简化 protocol.rs: 移除 NODE_STATUS 帧类型,抽取通用 encode_frame/decode_payload
Python 侧:
- 删除 tunnel_manager.py 及其 WebSocket 连接管理器 (HubConnectionManager)
- 简化 hub_transport.py 为 HTTP relay 调用
- 新增 src/api/internal/hub.py 接收 Hub 控制面回调 (heartbeat/node-status)
- hub_config.py 移除 WebSocket 相关配置,改为 HTTP relay URL
- service.py 新增 update_tunnel_status 方法
- 删除 src/api/admin/proxy_tunnel.py (旧管理接口)
- proxy_node 缓存 TTL 从 15s 降至 3s 加速状态感知
This commit is contained in:
804
aether-hub/Cargo.lock
generated
804
aether-hub/Cargo.lock
generated
File diff suppressed because it is too large
Load Diff
@@ -16,6 +16,9 @@ dashmap = "6"
|
||||
parking_lot = "0.12"
|
||||
flate2 = "1"
|
||||
futures-util = "0.3"
|
||||
bytes = "1"
|
||||
async-stream = "0.3"
|
||||
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls"] }
|
||||
|
||||
[profile.release]
|
||||
lto = true
|
||||
|
||||
85
aether-hub/src/control_plane.rs
Normal file
85
aether-hub/src/control_plane.rs
Normal file
@@ -0,0 +1,85 @@
|
||||
use reqwest::Client;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ControlPlaneClient {
|
||||
client: Option<Client>,
|
||||
base_url: String,
|
||||
}
|
||||
|
||||
impl ControlPlaneClient {
|
||||
pub fn new(base_url: String) -> Self {
|
||||
let client = Client::builder()
|
||||
.timeout(std::time::Duration::from_secs(10))
|
||||
.build()
|
||||
.ok();
|
||||
Self { client, base_url }
|
||||
}
|
||||
|
||||
pub fn disabled() -> Self {
|
||||
Self {
|
||||
client: None,
|
||||
base_url: String::new(),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn heartbeat_ack(&self, payload: &[u8]) -> Result<Vec<u8>, String> {
|
||||
let Some(client) = &self.client else {
|
||||
return Ok(b"{}".to_vec());
|
||||
};
|
||||
let url = format!(
|
||||
"{}/api/internal/hub/heartbeat",
|
||||
self.base_url.trim_end_matches('/')
|
||||
);
|
||||
let response = client
|
||||
.post(&url)
|
||||
.header("content-type", "application/json")
|
||||
.body(payload.to_vec())
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| format!("heartbeat callback request failed: {e}"))?;
|
||||
if !response.status().is_success() {
|
||||
return Err(format!(
|
||||
"heartbeat callback failed with status {}",
|
||||
response.status()
|
||||
));
|
||||
}
|
||||
response
|
||||
.bytes()
|
||||
.await
|
||||
.map(|bytes| bytes.to_vec())
|
||||
.map_err(|e| format!("heartbeat callback body read failed: {e}"))
|
||||
}
|
||||
|
||||
pub async fn push_node_status(
|
||||
&self,
|
||||
node_id: &str,
|
||||
connected: bool,
|
||||
conn_count: usize,
|
||||
) -> Result<(), String> {
|
||||
let Some(client) = &self.client else {
|
||||
return Ok(());
|
||||
};
|
||||
let url = format!(
|
||||
"{}/api/internal/hub/node-status",
|
||||
self.base_url.trim_end_matches('/')
|
||||
);
|
||||
let response = client
|
||||
.post(&url)
|
||||
.json(&serde_json::json!({
|
||||
"node_id": node_id,
|
||||
"connected": connected,
|
||||
"conn_count": conn_count,
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| format!("node-status callback request failed: {e}"))?;
|
||||
if response.status().is_success() {
|
||||
Ok(())
|
||||
} else {
|
||||
Err(format!(
|
||||
"node-status callback failed with status {}",
|
||||
response.status()
|
||||
))
|
||||
}
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
169
aether-hub/src/local_relay.rs
Normal file
169
aether-hub/src/local_relay.rs
Normal file
@@ -0,0 +1,169 @@
|
||||
use std::io;
|
||||
use std::net::SocketAddr;
|
||||
use std::time::Duration;
|
||||
|
||||
use async_stream::stream;
|
||||
use axum::body::{Body, Bytes};
|
||||
use axum::extract::{ConnectInfo, Path, State};
|
||||
use axum::http::{HeaderMap, HeaderName, HeaderValue, Response, StatusCode};
|
||||
use axum::response::IntoResponse;
|
||||
use tracing::warn;
|
||||
|
||||
use crate::hub::LocalBodyEvent;
|
||||
use crate::protocol;
|
||||
use crate::AppState;
|
||||
|
||||
pub const TUNNEL_ERROR_HEADER: &str = "x-aether-tunnel-error";
|
||||
|
||||
struct StreamGuard {
|
||||
hub: std::sync::Arc<crate::hub::HubRouter>,
|
||||
stream_id: u64,
|
||||
finished: bool,
|
||||
}
|
||||
|
||||
impl Drop for StreamGuard {
|
||||
fn drop(&mut self) {
|
||||
if !self.finished {
|
||||
self.hub
|
||||
.cancel_local_stream(self.stream_id, "local relay client dropped");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn relay_request(
|
||||
Path(node_id): Path<String>,
|
||||
State(state): State<AppState>,
|
||||
ConnectInfo(addr): ConnectInfo<SocketAddr>,
|
||||
body: Bytes,
|
||||
) -> impl IntoResponse {
|
||||
if !addr.ip().is_loopback() {
|
||||
return tunnel_error_response(
|
||||
StatusCode::FORBIDDEN,
|
||||
"forbidden",
|
||||
"local relay only accepts loopback requests",
|
||||
);
|
||||
}
|
||||
|
||||
let (meta, request_body) = match decode_envelope(body) {
|
||||
Ok(value) => value,
|
||||
Err(error) => {
|
||||
return tunnel_error_response(StatusCode::BAD_REQUEST, "bad_request", &error);
|
||||
}
|
||||
};
|
||||
|
||||
let stream = match state.hub.open_local_stream(&node_id, &meta, request_body) {
|
||||
Ok(stream) => stream,
|
||||
Err(error) => {
|
||||
return tunnel_error_response(StatusCode::SERVICE_UNAVAILABLE, "connect", &error);
|
||||
}
|
||||
};
|
||||
let request_guard = StreamGuard {
|
||||
hub: state.hub.clone(),
|
||||
stream_id: stream.id,
|
||||
finished: false,
|
||||
};
|
||||
|
||||
let wait_timeout = Duration::from_secs(meta.timeout.clamp(5, 300));
|
||||
let response_head = match stream.wait_headers(wait_timeout).await {
|
||||
Ok(response) => response,
|
||||
Err(error) => {
|
||||
state.hub.cancel_local_stream(stream.id, &error);
|
||||
return tunnel_error_response(StatusCode::GATEWAY_TIMEOUT, "timeout", &error);
|
||||
}
|
||||
};
|
||||
|
||||
let Some(mut body_rx) = stream.take_body_receiver() else {
|
||||
state
|
||||
.hub
|
||||
.cancel_local_stream(stream.id, "missing relay response body receiver");
|
||||
return tunnel_error_response(
|
||||
StatusCode::BAD_GATEWAY,
|
||||
"relay",
|
||||
"missing relay response body receiver",
|
||||
);
|
||||
};
|
||||
|
||||
let hub = state.hub.clone();
|
||||
let stream_id = stream.id;
|
||||
let body_stream = stream! {
|
||||
let mut guard = request_guard;
|
||||
guard.hub = hub;
|
||||
guard.stream_id = stream_id;
|
||||
while let Some(event) = body_rx.recv().await {
|
||||
match event {
|
||||
LocalBodyEvent::Chunk(chunk) => yield Ok::<Bytes, io::Error>(chunk),
|
||||
LocalBodyEvent::End => {
|
||||
guard.finished = true;
|
||||
break;
|
||||
}
|
||||
LocalBodyEvent::Error(error) => {
|
||||
guard.finished = true;
|
||||
yield Err(io::Error::other(error));
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
guard.finished = true;
|
||||
};
|
||||
|
||||
let mut builder = Response::builder().status(response_head.status);
|
||||
if let Some(headers) = builder.headers_mut() {
|
||||
append_headers(headers, &response_head.headers);
|
||||
}
|
||||
match builder.body(Body::from_stream(body_stream)) {
|
||||
Ok(response) => response,
|
||||
Err(error) => {
|
||||
warn!(error = %error, "failed to build relay response");
|
||||
tunnel_error_response(
|
||||
StatusCode::BAD_GATEWAY,
|
||||
"relay",
|
||||
"failed to build relay response",
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn decode_envelope(body: Bytes) -> Result<(protocol::RequestMeta, Bytes), String> {
|
||||
if body.len() < 4 {
|
||||
return Err("relay envelope too short".to_string());
|
||||
}
|
||||
let meta_len = u32::from_be_bytes([body[0], body[1], body[2], body[3]]) as usize;
|
||||
let meta_end = 4usize
|
||||
.checked_add(meta_len)
|
||||
.ok_or_else(|| "relay envelope length overflow".to_string())?;
|
||||
if body.len() < meta_end {
|
||||
return Err("relay envelope metadata truncated".to_string());
|
||||
}
|
||||
let meta = serde_json::from_slice::<protocol::RequestMeta>(&body[4..meta_end])
|
||||
.map_err(|e| format!("invalid relay metadata: {e}"))?;
|
||||
Ok((meta, body.slice(meta_end..)))
|
||||
}
|
||||
|
||||
fn append_headers(target: &mut HeaderMap, headers: &[(String, String)]) {
|
||||
for (name, value) in headers {
|
||||
let Ok(name) = HeaderName::from_bytes(name.as_bytes()) else {
|
||||
continue;
|
||||
};
|
||||
let Ok(value) = HeaderValue::from_str(value) else {
|
||||
continue;
|
||||
};
|
||||
target.append(name, value);
|
||||
}
|
||||
}
|
||||
|
||||
fn tunnel_error_response(status: StatusCode, kind: &str, message: &str) -> Response<Body> {
|
||||
let mut builder = Response::builder().status(status);
|
||||
if let Some(headers) = builder.headers_mut() {
|
||||
headers.insert(
|
||||
HeaderName::from_static(TUNNEL_ERROR_HEADER),
|
||||
HeaderValue::from_str(kind).unwrap_or_else(|_| HeaderValue::from_static("relay")),
|
||||
);
|
||||
headers.insert(
|
||||
axum::http::header::CONTENT_TYPE,
|
||||
HeaderValue::from_static("text/plain; charset=utf-8"),
|
||||
);
|
||||
}
|
||||
builder
|
||||
.body(Body::from(message.to_string()))
|
||||
.unwrap_or_else(|_| Response::new(Body::from("relay error")))
|
||||
}
|
||||
@@ -1,20 +1,23 @@
|
||||
mod control_plane;
|
||||
mod hub;
|
||||
mod local_relay;
|
||||
mod protocol;
|
||||
mod proxy_conn;
|
||||
mod worker_conn;
|
||||
|
||||
use std::sync::Arc;
|
||||
use std::net::SocketAddr;
|
||||
use std::time::Duration;
|
||||
|
||||
use axum::extract::ws::WebSocketUpgrade;
|
||||
use axum::extract::State;
|
||||
use axum::response::{IntoResponse, Json};
|
||||
use axum::routing::get;
|
||||
use axum::routing::{get, post};
|
||||
use axum::Router;
|
||||
use clap::Parser;
|
||||
use tracing::{info, warn};
|
||||
|
||||
use crate::control_plane::ControlPlaneClient;
|
||||
use crate::hub::{ConnConfig, HubRouter};
|
||||
use crate::local_relay::relay_request;
|
||||
|
||||
#[derive(Parser, Debug)]
|
||||
#[command(name = "aether-hub", about = "Tunnel Hub for Aether")]
|
||||
@@ -27,10 +30,6 @@ struct Args {
|
||||
#[arg(long, default_value_t = 0, env = "TUNNEL_HUB_PROXY_IDLE_TIMEOUT")]
|
||||
proxy_idle_timeout: u64,
|
||||
|
||||
/// Worker-side idle timeout in seconds (0 to disable)
|
||||
#[arg(long, default_value_t = 60, env = "TUNNEL_HUB_WORKER_IDLE_TIMEOUT")]
|
||||
worker_idle_timeout: u64,
|
||||
|
||||
/// Ping interval in seconds (for both sides)
|
||||
#[arg(long, default_value_t = 15, env = "TUNNEL_HUB_PING_INTERVAL")]
|
||||
ping_interval: u64,
|
||||
@@ -46,14 +45,21 @@ struct Args {
|
||||
env = "TUNNEL_HUB_OUTBOUND_QUEUE_CAPACITY"
|
||||
)]
|
||||
outbound_queue_capacity: usize,
|
||||
|
||||
/// Local Aether app base URL for control-plane callbacks
|
||||
#[arg(
|
||||
long,
|
||||
default_value = "http://127.0.0.1:8084",
|
||||
env = "TUNNEL_HUB_APP_BASE_URL"
|
||||
)]
|
||||
app_base_url: String,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct AppState {
|
||||
hub: Arc<HubRouter>,
|
||||
proxy_conn_cfg: ConnConfig,
|
||||
worker_conn_cfg: ConnConfig,
|
||||
max_streams: usize,
|
||||
pub struct AppState {
|
||||
pub hub: std::sync::Arc<HubRouter>,
|
||||
pub proxy_conn_cfg: ConnConfig,
|
||||
pub max_streams: usize,
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
@@ -68,7 +74,7 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
|
||||
let args = Args::parse();
|
||||
|
||||
let hub = HubRouter::new();
|
||||
let hub = HubRouter::new(ControlPlaneClient::new(args.app_base_url));
|
||||
let outbound_queue_capacity = args.outbound_queue_capacity.clamp(8, 4096);
|
||||
let ping_interval = Duration::from_secs(args.ping_interval);
|
||||
let state = AppState {
|
||||
@@ -78,11 +84,6 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
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,
|
||||
};
|
||||
|
||||
@@ -90,13 +91,17 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
.route("/health", get(health))
|
||||
.route("/stats", get(stats))
|
||||
.route("/proxy", get(ws_proxy))
|
||||
.route("/worker", get(ws_worker))
|
||||
.route("/local/relay/{node_id}", post(relay_request))
|
||||
.with_state(state);
|
||||
|
||||
let listener = tokio::net::TcpListener::bind(&args.bind).await?;
|
||||
info!(bind = %args.bind, "aether-hub started");
|
||||
|
||||
axum::serve(listener, app).await?;
|
||||
axum::serve(
|
||||
listener,
|
||||
app.into_make_service_with_connect_info::<SocketAddr>(),
|
||||
)
|
||||
.await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -160,10 +165,3 @@ async fn ws_proxy(
|
||||
})
|
||||
.into_response()
|
||||
}
|
||||
|
||||
async fn ws_worker(ws: WebSocketUpgrade, State(state): State<AppState>) -> impl IntoResponse {
|
||||
ws.max_frame_size(64 * 1024 * 1024)
|
||||
.on_upgrade(move |socket| {
|
||||
worker_conn::handle_worker_connection(socket, state.hub, state.worker_conn_cfg)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -22,8 +22,6 @@ pub const PONG: u8 = 0x11;
|
||||
pub const GOAWAY: u8 = 0x12;
|
||||
pub const HEARTBEAT_DATA: u8 = 0x13;
|
||||
pub const HEARTBEAT_ACK: u8 = 0x14;
|
||||
pub const NODE_STATUS: u8 = 0x15;
|
||||
|
||||
// Flags
|
||||
pub const FLAG_END_STREAM: u8 = 0x01;
|
||||
pub const FLAG_GZIP_COMPRESSED: u8 = 0x02;
|
||||
@@ -50,162 +48,89 @@ impl FrameHeader {
|
||||
payload_len: u32::from_be_bytes([data[6], data[7], data[8], data[9]]),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// Check if this is a stream-terminating frame
|
||||
#[inline]
|
||||
pub fn is_stream_terminal(&self) -> bool {
|
||||
self.msg_type == STREAM_END || self.msg_type == STREAM_ERROR
|
||||
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
|
||||
pub struct RequestMeta {
|
||||
pub method: String,
|
||||
pub url: String,
|
||||
pub headers: std::collections::HashMap<String, String>,
|
||||
#[serde(default = "default_timeout", deserialize_with = "deserialize_timeout")]
|
||||
pub timeout: u64,
|
||||
}
|
||||
|
||||
fn default_timeout() -> u64 {
|
||||
60
|
||||
}
|
||||
|
||||
fn deserialize_timeout<'de, D>(deserializer: D) -> Result<u64, D::Error>
|
||||
where
|
||||
D: serde::Deserializer<'de>,
|
||||
{
|
||||
#[derive(serde::Deserialize)]
|
||||
#[serde(untagged)]
|
||||
enum TimeoutValue {
|
||||
Int(u64),
|
||||
Float(f64),
|
||||
}
|
||||
|
||||
match <TimeoutValue as serde::Deserialize>::deserialize(deserializer)? {
|
||||
TimeoutValue::Int(v) => Ok(v),
|
||||
TimeoutValue::Float(v) => {
|
||||
if !v.is_finite() || v < 0.0 {
|
||||
return Err(serde::de::Error::custom(
|
||||
"timeout must be a non-negative finite number",
|
||||
));
|
||||
}
|
||||
if v.fract() != 0.0 {
|
||||
return Err(serde::de::Error::custom("timeout must be integer seconds"));
|
||||
}
|
||||
if v > (u64::MAX as f64) {
|
||||
return Err(serde::de::Error::custom("timeout is too large"));
|
||||
}
|
||||
Ok(v as u64)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct RequestHeadersExtracted {
|
||||
pub node_id: String,
|
||||
pub rebuilt_frame: Vec<u8>,
|
||||
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ResponseMeta {
|
||||
pub status: u16,
|
||||
pub headers: Vec<(String, String)>,
|
||||
}
|
||||
|
||||
pub fn encode_frame(stream_id: u32, msg_type: u8, flags: u8, payload: &[u8]) -> Vec<u8> {
|
||||
let mut buf = Vec::with_capacity(HEADER_SIZE + payload.len());
|
||||
buf.extend_from_slice(&stream_id.to_be_bytes());
|
||||
buf.push(msg_type);
|
||||
buf.push(flags);
|
||||
buf.extend_from_slice(&(payload.len() as u32).to_be_bytes());
|
||||
buf.extend_from_slice(payload);
|
||||
buf
|
||||
}
|
||||
|
||||
/// Encode a STREAM_ERROR frame for a given stream_id with an error message
|
||||
pub fn encode_stream_error(stream_id: u32, msg: &str) -> Vec<u8> {
|
||||
let payload = msg.as_bytes();
|
||||
let mut buf = Vec::with_capacity(HEADER_SIZE + payload.len());
|
||||
buf.extend_from_slice(&stream_id.to_be_bytes());
|
||||
buf.push(STREAM_ERROR);
|
||||
buf.push(0); // flags
|
||||
buf.extend_from_slice(&(payload.len() as u32).to_be_bytes());
|
||||
buf.extend_from_slice(payload);
|
||||
buf
|
||||
}
|
||||
|
||||
/// Encode a NODE_STATUS frame (stream_id=0, Hub-generated)
|
||||
pub fn encode_node_status(node_id: &str, connected: bool, conn_count: usize) -> Vec<u8> {
|
||||
let payload = serde_json::json!({
|
||||
"node_id": node_id,
|
||||
"connected": connected,
|
||||
"conn_count": conn_count,
|
||||
});
|
||||
let payload_bytes = payload.to_string().into_bytes();
|
||||
let mut buf = Vec::with_capacity(HEADER_SIZE + payload_bytes.len());
|
||||
buf.extend_from_slice(&0u32.to_be_bytes()); // stream_id = 0
|
||||
buf.push(NODE_STATUS);
|
||||
buf.push(0); // flags
|
||||
buf.extend_from_slice(&(payload_bytes.len() as u32).to_be_bytes());
|
||||
buf.extend_from_slice(&payload_bytes);
|
||||
buf
|
||||
encode_frame(stream_id, STREAM_ERROR, 0, msg.as_bytes())
|
||||
}
|
||||
|
||||
/// Encode a PING frame (stream_id=0)
|
||||
pub fn encode_ping() -> Vec<u8> {
|
||||
let mut buf = Vec::with_capacity(HEADER_SIZE);
|
||||
buf.extend_from_slice(&0u32.to_be_bytes());
|
||||
buf.push(PING);
|
||||
buf.push(0);
|
||||
buf.extend_from_slice(&0u32.to_be_bytes());
|
||||
buf
|
||||
encode_frame(0, PING, 0, &[])
|
||||
}
|
||||
|
||||
/// Encode a PONG frame (stream_id=0, echo payload)
|
||||
pub fn encode_pong(payload: &[u8]) -> Vec<u8> {
|
||||
let mut buf = Vec::with_capacity(HEADER_SIZE + payload.len());
|
||||
buf.extend_from_slice(&0u32.to_be_bytes());
|
||||
buf.push(PONG);
|
||||
buf.push(0);
|
||||
buf.extend_from_slice(&(payload.len() as u32).to_be_bytes());
|
||||
buf.extend_from_slice(payload);
|
||||
buf
|
||||
encode_frame(0, PONG, 0, payload)
|
||||
}
|
||||
|
||||
/// Encode a GOAWAY frame (stream_id=0)
|
||||
pub fn encode_goaway() -> Vec<u8> {
|
||||
let mut buf = Vec::with_capacity(HEADER_SIZE);
|
||||
buf.extend_from_slice(&0u32.to_be_bytes());
|
||||
buf.push(GOAWAY);
|
||||
buf.push(0);
|
||||
buf.extend_from_slice(&0u32.to_be_bytes());
|
||||
buf
|
||||
}
|
||||
|
||||
/// Rewrite the stream_id in raw frame bytes (first 4 bytes) -- near zero-copy
|
||||
#[inline]
|
||||
pub fn rewrite_stream_id(data: &mut [u8], new_stream_id: u32) {
|
||||
let bytes = new_stream_id.to_be_bytes();
|
||||
data[0] = bytes[0];
|
||||
data[1] = bytes[1];
|
||||
data[2] = bytes[2];
|
||||
data[3] = bytes[3];
|
||||
}
|
||||
|
||||
/// Get the payload portion of a raw frame (after the 10-byte header)
|
||||
#[inline]
|
||||
pub fn frame_payload(data: &[u8]) -> &[u8] {
|
||||
if data.len() > HEADER_SIZE {
|
||||
&data[HEADER_SIZE..]
|
||||
} else {
|
||||
&[]
|
||||
}
|
||||
}
|
||||
|
||||
/// Parse REQUEST_HEADERS payload, extract `node_id`, strip it from JSON,
|
||||
/// and rebuild a new REQUEST_HEADERS frame with `new_stream_id`.
|
||||
///
|
||||
/// If the source frame is gzip-compressed, this function will decode it first,
|
||||
/// then try to re-encode with gzip (only keeps compression when payload shrinks).
|
||||
pub fn rebuild_request_headers_without_node_id(
|
||||
data: &[u8],
|
||||
new_stream_id: u32,
|
||||
) -> Result<RequestHeadersExtracted, String> {
|
||||
let header = FrameHeader::parse(data).ok_or_else(|| "invalid frame header".to_string())?;
|
||||
if header.msg_type != REQUEST_HEADERS {
|
||||
return Err("frame is not REQUEST_HEADERS".to_string());
|
||||
}
|
||||
|
||||
let payload = frame_payload_by_header(data, &header)
|
||||
.ok_or_else(|| "incomplete REQUEST_HEADERS payload".to_string())?;
|
||||
|
||||
let decoded_payload = if header.flags & FLAG_GZIP_COMPRESSED != 0 {
|
||||
let mut decoder = GzDecoder::new(payload);
|
||||
let mut decoded = Vec::new();
|
||||
decoder
|
||||
.read_to_end(&mut decoded)
|
||||
.map_err(|e| format!("failed to decompress REQUEST_HEADERS: {e}"))?;
|
||||
decoded
|
||||
} else {
|
||||
payload.to_vec()
|
||||
};
|
||||
|
||||
let mut meta: serde_json::Value = serde_json::from_slice(&decoded_payload)
|
||||
.map_err(|e| format!("invalid REQUEST_HEADERS JSON: {e}"))?;
|
||||
let obj = meta
|
||||
.as_object_mut()
|
||||
.ok_or_else(|| "REQUEST_HEADERS payload must be a JSON object".to_string())?;
|
||||
|
||||
let node_id = obj
|
||||
.remove("node_id")
|
||||
.and_then(|v| v.as_str().map(|s| s.to_string()))
|
||||
.map(|s| s.trim().to_string())
|
||||
.filter(|s| !s.is_empty())
|
||||
.ok_or_else(|| "missing node_id in REQUEST_HEADERS".to_string())?;
|
||||
|
||||
let stripped_payload = serde_json::to_vec(&meta)
|
||||
.map_err(|e| format!("failed to encode REQUEST_HEADERS payload: {e}"))?;
|
||||
let (final_payload, flags) =
|
||||
maybe_recompress_payload(&stripped_payload, header.flags & FLAG_GZIP_COMPRESSED != 0)
|
||||
.map_err(|e| format!("failed to recompress REQUEST_HEADERS payload: {e}"))?;
|
||||
|
||||
let mut rebuilt = Vec::with_capacity(HEADER_SIZE + final_payload.len());
|
||||
rebuilt.extend_from_slice(&new_stream_id.to_be_bytes());
|
||||
rebuilt.push(REQUEST_HEADERS);
|
||||
rebuilt.push(flags);
|
||||
rebuilt.extend_from_slice(&(final_payload.len() as u32).to_be_bytes());
|
||||
rebuilt.extend_from_slice(&final_payload);
|
||||
|
||||
Ok(RequestHeadersExtracted {
|
||||
node_id,
|
||||
rebuilt_frame: rebuilt,
|
||||
})
|
||||
encode_frame(0, GOAWAY, 0, &[])
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn frame_payload_by_header<'a>(data: &'a [u8], header: &FrameHeader) -> Option<&'a [u8]> {
|
||||
pub fn frame_payload_by_header<'a>(data: &'a [u8], header: &FrameHeader) -> Option<&'a [u8]> {
|
||||
let payload_len = header.payload_len as usize;
|
||||
let end = HEADER_SIZE.checked_add(payload_len)?;
|
||||
if data.len() < end {
|
||||
@@ -214,6 +139,25 @@ fn frame_payload_by_header<'a>(data: &'a [u8], header: &FrameHeader) -> Option<&
|
||||
Some(&data[HEADER_SIZE..end])
|
||||
}
|
||||
|
||||
pub fn decode_payload(data: &[u8], header: &FrameHeader) -> Result<Vec<u8>, String> {
|
||||
let payload = frame_payload_by_header(data, header)
|
||||
.ok_or_else(|| "incomplete frame payload".to_string())?;
|
||||
if header.flags & FLAG_GZIP_COMPRESSED != 0 {
|
||||
let mut decoder = GzDecoder::new(payload);
|
||||
let mut decoded = Vec::new();
|
||||
decoder
|
||||
.read_to_end(&mut decoded)
|
||||
.map_err(|e| format!("failed to decompress payload: {e}"))?;
|
||||
Ok(decoded)
|
||||
} else {
|
||||
Ok(payload.to_vec())
|
||||
}
|
||||
}
|
||||
|
||||
pub fn compress_payload(payload: &[u8]) -> Result<(Vec<u8>, u8), std::io::Error> {
|
||||
maybe_recompress_payload(payload, true)
|
||||
}
|
||||
|
||||
fn maybe_recompress_payload(
|
||||
payload: &[u8],
|
||||
prefer_gzip: bool,
|
||||
|
||||
@@ -140,7 +140,7 @@ async fn run_proxy_reader(
|
||||
continue;
|
||||
}
|
||||
|
||||
hub.handle_proxy_frame(conn.id, &mut data);
|
||||
hub.handle_proxy_frame(conn.id, &mut data).await;
|
||||
}
|
||||
Some(Ok(Message::Close(_))) | None => {
|
||||
info!(conn_id = conn.id, node_id = %conn.node_id, "proxy WebSocket closed");
|
||||
|
||||
@@ -1,204 +0,0 @@
|
||||
/// Worker-side WebSocket connection handler
|
||||
///
|
||||
/// Handles the lifecycle of a single Gunicorn worker connection:
|
||||
/// accept -> read loop (route frames via Hub) -> cleanup
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
use std::sync::Arc;
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use axum::extract::ws::{Message, WebSocket};
|
||||
use futures_util::{SinkExt, StreamExt};
|
||||
use tokio::sync::{mpsc, watch};
|
||||
use tracing::{debug, info, warn};
|
||||
|
||||
use crate::hub::{ConnConfig, HubRouter, SendStatus, WorkerConn};
|
||||
use crate::protocol;
|
||||
|
||||
pub async fn handle_worker_connection(ws: WebSocket, hub: Arc<HubRouter>, cfg: ConnConfig) {
|
||||
let conn_id = hub.alloc_conn_id();
|
||||
let (mut ws_tx, ws_rx) = ws.split();
|
||||
|
||||
let (tx, mut rx) = mpsc::channel::<Message>(cfg.outbound_queue_capacity);
|
||||
let (close_tx, mut close_rx) = watch::channel(false);
|
||||
|
||||
let conn = Arc::new(WorkerConn::new(conn_id, tx, close_tx));
|
||||
hub.register_worker(conn.clone());
|
||||
|
||||
let writer = tokio::spawn(async move {
|
||||
loop {
|
||||
tokio::select! {
|
||||
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 liveness_clock = Instant::now();
|
||||
let last_seen_ms = Arc::new(AtomicU64::new(0));
|
||||
|
||||
let reader_hub = hub.clone();
|
||||
let reader_conn = conn.clone();
|
||||
let reader_last_seen_ms = last_seen_ms.clone();
|
||||
let liveness_conn = conn.clone();
|
||||
let mut reader = tokio::spawn(async move {
|
||||
run_worker_reader(
|
||||
ws_rx,
|
||||
reader_hub,
|
||||
conn_id,
|
||||
reader_conn,
|
||||
reader_last_seen_ms,
|
||||
liveness_clock,
|
||||
)
|
||||
.await;
|
||||
});
|
||||
|
||||
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;
|
||||
});
|
||||
|
||||
let reader_finished = tokio::select! {
|
||||
res = &mut reader => {
|
||||
if let Err(err) = res {
|
||||
warn!(worker_id = conn_id, error = %err, "worker reader task failed");
|
||||
}
|
||||
true
|
||||
}
|
||||
res = &mut liveness => {
|
||||
if let Err(err) = res {
|
||||
warn!(worker_id = conn_id, error = %err, "worker liveness task failed");
|
||||
}
|
||||
false
|
||||
}
|
||||
};
|
||||
|
||||
conn.request_close();
|
||||
if !reader_finished {
|
||||
reader.abort();
|
||||
let _ = reader.await;
|
||||
}
|
||||
if reader_finished {
|
||||
liveness.abort();
|
||||
let _ = liveness.await;
|
||||
}
|
||||
|
||||
hub.unregister_worker(conn_id);
|
||||
tokio::time::sleep(Duration::from_millis(100)).await;
|
||||
writer.abort();
|
||||
let _ = writer.await;
|
||||
}
|
||||
|
||||
async fn run_worker_reader(
|
||||
mut ws_rx: futures_util::stream::SplitStream<WebSocket>,
|
||||
hub: Arc<HubRouter>,
|
||||
conn_id: u64,
|
||||
conn: Arc<WorkerConn>,
|
||||
last_seen_ms: Arc<AtomicU64>,
|
||||
liveness_clock: Instant,
|
||||
) {
|
||||
loop {
|
||||
match ws_rx.next().await {
|
||||
Some(Ok(Message::Binary(data))) => {
|
||||
last_seen_ms.store(elapsed_millis(liveness_clock), Ordering::Relaxed);
|
||||
|
||||
let mut data = data.to_vec();
|
||||
if data.len() < protocol::HEADER_SIZE {
|
||||
debug!(worker_id = conn_id, "frame too small, skipping");
|
||||
continue;
|
||||
}
|
||||
|
||||
let header = match protocol::FrameHeader::parse(&data) {
|
||||
Some(h) => h,
|
||||
None => continue,
|
||||
};
|
||||
|
||||
if header.msg_type == protocol::HEARTBEAT_ACK {
|
||||
hub.handle_worker_heartbeat_ack(&mut data);
|
||||
continue;
|
||||
}
|
||||
|
||||
if let Some(err_msg) = hub.handle_worker_frame(conn_id, &mut data) {
|
||||
let err_frame = protocol::encode_stream_error(header.stream_id, &err_msg);
|
||||
let _ = conn.send(Message::Binary(err_frame.into()));
|
||||
}
|
||||
}
|
||||
Some(Ok(Message::Close(_))) | None => {
|
||||
info!(worker_id = conn_id, "worker WebSocket closed");
|
||||
break;
|
||||
}
|
||||
Some(Err(e)) => {
|
||||
warn!(worker_id = conn_id, error = %e, "worker WebSocket error");
|
||||
break;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
Reference in New Issue
Block a user