feat(tunnel): 引入 aether-hub 帧路由器,支持多 worker 共享 tunnel 连接

新增 Rust 实现的 aether-hub 服务,作为 Docker 容器内部 WebSocket 帧路由器,
解决多 Gunicorn worker 进程间 tunnel 连接隔离问题。

主要改动:
- 新增 aether-hub Rust 项目,实现 proxy/worker 双向帧路由与 stream_id 重映射
- 新增 HubConnectionManager/HubTunnelTransport,worker 通过 Hub 转发 tunnel 帧
- 新增 create_tunnel_transport 工厂函数,按运行环境自动选择 Hub 或直连模式
- 新增 NODE_STATUS 广播机制,Hub 实时通知所有 worker 节点连接状态变化
- CI/CD 新增 build-hub job,Dockerfile 集成 Hub 二进制,deploy.sh 适配 Hub 构建
- 默认 GUNICORN_WORKERS 从 4 降为 2
This commit is contained in:
fawney19
2026-03-02 02:43:14 +08:00
parent 97d42703da
commit 039a18c243
30 changed files with 3728 additions and 75 deletions

View File

@@ -32,9 +32,15 @@ ADMIN_PASSWORD=admin123456
# 应用端口(默认 8084 # 应用端口(默认 8084
# APP_PORT=8084 # APP_PORT=8084
# Gunicorn Worker 数量(默认 4 # Gunicorn Worker 数量(默认 2
# 建议最小设置为 2 # Docker 部署下 Tunnel Hub 为容器内部固定服务,可安全使用多 worker。
# GUNICORN_WORKERS=4 # 非 Docker 运行时若使用 ProxyNode tunnel建议设置为 1。
# 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 # Gunicorn Max Requests默认 4000
# Worker 处理指定数量请求后自动重启,防止内存泄漏 # Worker 处理指定数量请求后自动重启,防止内存泄漏

View File

@@ -14,6 +14,7 @@ on:
env: env:
REGISTRY: ghcr.io REGISTRY: ghcr.io
BASE_IMAGE_NAME: fawney19/aether-base BASE_IMAGE_NAME: fawney19/aether-base
HUB_IMAGE_NAME: fawney19/aether-hub
APP_IMAGE_NAME: fawney19/aether APP_IMAGE_NAME: fawney19/aether
# Base image hash inputs: # Base image hash inputs:
# - Dockerfile.base # - Dockerfile.base
@@ -174,9 +175,61 @@ jobs:
cache-to: type=gha,mode=max,scope=base cache-to: type=gha,mode=max,scope=base
platforms: linux/amd64,linux/arm64 platforms: linux/amd64,linux/arm64
build-hub:
runs-on: ubuntu-latest
permissions:
contents: read
packages: write
outputs:
hub_binary_image: ${{ steps.hub-ref.outputs.image }}
steps:
- uses: actions/checkout@v4
- name: Set up Docker Buildx
uses: docker/setup-buildx-action@v3
- name: Log in to Container Registry
uses: docker/login-action@v3
with:
registry: ${{ env.REGISTRY }}
username: ${{ github.actor }}
password: ${{ secrets.GITHUB_TOKEN }}
- name: Extract metadata for hub image
id: meta
uses: docker/metadata-action@v5
with:
images: ${{ env.REGISTRY }}/${{ env.HUB_IMAGE_NAME }}
tags: |
type=semver,pattern={{version}}
type=semver,pattern={{major}}.{{minor}}
type=raw,value=pre,enable=${{ contains(github.ref, '-') }}
type=raw,value=fix,enable=${{ contains(github.ref, '-fix') }}
type=raw,value=${{ github.sha }}
type=sha,prefix=
flavor: |
latest=auto
- name: Build and push hub image
uses: docker/build-push-action@v5
with:
context: ./aether-hub
file: ./aether-hub/Dockerfile
push: true
tags: ${{ steps.meta.outputs.tags }}
labels: ${{ steps.meta.outputs.labels }}
cache-from: type=gha,scope=hub
cache-to: type=gha,mode=max,scope=hub
platforms: linux/amd64,linux/arm64
- name: Export hub binary image reference
id: hub-ref
run: |
echo "image=${{ env.REGISTRY }}/${{ env.HUB_IMAGE_NAME }}:${{ github.sha }}" >> $GITHUB_OUTPUT
build-app: build-app:
needs: [check-base-changes, build-base] needs: [check-base-changes, build-base, build-hub]
if: always() && (needs.build-base.result == 'success' || needs.build-base.result == 'skipped') if: always() && (needs.build-base.result == 'success' || needs.build-base.result == 'skipped') && needs.build-hub.result == 'success'
runs-on: ubuntu-latest runs-on: ubuntu-latest
permissions: permissions:
contents: read contents: read
@@ -249,6 +302,8 @@ jobs:
context: . context: .
file: ./Dockerfile.app file: ./Dockerfile.app
push: true push: true
build-args: |
HUB_BINARY_IMAGE=${{ needs.build-hub.outputs.hub_binary_image }}
tags: ${{ steps.meta.outputs.tags }} tags: ${{ steps.meta.outputs.tags }}
labels: ${{ steps.meta.outputs.labels }} labels: ${{ steps.meta.outputs.labels }}
no-cache-filters: builder no-cache-filters: builder

View File

@@ -2,11 +2,15 @@
# 运行镜像:从 base 提取产物到精简运行时 # 运行镜像:从 base 提取产物到精简运行时
# 构建命令: docker build -f Dockerfile.app -t aether-app:latest . # 构建命令: docker build -f Dockerfile.app -t aether-app:latest .
# 用于 GitHub Actions CI官方源 # 用于 GitHub Actions CI官方源
ARG HUB_BINARY_IMAGE=ghcr.io/fawney19/aether-hub:latest
FROM ${HUB_BINARY_IMAGE} AS hub-bin
FROM aether-base:latest AS builder FROM aether-base:latest AS builder
WORKDIR /app WORKDIR /app
# 复制前端源码并构建CI 通过 no-cache-filters=builder 确保每次重建) # 复制前端源码并构建CI 通过 no-cache-filters=builder 确保每次重建)
COPY frontend/ ./frontend/ COPY frontend/ ./frontend/
RUN cd frontend && npm run build RUN cd frontend && npm run build
# ==================== 运行时镜像 ==================== # ==================== 运行时镜像 ====================
FROM python:3.13-slim FROM python:3.13-slim
WORKDIR /app WORKDIR /app
@@ -24,6 +28,7 @@ COPY --from=builder /usr/local/lib/python3.13/site-packages /usr/local/lib/pytho
COPY --from=builder /usr/local/bin/gunicorn /usr/local/bin/ COPY --from=builder /usr/local/bin/gunicorn /usr/local/bin/
COPY --from=builder /usr/local/bin/uvicorn /usr/local/bin/ COPY --from=builder /usr/local/bin/uvicorn /usr/local/bin/
COPY --from=builder /usr/local/bin/alembic /usr/local/bin/ COPY --from=builder /usr/local/bin/alembic /usr/local/bin/
COPY --from=hub-bin /usr/local/bin/aether-hub /usr/local/bin/
# 从 builder 阶段复制前端构建产物 # 从 builder 阶段复制前端构建产物
COPY --from=builder /app/frontend/dist /usr/share/nginx/html COPY --from=builder /app/frontend/dist /usr/share/nginx/html
RUN chmod -R 755 /usr/share/nginx/html RUN chmod -R 755 /usr/share/nginx/html
@@ -151,7 +156,16 @@ RUN printf '%s\n' \
'stdout_logfile_maxbytes=0' \ 'stdout_logfile_maxbytes=0' \
'stderr_logfile=/dev/stderr' \ 'stderr_logfile=/dev/stderr' \
'stderr_logfile_maxbytes=0' \ 'stderr_logfile_maxbytes=0' \
'environment=PYTHONUNBUFFERED=1,PYTHONIOENCODING=utf-8,LANG=C.UTF-8,LC_ALL=C.UTF-8,DOCKER_CONTAINER=true' > /etc/supervisor/conf.d/supervisord.conf 'environment=PYTHONUNBUFFERED=1,PYTHONIOENCODING=utf-8,LANG=C.UTF-8,LC_ALL=C.UTF-8,DOCKER_CONTAINER=true' \
'' \
'[program:tunnel-hub]' \
'command=/usr/local/bin/aether-hub --bind 0.0.0.0:8085' \
'autostart=true' \
'autorestart=true' \
'stdout_logfile=/dev/stdout' \
'stdout_logfile_maxbytes=0' \
'stderr_logfile=/dev/stderr' \
'stderr_logfile_maxbytes=0' > /etc/supervisor/conf.d/supervisord.conf
# 创建目录 # 创建目录
RUN mkdir -p /var/log/supervisor /app/logs /app/data RUN mkdir -p /var/log/supervisor /app/logs /app/data
# 入口脚本(启动前执行迁移) # 入口脚本(启动前执行迁移)
@@ -164,7 +178,7 @@ ENV PYTHONUNBUFFERED=1 \
LANG=C.UTF-8 \ LANG=C.UTF-8 \
LC_ALL=C.UTF-8 \ LC_ALL=C.UTF-8 \
PORT=8084 \ PORT=8084 \
GUNICORN_WORKERS=4 \ GUNICORN_WORKERS=2 \
MAX_REQUESTS=4000 MAX_REQUESTS=4000
EXPOSE 80 EXPOSE 80
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \ HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \

View File

@@ -2,6 +2,9 @@
# 运行镜像:从 base 提取产物到精简运行时(国内镜像源版本) # 运行镜像:从 base 提取产物到精简运行时(国内镜像源版本)
# 构建命令: docker build -f Dockerfile.app.local -t aether-app:latest . # 构建命令: docker build -f Dockerfile.app.local -t aether-app:latest .
# 用于本地/国内服务器部署 # 用于本地/国内服务器部署
ARG HUB_BINARY_IMAGE=aether-hub:local
FROM ${HUB_BINARY_IMAGE} AS hub-bin
FROM aether-base:latest AS builder FROM aether-base:latest AS builder
WORKDIR /app WORKDIR /app
@@ -32,6 +35,7 @@ COPY --from=builder /usr/local/lib/python3.13/site-packages /usr/local/lib/pytho
COPY --from=builder /usr/local/bin/gunicorn /usr/local/bin/ COPY --from=builder /usr/local/bin/gunicorn /usr/local/bin/
COPY --from=builder /usr/local/bin/uvicorn /usr/local/bin/ COPY --from=builder /usr/local/bin/uvicorn /usr/local/bin/
COPY --from=builder /usr/local/bin/alembic /usr/local/bin/ COPY --from=builder /usr/local/bin/alembic /usr/local/bin/
COPY --from=hub-bin /usr/local/bin/aether-hub /usr/local/bin/
# 从 builder 阶段复制前端构建产物 # 从 builder 阶段复制前端构建产物
COPY --from=builder /app/frontend/dist /usr/share/nginx/html COPY --from=builder /app/frontend/dist /usr/share/nginx/html
@@ -163,7 +167,16 @@ RUN printf '%s\n' \
'stdout_logfile_maxbytes=0' \ 'stdout_logfile_maxbytes=0' \
'stderr_logfile=/dev/stderr' \ 'stderr_logfile=/dev/stderr' \
'stderr_logfile_maxbytes=0' \ 'stderr_logfile_maxbytes=0' \
'environment=PYTHONUNBUFFERED=1,PYTHONIOENCODING=utf-8,LANG=C.UTF-8,LC_ALL=C.UTF-8,DOCKER_CONTAINER=true' > /etc/supervisor/conf.d/supervisord.conf 'environment=PYTHONUNBUFFERED=1,PYTHONIOENCODING=utf-8,LANG=C.UTF-8,LC_ALL=C.UTF-8,DOCKER_CONTAINER=true' \
'' \
'[program:tunnel-hub]' \
'command=/usr/local/bin/aether-hub --bind 0.0.0.0:8085' \
'autostart=true' \
'autorestart=true' \
'stdout_logfile=/dev/stdout' \
'stdout_logfile_maxbytes=0' \
'stderr_logfile=/dev/stderr' \
'stderr_logfile_maxbytes=0' > /etc/supervisor/conf.d/supervisord.conf
# 创建目录 # 创建目录
RUN mkdir -p /var/log/supervisor /app/logs /app/data RUN mkdir -p /var/log/supervisor /app/logs /app/data
@@ -179,7 +192,7 @@ ENV PYTHONUNBUFFERED=1 \
LANG=C.UTF-8 \ LANG=C.UTF-8 \
LC_ALL=C.UTF-8 \ LC_ALL=C.UTF-8 \
PORT=8084 \ PORT=8084 \
GUNICORN_WORKERS=4 \ GUNICORN_WORKERS=2 \
MAX_REQUESTS=4000 MAX_REQUESTS=4000
EXPOSE 80 EXPOSE 80

3
aether-hub/.dockerignore Normal file
View File

@@ -0,0 +1,3 @@
target/
.git/
.DS_Store

1229
aether-hub/Cargo.lock generated Normal file

File diff suppressed because it is too large Load Diff

23
aether-hub/Cargo.toml Normal file
View File

@@ -0,0 +1,23 @@
[package]
name = "aether-hub"
version = "0.1.0"
edition = "2021"
description = "Tunnel Hub for Aether - frame router between workers and proxies"
[dependencies]
tokio = { version = "1", features = ["full"] }
axum = { version = "0.8", features = ["ws"] }
serde = { version = "1", features = ["derive"] }
serde_json = "1"
tracing = "0.1"
tracing-subscriber = { version = "0.3", features = ["env-filter"] }
clap = { version = "4", features = ["derive", "env"] }
dashmap = "6"
parking_lot = "0.12"
flate2 = "1"
futures-util = "0.3"
[profile.release]
lto = true
strip = true
codegen-units = 1

27
aether-hub/Dockerfile Normal file
View File

@@ -0,0 +1,27 @@
# syntax=docker/dockerfile:1
FROM rust:1.85-slim AS builder
WORKDIR /build/aether-hub
# 先构建依赖层,最大化后续代码变更时的缓存命中
COPY Cargo.toml Cargo.lock ./
RUN mkdir src && printf 'fn main() {}\n' > src/main.rs
RUN --mount=type=cache,target=/usr/local/cargo/registry,sharing=locked \
--mount=type=cache,target=/build/aether-hub/target,sharing=locked \
cargo build --release --locked
RUN rm -rf src
COPY src ./src
RUN --mount=type=cache,target=/usr/local/cargo/registry,sharing=locked \
--mount=type=cache,target=/build/aether-hub/target,sharing=locked \
cargo build --release --locked && \
cp target/release/aether-hub /tmp/aether-hub
FROM debian:bookworm-slim
RUN apt-get update && apt-get install -y --no-install-recommends ca-certificates && \
rm -rf /var/lib/apt/lists/*
COPY --from=builder /tmp/aether-hub /usr/local/bin/aether-hub
EXPOSE 8085
ENTRYPOINT ["/usr/local/bin/aether-hub"]
CMD ["--bind", "0.0.0.0:8085"]

695
aether-hub/src/hub.rs Normal file
View File

@@ -0,0 +1,695 @@
/// HubRouter -- central frame routing engine
///
/// Manages proxy connections (node_id -> [ProxyConn]) and worker connections (conn_id -> WorkerConn).
/// Routes frames between workers and proxies with stream_id remapping.
use std::sync::atomic::{AtomicU32, AtomicU64, AtomicUsize, Ordering};
use std::sync::Arc;
use axum::extract::ws::Message;
use dashmap::DashMap;
use parking_lot::RwLock;
use tokio::sync::mpsc;
use tracing::{debug, info, warn};
use crate::protocol;
// ---------------------------------------------------------------------------
// Proxy connection
// ---------------------------------------------------------------------------
pub struct ProxyConn {
pub id: u64,
pub node_id: String,
pub node_name: String,
pub tx: mpsc::UnboundedSender<Message>,
next_stream_id: AtomicU32,
pub stream_count: AtomicUsize,
pub max_streams: usize,
}
impl ProxyConn {
pub fn new(
id: u64,
node_id: String,
node_name: String,
tx: mpsc::UnboundedSender<Message>,
max_streams: usize,
) -> Self {
Self {
id,
node_id,
node_name,
tx,
next_stream_id: AtomicU32::new(2), // even IDs, start at 2
stream_count: AtomicUsize::new(0),
max_streams,
}
}
/// Allocate a proxy-side stream_id (even numbers)
pub fn alloc_stream_id(&self) -> Option<u32> {
// Reserve one stream slot first (CAS to honor max_streams under contention).
let mut current = self.stream_count.load(Ordering::Relaxed);
loop {
if current >= self.max_streams {
return None;
}
match self.stream_count.compare_exchange_weak(
current,
current + 1,
Ordering::AcqRel,
Ordering::Relaxed,
) {
Ok(_) => break,
Err(observed) => current = observed,
}
}
let sid = loop {
let current_sid = self.next_stream_id.load(Ordering::Relaxed);
let next_sid = if current_sid >= 0xFFFF_FFFE {
2
} else {
current_sid + 2
};
if self
.next_stream_id
.compare_exchange_weak(current_sid, next_sid, Ordering::AcqRel, Ordering::Relaxed)
.is_ok()
{
break current_sid;
}
};
Some(sid)
}
pub fn release_stream(&self) {
let mut current = self.stream_count.load(Ordering::Relaxed);
while current > 0 {
match self.stream_count.compare_exchange_weak(
current,
current - 1,
Ordering::AcqRel,
Ordering::Relaxed,
) {
Ok(_) => return,
Err(observed) => current = observed,
}
}
}
pub fn send(&self, msg: Message) -> bool {
self.tx.send(msg).is_ok()
}
}
// ---------------------------------------------------------------------------
// Worker connection
// ---------------------------------------------------------------------------
pub struct WorkerConn {
pub id: u64,
pub tx: mpsc::UnboundedSender<Message>,
}
impl WorkerConn {
pub fn new(id: u64, tx: mpsc::UnboundedSender<Message>) -> Self {
Self { id, tx }
}
pub fn send(&self, msg: Message) -> bool {
self.tx.send(msg).is_ok()
}
}
// ---------------------------------------------------------------------------
// Stream mapping entry
// ---------------------------------------------------------------------------
#[derive(Debug, Clone, Copy)]
struct ProxySide {
proxy_conn_id: u64,
proxy_stream_id: u32,
}
#[derive(Debug, Clone, Copy)]
struct WorkerSide {
worker_conn_id: u64,
worker_stream_id: u32,
}
// ---------------------------------------------------------------------------
// HubRouter
// ---------------------------------------------------------------------------
pub struct HubRouter {
/// node_id -> list of proxy connections
proxy_conns: RwLock<std::collections::HashMap<String, Vec<Arc<ProxyConn>>>>,
/// proxy_conn_id -> Arc<ProxyConn> (for reverse lookup)
proxy_conns_by_id: DashMap<u64, Arc<ProxyConn>>,
/// worker_conn_id -> Arc<WorkerConn>
worker_conns: DashMap<u64, Arc<WorkerConn>>,
/// (worker_conn_id, worker_stream_id) -> ProxySide
worker_to_proxy: DashMap<(u64, u32), ProxySide>,
/// (proxy_conn_id, proxy_stream_id) -> WorkerSide
proxy_to_worker: DashMap<(u64, u32), WorkerSide>,
/// Connection ID generator
next_conn_id: AtomicU64,
/// Round-robin counter for heartbeat forwarding
heartbeat_rr: AtomicU64,
/// Heartbeat tag -> proxy_conn_id mapping (u32 tag fits in stream_id field)
heartbeat_tags: DashMap<u32, u64>,
/// Next heartbeat tag (wrapping u32)
next_heartbeat_tag: AtomicU32,
}
impl HubRouter {
pub fn new() -> Arc<Self> {
Arc::new(Self {
proxy_conns: RwLock::new(std::collections::HashMap::new()),
proxy_conns_by_id: DashMap::new(),
worker_conns: DashMap::new(),
worker_to_proxy: DashMap::new(),
proxy_to_worker: DashMap::new(),
next_conn_id: AtomicU64::new(1),
heartbeat_rr: AtomicU64::new(0),
heartbeat_tags: DashMap::new(),
next_heartbeat_tag: AtomicU32::new(1),
})
}
pub fn alloc_conn_id(&self) -> u64 {
self.next_conn_id.fetch_add(1, Ordering::Relaxed)
}
// -----------------------------------------------------------------------
// Proxy connection management
// -----------------------------------------------------------------------
pub fn register_proxy(&self, conn: Arc<ProxyConn>) {
let node_id = conn.node_id.clone();
let node_name = conn.node_name.clone();
let conn_id = conn.id;
self.proxy_conns_by_id.insert(conn_id, conn.clone());
let mut map = self.proxy_conns.write();
map.entry(node_id.clone()).or_default().push(conn);
let pool_size = map.get(&node_id).map(|v| v.len()).unwrap_or(0);
info!(
node_id = %node_id,
node_name = %node_name,
conn_id = conn_id,
pool_size = pool_size,
"proxy connected"
);
drop(map);
self.broadcast_node_status(&node_id);
}
pub fn unregister_proxy(&self, conn_id: u64, node_id: &str) {
self.proxy_conns_by_id.remove(&conn_id);
let mut map = self.proxy_conns.write();
if let Some(conns) = map.get_mut(node_id) {
conns.retain(|c| c.id != conn_id);
if conns.is_empty() {
map.remove(node_id);
}
}
let pool_size = map.get(node_id).map(|v| v.len()).unwrap_or(0);
info!(
node_id = %node_id,
conn_id = conn_id,
remaining = pool_size,
"proxy disconnected"
);
drop(map);
// Cancel all in-flight streams on this proxy connection
self.cancel_streams_for_proxy(conn_id);
self.broadcast_node_status(node_id);
}
/// Get least-loaded proxy connection for a node
fn get_proxy_conn(&self, node_id: &str) -> Option<Arc<ProxyConn>> {
let map = self.proxy_conns.read();
let conns = map.get(node_id)?;
conns
.iter()
.min_by_key(|c| c.stream_count.load(Ordering::Relaxed))
.cloned()
}
/// Get pool size for a node
fn proxy_conn_count(&self, node_id: &str) -> usize {
let map = self.proxy_conns.read();
map.get(node_id).map(|v| v.len()).unwrap_or(0)
}
// -----------------------------------------------------------------------
// Worker connection management
// -----------------------------------------------------------------------
pub fn register_worker(&self, conn: Arc<WorkerConn>) {
info!(worker_id = conn.id, "worker connected");
self.worker_conns.insert(conn.id, conn);
}
pub fn unregister_worker(&self, conn_id: u64) {
self.worker_conns.remove(&conn_id);
info!(worker_id = conn_id, "worker disconnected");
// Clean up all stream mappings for this worker
let to_remove: Vec<(u64, u32)> = self
.worker_to_proxy
.iter()
.filter(|e| e.key().0 == conn_id)
.map(|e| *e.key())
.collect();
for key in &to_remove {
if let Some((_, proxy_side)) = self.worker_to_proxy.remove(key) {
self.proxy_to_worker
.remove(&(proxy_side.proxy_conn_id, proxy_side.proxy_stream_id));
// Release stream count on proxy side
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_side.proxy_conn_id) {
pc.release_stream();
}
}
}
if !to_remove.is_empty() {
debug!(
worker_id = conn_id,
streams_cleaned = to_remove.len(),
"cleaned up worker streams"
);
}
}
// -----------------------------------------------------------------------
// Frame routing: Worker -> Proxy
// -----------------------------------------------------------------------
/// Handle a frame from a worker. Returns error message if routing fails.
pub fn handle_worker_frame(&self, worker_conn_id: u64, data: &mut [u8]) -> Option<String> {
let header = match protocol::FrameHeader::parse(data) {
Some(h) => h,
None => return Some("invalid frame".to_string()),
};
let expected_len = protocol::HEADER_SIZE + header.payload_len as usize;
if data.len() < expected_len {
return Some("incomplete frame payload".to_string());
}
match header.msg_type {
protocol::REQUEST_HEADERS => {
self.route_request_headers(worker_conn_id, header.stream_id, data)
}
protocol::REQUEST_BODY => {
if header.flags & protocol::FLAG_END_STREAM != 0 {
debug!(
worker_conn_id = worker_conn_id,
stream_id = header.stream_id,
"worker sent REQUEST_BODY with END_STREAM"
);
}
self.route_worker_to_proxy(worker_conn_id, header.stream_id, data, false);
None
}
protocol::STREAM_END | protocol::STREAM_ERROR => {
self.route_worker_to_proxy(worker_conn_id, header.stream_id, data, true);
None
}
protocol::GOAWAY => {
warn!(
worker_conn_id = worker_conn_id,
"received GOAWAY from worker connection"
);
None
}
protocol::PING => {
let payload = protocol::frame_payload(data).to_vec();
let pong = protocol::encode_pong(&payload);
if let Some(wc) = self.worker_conns.get(&worker_conn_id) {
let _ = wc.send(Message::Binary(pong.into()));
}
None
}
protocol::PONG => None, // Worker responded to our ping, nothing to do
_ => {
debug!(
msg_type = header.msg_type,
"unexpected frame type from worker"
);
None
}
}
}
/// Route REQUEST_HEADERS: extract node_id, allocate proxy stream, create mapping
fn route_request_headers(
&self,
worker_conn_id: u64,
worker_stream_id: u32,
data: &mut [u8],
) -> Option<String> {
// Parse payload to extract node_id, and pre-build frame with node_id stripped.
// stream_id is set to 0 first; we'll rewrite to proxy_stream_id after allocation.
let extracted = match protocol::rebuild_request_headers_without_node_id(data, 0) {
Ok(v) => v,
Err(e) => return Some(e),
};
let node_id = extracted.node_id;
// Find a proxy connection for this node
let proxy_conn = match self.get_proxy_conn(&node_id) {
Some(c) => c,
None => {
return Some(format!("no proxy connection for node {}", node_id));
}
};
// Allocate proxy-side stream_id
let proxy_stream_id = match proxy_conn.alloc_stream_id() {
Some(sid) => sid,
None => {
return Some(format!("stream limit reached for node {}", node_id));
}
};
let mut rebuilt_frame = extracted.rebuilt_frame;
protocol::rewrite_stream_id(&mut rebuilt_frame, proxy_stream_id);
// Record bidirectional mapping
self.worker_to_proxy.insert(
(worker_conn_id, worker_stream_id),
ProxySide {
proxy_conn_id: proxy_conn.id,
proxy_stream_id,
},
);
self.proxy_to_worker.insert(
(proxy_conn.id, proxy_stream_id),
WorkerSide {
worker_conn_id,
worker_stream_id,
},
);
if !proxy_conn.send(Message::Binary(rebuilt_frame.into())) {
// Send failed, clean up mapping
self.worker_to_proxy
.remove(&(worker_conn_id, worker_stream_id));
self.proxy_to_worker
.remove(&(proxy_conn.id, proxy_stream_id));
proxy_conn.release_stream();
return Some("proxy connection send failed".to_string());
}
None
}
/// Route non-header frames from worker to proxy (REQUEST_BODY etc.)
fn route_worker_to_proxy(
&self,
worker_conn_id: u64,
worker_stream_id: u32,
data: &mut [u8],
terminal: bool,
) {
let proxy_side = if terminal {
match self
.worker_to_proxy
.remove(&(worker_conn_id, worker_stream_id))
{
Some((_, ps)) => {
self.proxy_to_worker
.remove(&(ps.proxy_conn_id, ps.proxy_stream_id));
if let Some(pc) = self.proxy_conns_by_id.get(&ps.proxy_conn_id) {
pc.release_stream();
}
ps
}
None => return, // Silently discard -- mapping already removed (race condition)
}
} else {
match self
.worker_to_proxy
.get(&(worker_conn_id, worker_stream_id))
{
Some(entry) => *entry.value(),
None => return, // Silently discard -- mapping already removed (race condition)
}
};
// Rewrite stream_id
protocol::rewrite_stream_id(data, proxy_side.proxy_stream_id);
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_side.proxy_conn_id) {
let _ = pc.send(Message::Binary(data.to_vec().into()));
}
}
// -----------------------------------------------------------------------
// Frame routing: Proxy -> Worker
// -----------------------------------------------------------------------
/// Handle a frame from a proxy connection
pub fn handle_proxy_frame(&self, proxy_conn_id: u64, data: &mut [u8]) {
let header = match protocol::FrameHeader::parse(data) {
Some(h) => h,
None => return,
};
let expected_len = protocol::HEADER_SIZE + header.payload_len as usize;
if data.len() < expected_len {
return;
}
match header.msg_type {
protocol::RESPONSE_HEADERS | protocol::RESPONSE_BODY => {
self.route_proxy_to_worker(proxy_conn_id, header.stream_id, data, false);
}
_ if header.is_stream_terminal() => {
self.route_proxy_to_worker(proxy_conn_id, header.stream_id, data, true);
}
protocol::HEARTBEAT_DATA => {
self.forward_heartbeat_to_worker(proxy_conn_id, data);
}
protocol::PONG => {} // Proxy responded to our ping
protocol::GOAWAY => {
warn!(
proxy_conn_id = proxy_conn_id,
"received GOAWAY from proxy connection"
);
}
protocol::PING => {
// Proxy sent a ping, reply with pong
let payload = if data.len() > protocol::HEADER_SIZE {
&data[protocol::HEADER_SIZE..]
} else {
&[]
};
let pong = protocol::encode_pong(payload);
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
let _ = pc.send(Message::Binary(pong.into()));
}
}
_ => {
debug!(
msg_type = header.msg_type,
proxy_conn_id = proxy_conn_id,
"unexpected frame type from proxy"
);
}
}
}
/// Route response frames from proxy to worker
fn route_proxy_to_worker(
&self,
proxy_conn_id: u64,
proxy_stream_id: u32,
data: &mut [u8],
terminal: bool,
) {
let worker_side = if terminal {
// Remove mapping on terminal frames
match self
.proxy_to_worker
.remove(&(proxy_conn_id, proxy_stream_id))
{
Some((_, ws)) => {
self.worker_to_proxy
.remove(&(ws.worker_conn_id, ws.worker_stream_id));
// Release stream count
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
pc.release_stream();
}
ws
}
None => return, // Silently discard
}
} else {
match self.proxy_to_worker.get(&(proxy_conn_id, proxy_stream_id)) {
Some(entry) => *entry.value(),
None => return, // Silently discard
}
};
// Rewrite stream_id to worker-side
protocol::rewrite_stream_id(data, worker_side.worker_stream_id);
if let Some(wc) = self.worker_conns.get(&worker_side.worker_conn_id) {
let _ = wc.send(Message::Binary(data.to_vec().into()));
}
}
/// Forward HEARTBEAT_DATA to a worker (round-robin)
fn forward_heartbeat_to_worker(&self, proxy_conn_id: u64, data: &[u8]) {
// Pick a worker via round-robin
let workers: Vec<Arc<WorkerConn>> = self
.worker_conns
.iter()
.map(|e| e.value().clone())
.collect();
if workers.is_empty() {
debug!("no workers to forward heartbeat to");
return;
}
let idx = self.heartbeat_rr.fetch_add(1, Ordering::Relaxed) as usize % workers.len();
let worker = &workers[idx];
// Use a u32 tag in the stream_id field to identify the proxy connection.
// The tag maps to the full u64 proxy_conn_id via heartbeat_tags DashMap,
// avoiding truncation of u64 conn_id to u32.
// Skip 0 (reserved for control frames) via CAS loop.
let tag = loop {
let t = self.next_heartbeat_tag.fetch_add(1, Ordering::Relaxed);
if t != 0 {
break t;
}
};
self.heartbeat_tags.insert(tag, proxy_conn_id);
let mut forwarded = data.to_vec();
protocol::rewrite_stream_id(&mut forwarded, tag);
let _ = worker.send(Message::Binary(forwarded.into()));
}
/// Handle HEARTBEAT_ACK from worker -- route back to the proxy
pub fn handle_worker_heartbeat_ack(&self, data: &mut [u8]) {
let header = match protocol::FrameHeader::parse(data) {
Some(h) => h,
None => return,
};
// Recover the original proxy_conn_id from the tag stored in stream_id
let tag = header.stream_id;
let proxy_conn_id = match self.heartbeat_tags.remove(&tag) {
Some((_, id)) => id,
None => return,
};
// Reset stream_id to 0 before forwarding to proxy
protocol::rewrite_stream_id(data, 0);
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
let _ = pc.send(Message::Binary(data.to_vec().into()));
}
}
// -----------------------------------------------------------------------
// Stream cleanup
// -----------------------------------------------------------------------
/// Cancel all in-flight streams for a disconnected proxy connection
fn cancel_streams_for_proxy(&self, proxy_conn_id: u64) {
let to_remove: Vec<((u64, u32), WorkerSide)> = self
.proxy_to_worker
.iter()
.filter(|e| e.key().0 == proxy_conn_id)
.map(|e| (*e.key(), *e.value()))
.collect();
for ((p_conn_id, p_sid), worker_side) in &to_remove {
self.proxy_to_worker.remove(&(*p_conn_id, *p_sid));
self.worker_to_proxy
.remove(&(worker_side.worker_conn_id, worker_side.worker_stream_id));
// Send STREAM_ERROR to worker
let err_frame =
protocol::encode_stream_error(worker_side.worker_stream_id, "proxy disconnected");
if let Some(wc) = self.worker_conns.get(&worker_side.worker_conn_id) {
let _ = wc.send(Message::Binary(err_frame.into()));
}
}
if !to_remove.is_empty() {
warn!(
proxy_conn_id = proxy_conn_id,
streams_cancelled = to_remove.len(),
"cancelled in-flight streams due to proxy disconnect"
);
}
}
// -----------------------------------------------------------------------
// NODE_STATUS broadcast
// -----------------------------------------------------------------------
fn broadcast_node_status(&self, node_id: &str) {
let conn_count = self.proxy_conn_count(node_id);
let connected = conn_count > 0;
let frame = protocol::encode_node_status(node_id, connected, conn_count);
let msg = Message::Binary(frame.into());
let mut sent = 0usize;
for entry in self.worker_conns.iter() {
if entry.value().send(msg.clone()) {
sent += 1;
}
}
debug!(
node_id = %node_id,
connected = connected,
conn_count = conn_count,
workers_notified = sent,
"broadcast NODE_STATUS"
);
}
// -----------------------------------------------------------------------
// Stats
// -----------------------------------------------------------------------
pub fn stats(&self) -> HubStats {
let proxy_conns = self.proxy_conns.read();
let total_proxy = proxy_conns.values().map(|v| v.len()).sum();
let nodes = proxy_conns.len();
drop(proxy_conns);
HubStats {
proxy_connections: total_proxy,
worker_connections: self.worker_conns.len(),
nodes,
active_streams: self.worker_to_proxy.len(),
}
}
}
#[derive(serde::Serialize)]
pub struct HubStats {
pub proxy_connections: usize,
pub worker_connections: usize,
pub nodes: usize,
pub active_streams: usize,
}

159
aether-hub/src/main.rs Normal file
View File

@@ -0,0 +1,159 @@
mod hub;
mod protocol;
mod proxy_conn;
mod worker_conn;
use std::sync::Arc;
use std::time::Duration;
use axum::extract::ws::WebSocketUpgrade;
use axum::extract::State;
use axum::response::{IntoResponse, Json};
use axum::routing::get;
use axum::Router;
use clap::Parser;
use tracing::{info, warn};
use crate::hub::HubRouter;
#[derive(Parser, Debug)]
#[command(name = "aether-hub", about = "Tunnel Hub for Aether")]
struct Args {
/// Bind address
#[arg(long, default_value = "0.0.0.0:8085", env = "TUNNEL_HUB_BIND")]
bind: String,
/// Proxy-side idle timeout in seconds
#[arg(long, default_value_t = 90, env = "TUNNEL_HUB_PROXY_IDLE_TIMEOUT")]
proxy_idle_timeout: u64,
/// Worker-side idle timeout in seconds
#[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,
/// Max concurrent streams per proxy connection
#[arg(long, default_value_t = 2048, env = "TUNNEL_HUB_MAX_STREAMS")]
max_streams: usize,
}
#[derive(Clone)]
struct AppState {
hub: Arc<HubRouter>,
proxy_idle_timeout: Duration,
worker_idle_timeout: Duration,
ping_interval: Duration,
max_streams: usize,
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
// Initialize tracing
tracing_subscriber::fmt()
.with_env_filter(
tracing_subscriber::EnvFilter::try_from_default_env()
.unwrap_or_else(|_| "aether_hub=info".into()),
)
.init();
let args = Args::parse();
let hub = HubRouter::new();
let state = AppState {
hub,
proxy_idle_timeout: Duration::from_secs(args.proxy_idle_timeout),
worker_idle_timeout: Duration::from_secs(args.worker_idle_timeout),
ping_interval: Duration::from_secs(args.ping_interval),
max_streams: args.max_streams,
};
let app = Router::new()
.route("/health", get(health))
.route("/stats", get(stats))
.route("/proxy", get(ws_proxy))
.route("/worker", get(ws_worker))
.with_state(state);
let listener = tokio::net::TcpListener::bind(&args.bind).await?;
info!(bind = %args.bind, "aether-hub started");
axum::serve(listener, app).await?;
Ok(())
}
// ---------------------------------------------------------------------------
// HTTP endpoints
// ---------------------------------------------------------------------------
async fn health() -> impl IntoResponse {
Json(serde_json::json!({"status": "ok"}))
}
async fn stats(State(state): State<AppState>) -> impl IntoResponse {
Json(state.hub.stats())
}
// ---------------------------------------------------------------------------
// WebSocket endpoints
// ---------------------------------------------------------------------------
async fn ws_proxy(
ws: WebSocketUpgrade,
State(state): State<AppState>,
headers: axum::http::HeaderMap,
) -> impl IntoResponse {
let node_id = headers
.get("x-node-id")
.and_then(|v| v.to_str().ok())
.unwrap_or("")
.trim()
.to_string();
let node_name = headers
.get("x-node-name")
.and_then(|v| v.to_str().ok())
.unwrap_or(&node_id)
.trim()
.to_string();
let max_streams: usize = headers
.get("x-tunnel-max-streams")
.and_then(|v| v.to_str().ok())
.and_then(|v| v.parse().ok())
.unwrap_or(state.max_streams)
.clamp(64, 2048);
if node_id.is_empty() {
warn!("proxy connection rejected: missing X-Node-ID header");
return axum::http::StatusCode::BAD_REQUEST.into_response();
}
ws.max_frame_size(64 * 1024 * 1024)
.on_upgrade(move |socket| {
proxy_conn::handle_proxy_connection(
socket,
state.hub,
node_id,
node_name,
max_streams,
state.ping_interval,
state.proxy_idle_timeout,
)
})
.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.ping_interval,
state.worker_idle_timeout,
)
})
}

232
aether-hub/src/protocol.rs Normal file
View File

@@ -0,0 +1,232 @@
/// Tunnel binary frame protocol
///
/// Frame format (10-byte header + payload):
/// | stream_id (4B) | msg_type (1B) | flags (1B) | payload_len (4B) | payload (NB) |
use std::io::Read;
use flate2::read::GzDecoder;
use flate2::write::GzEncoder;
use flate2::Compression;
pub const HEADER_SIZE: usize = 10;
// Message types
pub const REQUEST_HEADERS: u8 = 0x01;
pub const REQUEST_BODY: u8 = 0x02;
pub const RESPONSE_HEADERS: u8 = 0x03;
pub const RESPONSE_BODY: u8 = 0x04;
pub const STREAM_END: u8 = 0x05;
pub const STREAM_ERROR: u8 = 0x06;
pub const PING: u8 = 0x10;
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;
#[derive(Debug, Clone, Copy)]
pub struct FrameHeader {
pub stream_id: u32,
pub msg_type: u8,
pub flags: u8,
pub payload_len: u32,
}
impl FrameHeader {
/// Parse frame header from raw bytes (must be >= HEADER_SIZE)
#[inline]
pub fn parse(data: &[u8]) -> Option<Self> {
if data.len() < HEADER_SIZE {
return None;
}
Some(Self {
stream_id: u32::from_be_bytes([data[0], data[1], data[2], data[3]]),
msg_type: data[4],
flags: data[5],
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)]
pub struct RequestHeadersExtracted {
pub node_id: String,
pub rebuilt_frame: Vec<u8>,
}
/// 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 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 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 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,
})
}
#[inline]
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 {
return None;
}
Some(&data[HEADER_SIZE..end])
}
fn maybe_recompress_payload(
payload: &[u8],
prefer_gzip: bool,
) -> Result<(Vec<u8>, u8), std::io::Error> {
if !prefer_gzip {
return Ok((payload.to_vec(), 0));
}
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
std::io::Write::write_all(&mut encoder, payload)?;
let compressed = encoder.finish()?;
if compressed.len() < payload.len() {
Ok((compressed, FLAG_GZIP_COMPRESSED))
} else {
Ok((payload.to_vec(), 0))
}
}

View File

@@ -0,0 +1,144 @@
/// Proxy-side WebSocket connection handler
///
/// Handles the lifecycle of a single aether-proxy connection:
/// accept -> authenticate (headers) -> read loop -> cleanup
use std::sync::Arc;
use std::time::Duration;
use axum::extract::ws::{Message, WebSocket};
use futures_util::{SinkExt, StreamExt};
use tokio::sync::mpsc;
use tracing::{debug, info, warn};
use crate::hub::{HubRouter, ProxyConn};
use crate::protocol;
/// Maximum single frame size: 64 MB
const MAX_FRAME_SIZE: usize = 64 * 1024 * 1024;
pub async fn handle_proxy_connection(
ws: WebSocket,
hub: Arc<HubRouter>,
node_id: String,
node_name: String,
max_streams: usize,
ping_interval: Duration,
idle_timeout: Duration,
) {
let conn_id = hub.alloc_conn_id();
let (mut ws_tx, ws_rx) = ws.split();
// Create channel for outbound messages
let (tx, mut rx) = mpsc::unbounded_channel::<Message>();
let conn = Arc::new(ProxyConn::new(
conn_id,
node_id.clone(),
node_name.clone(),
tx,
max_streams,
));
hub.register_proxy(conn.clone());
// Spawn writer task: drains channel -> WebSocket
let writer = tokio::spawn(async move {
while let Some(msg) = rx.recv().await {
if ws_tx.send(msg).await.is_err() {
break;
}
}
let _ = ws_tx.close().await;
});
// Spawn ping task
let ping_tx = conn.tx.clone();
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_node_id = node_id.clone();
let reader_tx = conn.tx.clone();
let reader = tokio::spawn(async move {
run_proxy_reader(
ws_rx,
reader_hub,
conn_id,
reader_node_id,
reader_tx,
idle_timeout,
)
.await;
});
// Wait for reader to end, then cleanup writer/ping and unregister from hub.
let _ = reader.await;
ping_task.abort();
writer.abort();
hub.unregister_proxy(conn_id, &node_id);
}
async fn run_proxy_reader(
mut ws_rx: futures_util::stream::SplitStream<WebSocket>,
hub: Arc<HubRouter>,
conn_id: u64,
node_id: String,
tx: mpsc::UnboundedSender<Message>,
idle_timeout: Duration,
) {
let mut oversized_count = 0u32;
loop {
let msg = tokio::select! {
msg = ws_rx.next() => msg,
_ = tokio::time::sleep(idle_timeout) => {
warn!(conn_id = conn_id, node_id = %node_id, "proxy idle timeout");
let _ = tx.send(Message::Binary(protocol::encode_goaway().into()));
break;
}
};
match msg {
Some(Ok(Message::Binary(data))) => {
let mut data = data.to_vec();
if data.len() > MAX_FRAME_SIZE {
oversized_count += 1;
warn!(
conn_id = conn_id,
size = data.len(),
"oversized frame from proxy"
);
if oversized_count >= 5 {
warn!(conn_id = conn_id, "too many oversized frames, closing");
break;
}
continue;
}
oversized_count = 0;
if data.len() < protocol::HEADER_SIZE {
debug!(conn_id = conn_id, "frame too small, skipping");
continue;
}
hub.handle_proxy_frame(conn_id, &mut data);
}
Some(Ok(Message::Close(_))) | None => {
info!(conn_id = conn_id, node_id = %node_id, "proxy WebSocket closed");
break;
}
Some(Err(e)) => {
warn!(conn_id = conn_id, error = %e, "proxy WebSocket error");
break;
}
_ => {} // Ignore text/ping/pong at WS level
}
}
}

View File

@@ -0,0 +1,130 @@
/// 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::Arc;
use std::time::Duration;
use axum::extract::ws::{Message, WebSocket};
use futures_util::{SinkExt, StreamExt};
use tokio::sync::mpsc;
use tracing::{debug, info, warn};
use crate::hub::{HubRouter, WorkerConn};
use crate::protocol;
pub async fn handle_worker_connection(
ws: WebSocket,
hub: Arc<HubRouter>,
ping_interval: Duration,
idle_timeout: Duration,
) {
let conn_id = hub.alloc_conn_id();
let (mut ws_tx, ws_rx) = ws.split();
// Create channel for outbound messages
let (tx, mut rx) = mpsc::unbounded_channel::<Message>();
let conn = Arc::new(WorkerConn::new(conn_id, tx));
hub.register_worker(conn.clone());
// Spawn writer task
let writer = tokio::spawn(async move {
while let Some(msg) = rx.recv().await {
if ws_tx.send(msg).await.is_err() {
break;
}
}
let _ = ws_tx.close().await;
});
// Spawn ping task
let ping_tx = conn.tx.clone();
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_tx = conn.tx.clone();
let reader = tokio::spawn(async move {
run_worker_reader(
ws_rx,
reader_hub,
conn_id,
conn.clone(),
reader_tx,
idle_timeout,
)
.await;
});
// Wait for reader to end, then cleanup writer/ping and unregister from hub.
let _ = reader.await;
ping_task.abort();
writer.abort();
hub.unregister_worker(conn_id);
}
async fn run_worker_reader(
mut ws_rx: futures_util::stream::SplitStream<WebSocket>,
hub: Arc<HubRouter>,
conn_id: u64,
conn: Arc<WorkerConn>,
tx: mpsc::UnboundedSender<Message>,
idle_timeout: Duration,
) {
loop {
let msg = 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;
}
};
match msg {
Some(Ok(Message::Binary(data))) => {
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,
};
// HEARTBEAT_ACK from worker -> route back to proxy
if header.msg_type == protocol::HEARTBEAT_ACK {
hub.handle_worker_heartbeat_ack(&mut data);
continue;
}
// Regular frames: route via hub
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 _ = 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;
}
_ => {} // Ignore text/ping/pong at WS level
}
}
}

View File

@@ -4,6 +4,7 @@
# 用法: # 用法:
# 部署/更新: ./deploy.sh (自动检测所有变化) # 部署/更新: ./deploy.sh (自动检测所有变化)
# 强制重建: ./deploy.sh --rebuild-base # 强制重建: ./deploy.sh --rebuild-base
# 强制重建 Hub: ./deploy.sh --rebuild-hub
# 强制全部重建: ./deploy.sh --force # 强制全部重建: ./deploy.sh --force
set -e set -e
@@ -12,12 +13,23 @@ cd "$(dirname "$0")"
# 兼容 docker-compose 和 docker compose # 兼容 docker-compose 和 docker compose
if command -v docker-compose &> /dev/null; then if command -v docker-compose &> /dev/null; then
DC="docker-compose -f docker-compose.build.yml" DC="docker-compose -f docker-compose.build.yml"
USE_LEGACY_COMPOSE=true
else else
DC="docker compose -f docker-compose.build.yml" DC="docker compose -f docker-compose.build.yml"
USE_LEGACY_COMPOSE=false
fi fi
compose_up() {
if [ "$USE_LEGACY_COMPOSE" = true ]; then
$DC up -d --no-build "$@"
else
$DC up -d --no-build --pull never "$@"
fi
}
# 缓存文件 # 缓存文件
HASH_FILE=".deps-hash" HASH_FILE=".deps-hash"
HUB_HASH_FILE=".hub-hash"
CODE_HASH_FILE=".code-hash" CODE_HASH_FILE=".code-hash"
MIGRATION_HASH_FILE=".migration-hash" MIGRATION_HASH_FILE=".migration-hash"
@@ -63,6 +75,16 @@ calc_code_hash() {
} | md5sum | cut -d' ' -f1 } | md5sum | cut -d' ' -f1
} }
# 计算 Hub 文件的哈希值(本地构建 aether-hub:local
calc_hub_hash() {
{
cat aether-hub/Dockerfile 2>/dev/null
cat aether-hub/Cargo.toml 2>/dev/null
cat aether-hub/Cargo.lock 2>/dev/null
find aether-hub/src -type f -name "*.rs" 2>/dev/null | sort | xargs cat 2>/dev/null
} | md5sum | cut -d' ' -f1
}
# 计算迁移文件的哈希值 # 计算迁移文件的哈希值
calc_migration_hash() { calc_migration_hash() {
find alembic/versions -name "*.py" -type f 2>/dev/null | sort | xargs cat 2>/dev/null | md5sum | cut -d' ' -f1 find alembic/versions -name "*.py" -type f 2>/dev/null | sort | xargs cat 2>/dev/null | md5sum | cut -d' ' -f1
@@ -92,6 +114,18 @@ check_code_changed() {
return 0 return 0
} }
# 检查 Hub 是否变化
check_hub_changed() {
local current_hash=$(calc_hub_hash)
if [ -f "$HUB_HASH_FILE" ]; then
local saved_hash=$(cat "$HUB_HASH_FILE")
if [ "$current_hash" = "$saved_hash" ]; then
return 1
fi
fi
return 0
}
# 检查迁移是否变化 # 检查迁移是否变化
check_migration_changed() { check_migration_changed() {
local current_hash=$(calc_migration_hash) local current_hash=$(calc_migration_hash)
@@ -106,16 +140,24 @@ check_migration_changed() {
# 保存哈希 # 保存哈希
save_deps_hash() { calc_deps_hash > "$HASH_FILE"; } save_deps_hash() { calc_deps_hash > "$HASH_FILE"; }
save_hub_hash() { calc_hub_hash > "$HUB_HASH_FILE"; }
save_code_hash() { calc_code_hash > "$CODE_HASH_FILE"; } save_code_hash() { calc_code_hash > "$CODE_HASH_FILE"; }
save_migration_hash() { calc_migration_hash > "$MIGRATION_HASH_FILE"; } save_migration_hash() { calc_migration_hash > "$MIGRATION_HASH_FILE"; }
# 构建基础镜像 # 构建基础镜像
build_base() { build_base() {
echo ">>> Building base image (dependencies)..." echo ">>> Building base image (dependencies)..."
docker build -f Dockerfile.base.local -t aether-base:latest . docker build --pull=false -f Dockerfile.base.local -t aether-base:latest .
save_deps_hash save_deps_hash
} }
# 构建 Hub 镜像(本地)
build_hub() {
echo ">>> Building hub image (local)..."
docker build --pull=false -f aether-hub/Dockerfile -t aether-hub:local ./aether-hub
save_hub_hash
}
# 生成版本文件 # 生成版本文件
generate_version_file() { generate_version_file() {
# 从 git 获取版本号 # 从 git 获取版本号
@@ -138,7 +180,7 @@ EOF
build_app() { build_app() {
echo ">>> Building app image (code only)..." echo ">>> Building app image (code only)..."
generate_version_file generate_version_file
docker build -f Dockerfile.app.local -t aether-app:latest . docker build --pull=false --build-arg HUB_BINARY_IMAGE=aether-hub:local -f Dockerfile.app.local -t aether-app:latest .
save_code_hash save_code_hash
} }
@@ -189,8 +231,9 @@ print('Old version cleared')
if [ "$1" = "--force" ] || [ "$1" = "-f" ]; then if [ "$1" = "--force" ] || [ "$1" = "-f" ]; then
echo ">>> Force rebuilding everything..." echo ">>> Force rebuilding everything..."
build_base build_base
build_hub
build_app build_app
$DC up -d --force-recreate compose_up --force-recreate
sleep 3 sleep 3
run_migration run_migration
docker image prune -f docker image prune -f
@@ -206,13 +249,19 @@ if [ "$1" = "--rebuild-base" ] || [ "$1" = "-r" ]; then
exit 0 exit 0
fi fi
# 拉取最新代码 # 强制重建 Hub 镜像
echo ">>> Pulling latest code..." if [ "$1" = "--rebuild-hub" ]; then
git pull build_hub
echo ">>> Hub image rebuilt. Run ./deploy.sh to deploy."
exit 0
fi
echo ">>> Local-only mode: skip git pull."
# 标记是否需要重启 # 标记是否需要重启
NEED_RESTART=false NEED_RESTART=false
BASE_REBUILT=false BASE_REBUILT=false
HUB_REBUILT=false
# 检查基础镜像是否存在,或依赖是否变化 # 检查基础镜像是否存在,或依赖是否变化
if ! docker image inspect aether-base:latest >/dev/null 2>&1; then if ! docker image inspect aether-base:latest >/dev/null 2>&1; then
@@ -229,6 +278,21 @@ else
echo ">>> Dependencies unchanged." echo ">>> Dependencies unchanged."
fi fi
# 检查 Hub 镜像是否存在,或 Hub 代码是否变化
if ! docker image inspect aether-hub:local >/dev/null 2>&1; then
echo ">>> Hub image not found, building..."
build_hub
HUB_REBUILT=true
NEED_RESTART=true
elif check_hub_changed; then
echo ">>> Hub changed, rebuilding hub image..."
build_hub
HUB_REBUILT=true
NEED_RESTART=true
else
echo ">>> Hub unchanged."
fi
# 检查代码或迁移是否变化,或者 base 重建了app 依赖 base # 检查代码或迁移是否变化,或者 base 重建了app 依赖 base
# 注意:迁移文件打包在镜像中,所以迁移变化也需要重建 app 镜像 # 注意:迁移文件打包在镜像中,所以迁移变化也需要重建 app 镜像
MIGRATION_CHANGED=false MIGRATION_CHANGED=false
@@ -244,6 +308,10 @@ elif [ "$BASE_REBUILT" = true ]; then
echo ">>> Base image rebuilt, rebuilding app image..." echo ">>> Base image rebuilt, rebuilding app image..."
build_app build_app
NEED_RESTART=true NEED_RESTART=true
elif [ "$HUB_REBUILT" = true ]; then
echo ">>> Hub image rebuilt, rebuilding app image..."
build_app
NEED_RESTART=true
elif check_code_changed; then elif check_code_changed; then
echo ">>> Code changed, rebuilding app image..." echo ">>> Code changed, rebuilding app image..."
build_app build_app
@@ -265,10 +333,10 @@ fi
# 有变化时重启,或容器未运行时启动 # 有变化时重启,或容器未运行时启动
if [ "$NEED_RESTART" = true ]; then if [ "$NEED_RESTART" = true ]; then
echo ">>> Restarting services..." echo ">>> Restarting services..."
$DC up -d compose_up
elif [ "$CONTAINERS_RUNNING" = false ]; then elif [ "$CONTAINERS_RUNNING" = false ]; then
echo ">>> Containers not running, starting services..." echo ">>> Containers not running, starting services..."
$DC up -d compose_up
else else
echo ">>> No changes detected, skipping restart." echo ">>> No changes detected, skipping restart."
fi fi

View File

@@ -1,6 +1,7 @@
# Aether 部署配置 - 本地构建 # Aether 部署配置 - 本地构建
# 使用方法: # 使用方法:
# 首次构建 base: docker build -f Dockerfile.base -t aether-base:latest . # 首次构建 base: docker build -f Dockerfile.base -t aether-base:latest .
# 首次构建 hub: docker build -f aether-hub/Dockerfile -t aether-hub:local ./aether-hub
# 启动服务: docker compose -f docker-compose.build.yml up -d --build # 启动服务: docker compose -f docker-compose.build.yml up -d --build
services: services:
@@ -42,6 +43,8 @@ services:
build: build:
context: . context: .
dockerfile: Dockerfile.app.local dockerfile: Dockerfile.app.local
args:
HUB_BINARY_IMAGE: ${HUB_BINARY_IMAGE:-aether-hub:local}
image: aether-app:latest image: aether-app:latest
container_name: aether-app container_name: aether-app
env_file: env_file:

View File

@@ -495,7 +495,7 @@ class HTTPClientPool:
调用方不应关闭此 client其生命周期由 HTTPClientPool 管理。 调用方不应关闭此 client其生命周期由 HTTPClientPool 管理。
当 timeout 非 None 时(流式请求),每次创建新 client由调用方负责关闭。 当 timeout 非 None 时(流式请求),每次创建新 client由调用方负责关闭。
""" """
from src.services.proxy_node.tunnel_transport import TunnelTransport from src.services.proxy_node.tunnel_transport import create_tunnel_transport
t = timeout or httpx.Timeout( t = timeout or httpx.Timeout(
connect=config.http_connect_timeout, connect=config.http_connect_timeout,
@@ -507,7 +507,7 @@ class HTTPClientPool:
# 流式请求:每次创建新 client调用方负责关闭 # 流式请求:每次创建新 client调用方负责关闭
if timeout is not None: if timeout is not None:
transport = TunnelTransport(node_id, timeout=timeout_secs or 60.0) transport = create_tunnel_transport(node_id, timeout=timeout_secs or 60.0)
return httpx.AsyncClient(transport=transport, timeout=t) return httpx.AsyncClient(transport=transport, timeout=t)
# 非流式请求:复用缓存的 client加锁与 proxy_clients 保持一致) # 非流式请求:复用缓存的 client加锁与 proxy_clients 保持一致)
@@ -517,7 +517,7 @@ class HTTPClientPool:
if existing and not existing.is_closed: if existing and not existing.is_closed:
return existing return existing
transport = TunnelTransport(node_id, timeout=timeout_secs or 60.0) transport = create_tunnel_transport(node_id, timeout=timeout_secs or 60.0)
client = httpx.AsyncClient(transport=transport, timeout=t) client = httpx.AsyncClient(transport=transport, timeout=t)
cls._tunnel_clients[node_id] = client cls._tunnel_clients[node_id] = client
return client return client

View File

@@ -78,14 +78,22 @@ async def _on_startup() -> None:
from src.config import config from src.config import config
from src.services.proxy_node.health_scheduler import get_proxy_node_health_scheduler from src.services.proxy_node.health_scheduler import get_proxy_node_health_scheduler
from src.services.proxy_node.hub_config import get_hub_config
from src.utils.task_coordinator import StartupTaskCoordinator from src.utils.task_coordinator import StartupTaskCoordinator
logger = logging.getLogger("aether.modules.proxy_nodes") logger = logging.getLogger("aether.modules.proxy_nodes")
hub_enabled = get_hub_config().enabled
if config.worker_processes > 1: if config.worker_processes > 1 and not hub_enabled:
logger.warning( logger.warning(
"检测到 WEB_CONCURRENCY={}。Proxy tunnel 连接是进程内资源," "检测到 WEB_CONCURRENCY={}。Proxy tunnel 连接是进程内资源,"
"多 worker 场景可能出现节点显示 ONLINE 但当前 worker 无可用 tunnel 的情况。", "多 worker 场景可能出现节点显示 ONLINE 但当前 worker 无可用 tunnel 的情况。"
"建议设置 GUNICORN_WORKERS/WEB_CONCURRENCY=1。",
config.worker_processes,
)
elif config.worker_processes > 1 and hub_enabled:
logger.info(
"检测到 WEB_CONCURRENCY={}Hub 模式已启用,允许多 worker 共享 tunnel。",
config.worker_processes, config.worker_processes,
) )
@@ -111,14 +119,22 @@ async def _on_shutdown() -> None:
import logging import logging
from src.services.proxy_node.health_scheduler import get_proxy_node_health_scheduler from src.services.proxy_node.health_scheduler import get_proxy_node_health_scheduler
from src.services.proxy_node.tunnel_manager import get_tunnel_manager from src.services.proxy_node.hub_config import get_hub_config
from src.utils.task_coordinator import StartupTaskCoordinator from src.utils.task_coordinator import StartupTaskCoordinator
logger = logging.getLogger("aether.modules.proxy_nodes") logger = logging.getLogger("aether.modules.proxy_nodes")
# 先向所有 tunnel 连接发送 GoAway让 proxy 端立即重连到其他 worker hub_enabled = get_hub_config().enabled
manager = get_tunnel_manager() if hub_enabled:
await manager.shutdown_all() from src.services.proxy_node.hub_transport import shutdown_hub_connection_manager
await shutdown_hub_connection_manager()
else:
# 先向所有 tunnel 连接发送 GoAway让 proxy 端立即重连到其他 worker
from src.services.proxy_node.tunnel_manager import get_tunnel_manager
manager = get_tunnel_manager()
await manager.shutdown_all()
from src.clients import get_redis_client from src.clients import get_redis_client

View File

@@ -214,9 +214,9 @@ async def _get_acw_cookie(
"verify": get_ssl_context(), "verify": get_ssl_context(),
} }
if tunnel_node_id: if tunnel_node_id:
from src.services.proxy_node.tunnel_transport import TunnelTransport from src.services.proxy_node.tunnel_transport import create_tunnel_transport
client_kwargs["transport"] = TunnelTransport(tunnel_node_id, timeout=timeout) client_kwargs["transport"] = create_tunnel_transport(tunnel_node_id, timeout=timeout)
elif proxy: elif proxy:
client_kwargs["proxy"] = proxy client_kwargs["proxy"] = proxy
logger.debug(f"获取 acw_sc__v2 Cookie 使用代理: {proxy}") logger.debug(f"获取 acw_sc__v2 Cookie 使用代理: {proxy}")

View File

@@ -127,9 +127,9 @@ class ProviderConnector(ABC):
""" """
transport = None transport = None
if self._tunnel_node_id: if self._tunnel_node_id:
from src.services.proxy_node.tunnel_transport import TunnelTransport from src.services.proxy_node.tunnel_transport import create_tunnel_transport
transport = TunnelTransport(self._tunnel_node_id, timeout=self._timeout) transport = create_tunnel_transport(self._tunnel_node_id, timeout=self._timeout)
elif self._proxy: elif self._proxy:
transport = httpx.AsyncHTTPTransport(proxy=self._proxy) transport = httpx.AsyncHTTPTransport(proxy=self._proxy)

View File

@@ -197,9 +197,9 @@ class NekoCodeArchitecture(ProviderArchitecture):
proxy, tunnel_node_id = resolve_ops_proxy_config(config) proxy, tunnel_node_id = resolve_ops_proxy_config(config)
if tunnel_node_id: if tunnel_node_id:
from src.services.proxy_node.tunnel_transport import TunnelTransport from src.services.proxy_node.tunnel_transport import create_tunnel_transport
client_kwargs["transport"] = TunnelTransport(tunnel_node_id, timeout=10.0) client_kwargs["transport"] = create_tunnel_transport(tunnel_node_id, timeout=10.0)
elif proxy: elif proxy:
client_kwargs["proxy"] = proxy client_kwargs["proxy"] = proxy

View File

@@ -106,9 +106,9 @@ class _Sub2ApiTokenMixin:
"""获取不带 auth hook 的裸 HTTP 客户端(用于登录/刷新 token""" """获取不带 auth hook 的裸 HTTP 客户端(用于登录/刷新 token"""
transport = None transport = None
if self._tunnel_node_id: if self._tunnel_node_id:
from src.services.proxy_node.tunnel_transport import TunnelTransport from src.services.proxy_node.tunnel_transport import create_tunnel_transport
transport = TunnelTransport(self._tunnel_node_id, timeout=self._timeout) transport = create_tunnel_transport(self._tunnel_node_id, timeout=self._timeout)
elif self._proxy: elif self._proxy:
transport = httpx.AsyncHTTPTransport(proxy=self._proxy) transport = httpx.AsyncHTTPTransport(proxy=self._proxy)
async with httpx.AsyncClient( async with httpx.AsyncClient(
@@ -488,9 +488,9 @@ class Sub2ApiArchitecture(ProviderArchitecture):
"verify": get_ssl_context(), "verify": get_ssl_context(),
} }
if tunnel_node_id: if tunnel_node_id:
from src.services.proxy_node.tunnel_transport import TunnelTransport from src.services.proxy_node.tunnel_transport import create_tunnel_transport
client_kwargs["transport"] = TunnelTransport(tunnel_node_id, timeout=30.0) client_kwargs["transport"] = create_tunnel_transport(tunnel_node_id, timeout=30.0)
elif proxy: elif proxy:
client_kwargs["proxy"] = proxy client_kwargs["proxy"] = proxy

View File

@@ -245,9 +245,9 @@ class YesCodeArchitecture(ProviderArchitecture):
"verify": get_ssl_context(), "verify": get_ssl_context(),
} }
if tunnel_node_id: if tunnel_node_id:
from src.services.proxy_node.tunnel_transport import TunnelTransport from src.services.proxy_node.tunnel_transport import create_tunnel_transport
client_kwargs["transport"] = TunnelTransport(tunnel_node_id, timeout=10.0) client_kwargs["transport"] = create_tunnel_transport(tunnel_node_id, timeout=10.0)
elif proxy: elif proxy:
client_kwargs["proxy"] = proxy client_kwargs["proxy"] = proxy

View File

@@ -1045,9 +1045,9 @@ class ProviderOpsService:
"verify": get_ssl_context(), "verify": get_ssl_context(),
} }
if tunnel_node_id: if tunnel_node_id:
from src.services.proxy_node.tunnel_transport import TunnelTransport from src.services.proxy_node.tunnel_transport import create_tunnel_transport
client_kwargs["transport"] = TunnelTransport(tunnel_node_id, timeout=30.0) client_kwargs["transport"] = create_tunnel_transport(tunnel_node_id, timeout=30.0)
logger.debug("使用 tunnel 代理: node_id={}", tunnel_node_id) logger.debug("使用 tunnel 代理: node_id={}", tunnel_node_id)
elif proxy: elif proxy:
client_kwargs["proxy"] = proxy client_kwargs["proxy"] = proxy

View File

@@ -0,0 +1,70 @@
"""
Tunnel Hub 配置
控制 worker 是否通过 aether-hub 转发 tunnel 帧。
设计约束:
- Hub 作为 Docker 内部固定服务运行
- 不对外暴露运行时配置项(不依赖 TUNNEL_HUB_* 环境变量)
"""
from __future__ import annotations
import os
from dataclasses import dataclass
_DOCKER_HUB_URL = "ws://127.0.0.1:8085"
_DOCKER_HUB_CONNECT_TIMEOUT_SECONDS = 5.0
_DOCKER_HUB_PING_INTERVAL_SECONDS = 15.0
_DOCKER_HUB_SEND_TIMEOUT_SECONDS = 10.0
_DOCKER_HUB_MAX_STREAMS = 2048
_DOCKER_HUB_MAX_FRAME_SIZE = 64 * 1024 * 1024
@dataclass(frozen=True)
class HubConfig:
enabled: bool
url: str
connect_timeout_seconds: float
ping_interval_seconds: float
send_timeout_seconds: float
max_streams: int
max_frame_size: int
@property
def worker_ws_url(self) -> str:
return f"{self.url.rstrip('/')}/worker"
_hub_config: HubConfig | None = None
def _is_docker_runtime() -> bool:
if os.getenv("DOCKER_CONTAINER", "").strip().lower() == "true":
return True
return os.path.exists("/.dockerenv")
def get_hub_config() -> HubConfig:
"""读取 Hub 配置(进程内缓存)。"""
global _hub_config
if _hub_config is not None:
return _hub_config
docker_runtime = _is_docker_runtime()
_hub_config = HubConfig(
enabled=docker_runtime,
url=_DOCKER_HUB_URL,
connect_timeout_seconds=_DOCKER_HUB_CONNECT_TIMEOUT_SECONDS,
ping_interval_seconds=_DOCKER_HUB_PING_INTERVAL_SECONDS,
send_timeout_seconds=_DOCKER_HUB_SEND_TIMEOUT_SECONDS,
max_streams=_DOCKER_HUB_MAX_STREAMS,
max_frame_size=_DOCKER_HUB_MAX_FRAME_SIZE,
)
return _hub_config
def reset_hub_config_cache() -> None:
"""测试或热更新场景下清理配置缓存。"""
global _hub_config
_hub_config = None

View File

@@ -0,0 +1,639 @@
"""
Hub 模式 tunnel transport
Worker 通过单条到 aether-hub 的 WebSocket 长连接转发 tunnel 帧。
"""
from __future__ import annotations
import asyncio
import gzip
import json
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Any
import aiohttp
import httpx
from aiohttp import WSMsgType
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
if TYPE_CHECKING:
from collections.abc import AsyncGenerator, Coroutine
_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)
_HOP_BY_HOP_HEADERS = frozenset(
{
"host",
"transfer-encoding",
"content-length",
"connection",
"upgrade",
"keep-alive",
"proxy-authorization",
"proxy-connection",
"te",
"trailer",
}
)
_HOP_BY_HOP_HEADERS_BYTES = frozenset(h.encode("ascii") for h in _HOP_BY_HOP_HEADERS)
class HubConnectionManager:
"""Worker 进程级 Hub 连接管理器(单例)。"""
def __init__(self, config: HubConfig | None = None) -> None:
self._config = config or get_hub_config()
self._session: aiohttp.ClientSession | None = None
self._ws: aiohttp.ClientWebSocketResponse | None = None
self._connect_lock = asyncio.Lock()
self._write_lock = asyncio.Lock()
self._next_stream_id = 2
self._pending_streams: dict[int, _StreamState] = {}
self._reader_task: asyncio.Task[None] | None = None
self._ping_task: asyncio.Task[None] | None = None
self._reconnect_task: asyncio.Task[None] | None = None
self._background_tasks: set[asyncio.Task[None]] = set()
self._closing = False
@property
def is_connected(self) -> bool:
ws = self._ws
return ws is not None and not ws.closed
def _background(self, coro: Coroutine[Any, Any, None]) -> None:
task = asyncio.create_task(coro)
self._background_tasks.add(task)
task.add_done_callback(self._background_tasks.discard)
async def ensure_connected(self) -> None:
if self._closing:
raise TunnelStreamError("hub connection manager is shutting down")
if self.is_connected:
return
async with self._connect_lock:
if self._closing:
raise TunnelStreamError("hub connection manager is shutting down")
if self.is_connected:
return
try:
await self._connect_once()
except Exception as e:
self._start_reconnect_loop()
raise TunnelStreamError(f"failed to connect hub worker channel: {e}") from e
async def _ensure_session(self) -> aiohttp.ClientSession:
if self._session is None or self._session.closed:
self._session = aiohttp.ClientSession()
return self._session
async def _connect_once(self) -> None:
session = await self._ensure_session()
ws = await session.ws_connect(
self._config.worker_ws_url,
timeout=self._config.connect_timeout_seconds,
autoping=False,
heartbeat=None,
max_msg_size=self._config.max_frame_size,
)
old_ws = self._ws
self._ws = ws
if old_ws is not None and not old_ws.closed:
try:
await old_ws.close()
except Exception:
pass
if self._reader_task is not None:
self._reader_task.cancel()
if self._ping_task is not None:
self._ping_task.cancel()
self._reader_task = asyncio.create_task(self._reader_loop(ws))
self._ping_task = asyncio.create_task(self._ping_loop(ws))
logger.info("Hub worker channel connected: {}", self._config.worker_ws_url)
def _start_reconnect_loop(self) -> None:
if self._closing:
return
if self._reconnect_task is not None and not self._reconnect_task.done():
return
self._reconnect_task = asyncio.create_task(self._reconnect_loop())
async def _reconnect_loop(self) -> None:
attempt = 0
while not self._closing and not self.is_connected:
delay = _RECONNECT_DELAYS_SECONDS[min(attempt, len(_RECONNECT_DELAYS_SECONDS) - 1)]
if delay > 0:
await asyncio.sleep(delay)
try:
async with self._connect_lock:
if self._closing or self.is_connected:
break
await self._connect_once()
if self.is_connected:
logger.info("Hub worker channel reconnected")
break
except Exception as e:
attempt += 1
logger.debug("Hub reconnect attempt {} failed: {}", attempt, e)
async def _handle_disconnect(
self,
reason: str,
*,
ws: aiohttp.ClientWebSocketResponse | None = None,
) -> None:
current: aiohttp.ClientWebSocketResponse | None = None
async with self._connect_lock:
if self._ws is None:
return
if ws is not None and self._ws is not ws:
return
current = self._ws
self._ws = None
if current is not None and not current.closed:
try:
await current.close()
except Exception:
pass
if self._pending_streams:
for state in self._pending_streams.values():
state.set_error("hub disconnected")
self._pending_streams.clear()
if not self._closing:
logger.warning("Hub worker channel disconnected: {}", reason)
self._start_reconnect_loop()
async def _send_frame(self, frame: Frame) -> None:
ws = self._ws
if ws is None or ws.closed:
raise TunnelStreamError("hub not connected")
try:
async with asyncio.timeout(self._config.send_timeout_seconds):
async with self._write_lock:
await ws.send_bytes(frame.encode())
except TimeoutError as e:
await self._handle_disconnect("send timeout", ws=ws)
raise TunnelStreamError("hub frame send timeout") from e
except Exception as e:
await self._handle_disconnect(f"send failed: {e}", ws=ws)
raise TunnelStreamError(f"hub frame send failed: {e}") from e
async def _reader_loop(self, ws: aiohttp.ClientWebSocketResponse) -> None:
try:
while not self._closing:
msg = await ws.receive()
if msg.type == WSMsgType.BINARY:
raw = msg.data
if isinstance(raw, memoryview):
raw = raw.tobytes()
elif isinstance(raw, bytearray):
raw = bytes(raw)
if not isinstance(raw, bytes):
continue
try:
frame = Frame.decode(raw)
except Exception as e:
logger.debug("invalid frame from hub: {}", e)
continue
await self._handle_incoming_frame(frame)
continue
if msg.type == WSMsgType.CLOSE or msg.type == WSMsgType.CLOSED:
break
if msg.type == WSMsgType.ERROR:
logger.debug("hub ws reader error: {}", ws.exception())
break
if msg.type == WSMsgType.PING:
payload = msg.data if isinstance(msg.data, bytes) else b""
self._background(self._send_pong(payload))
continue
# TEXT / PONG / 其他类型直接忽略
except asyncio.CancelledError:
return
except Exception as e:
logger.debug("hub reader loop aborted: {}", e)
finally:
await self._handle_disconnect("reader ended", ws=ws)
async def _ping_loop(self, ws: aiohttp.ClientWebSocketResponse) -> None:
try:
while not self._closing:
await asyncio.sleep(self._config.ping_interval_seconds)
if self._ws is not ws or ws.closed:
break
try:
await self._send_frame(Frame(0, MsgType.PING, 0, b""))
except TunnelStreamError:
break
except asyncio.CancelledError:
return
async def _send_pong(self, payload: bytes) -> None:
try:
await self._send_frame(Frame(0, MsgType.PONG, 0, payload))
except TunnelStreamError:
pass
async def _handle_incoming_frame(self, frame: Frame) -> None:
match frame.msg_type:
# -- stream-level frames --
case MsgType.RESPONSE_HEADERS:
stream = self._pending_streams.get(frame.stream_id)
if not stream:
return
try:
payload = _decompress_frame_payload(frame)
meta = json.loads(payload)
stream.set_response_headers(meta["status"], meta.get("headers", []))
except Exception as e:
stream.set_error(f"invalid response headers: {e}")
self._pending_streams.pop(frame.stream_id, None)
case MsgType.RESPONSE_BODY:
stream = self._pending_streams.get(frame.stream_id)
if stream:
stream.push_body_chunk(_decompress_frame_payload(frame))
case MsgType.STREAM_END:
stream = self._pending_streams.pop(frame.stream_id, None)
if stream:
stream.set_done()
case MsgType.STREAM_ERROR:
stream = self._pending_streams.pop(frame.stream_id, None)
if stream:
message = (
frame.payload.decode(errors="replace") if frame.payload else "stream error"
)
stream.set_error(message)
# -- connection-level frames --
case MsgType.PING:
self._background(self._send_pong(frame.payload))
case MsgType.PONG:
pass
case MsgType.GOAWAY:
await self._handle_disconnect("received GOAWAY")
case MsgType.HEARTBEAT_DATA:
self._background(self._handle_heartbeat(frame))
case MsgType.HEARTBEAT_ACK:
pass
case MsgType.NODE_STATUS:
self._background(self._handle_node_status(frame.payload))
async def _handle_heartbeat(self, frame: Frame) -> None:
try:
data = json.loads(frame.payload) if frame.payload else {}
except Exception:
data = {}
node_id = str(data.get("node_id") or "").strip()
def _sync_heartbeat() -> dict[str, object]:
from src.database import create_session
from src.services.proxy_node.service import ProxyNodeService
if not node_id:
return {}
db = create_session()
try:
node = ProxyNodeService.heartbeat(
db,
node_id=node_id,
active_connections=data.get("active_connections"),
total_requests=data.get("total_requests"),
avg_latency_ms=data.get("avg_latency_ms"),
failed_requests=data.get("failed_requests"),
dns_failures=data.get("dns_failures"),
stream_errors=data.get("stream_errors"),
)
result: dict[str, object] = {}
if node.remote_config:
result["remote_config"] = node.remote_config
result["config_version"] = node.config_version or 0
return result
finally:
db.close()
try:
ack = await asyncio.to_thread(_sync_heartbeat)
except Exception as e:
logger.warning("hub heartbeat DB update failed: {}", e)
ack = {}
try:
await self._send_frame(
Frame(
frame.stream_id,
MsgType.HEARTBEAT_ACK,
0,
json.dumps(ack, ensure_ascii=False).encode("utf-8"),
)
)
except TunnelStreamError:
logger.debug("hub heartbeat ACK send failed")
async def _handle_node_status(self, payload: bytes) -> None:
try:
data = json.loads(payload) if payload else {}
except Exception:
return
node_id = str(data.get("node_id") or "").strip()
if not node_id:
return
connected = bool(data.get("connected"))
conn_count = int(data.get("conn_count") or 0)
# 所有 worker 都需要立即失效本地缓存,保证请求路由正确
try:
from src.services.proxy_node.resolver import invalidate_proxy_node_cache
invalidate_proxy_node_cache(node_id)
except Exception:
pass
# 使用 Redis SETNX 去重:同一次 NODE_STATUS 广播只有一个 worker 执行 DB 写入,
# 避免 N 个 worker 并发写同一行并产生 N 条重复事件记录。
dedup_key = f"hub:node_status:{node_id}:{connected}:{conn_count}"
try:
from src.clients import get_redis_client
redis = await get_redis_client()
if redis:
acquired = await redis.set(dedup_key, "1", ex=10, nx=True)
if not acquired:
return
except Exception:
# Redis 不可用时不去重,允许重复写入(幂等)
pass
def _sync_update() -> None:
from src.database import create_session
from src.models.database import ProxyNode, ProxyNodeEvent, ProxyNodeStatus
db = create_session()
try:
node = db.query(ProxyNode).filter(ProxyNode.id == node_id).first()
if not node:
return
now = datetime.now(timezone.utc)
node.tunnel_connected = connected
if connected:
node.tunnel_connected_at = now
node.status = ProxyNodeStatus.ONLINE if connected else ProxyNodeStatus.OFFLINE
node.updated_at = now
event = ProxyNodeEvent(
node_id=node_id,
event_type="connected" if connected else "disconnected",
detail=f"[hub_node_status] conn_count={conn_count}",
)
db.add(event)
db.commit()
except Exception:
try:
db.rollback()
except Exception:
pass
raise
finally:
db.close()
try:
await asyncio.to_thread(_sync_update)
except Exception as e:
logger.warning("hub NODE_STATUS DB update failed: node_id={}, error={}", node_id, e)
async def send_request(
self,
node_id: str,
*,
method: str,
url: str,
headers: dict[str, str],
body: bytes | None = None,
timeout: float = 60.0,
) -> _StreamState:
await self.ensure_connected()
if len(self._pending_streams) >= self._config.max_streams:
raise TunnelStreamError(
f"hub stream limit reached ({self._config.max_streams}) for node {node_id}"
)
stream_id = self._alloc_stream_id()
stream_state = _StreamState(stream_id)
self._pending_streams[stream_id] = stream_state
try:
meta = json.dumps(
{
"node_id": node_id,
"method": method,
"url": url,
"headers": headers,
"timeout": timeout,
},
ensure_ascii=False,
separators=(",", ":"),
).encode("utf-8")
meta_payload, meta_flags = _compress_frame_payload(meta)
await self._send_frame(
Frame(stream_id, MsgType.REQUEST_HEADERS, meta_flags, meta_payload)
)
body_data = body or b""
if body_data:
body_payload, body_flags = _compress_frame_payload(body_data)
else:
body_payload, body_flags = body_data, 0
body_flags |= FrameFlags.END_STREAM
await self._send_frame(Frame(stream_id, MsgType.REQUEST_BODY, body_flags, body_payload))
except Exception:
self._pending_streams.pop(stream_id, None)
raise
return stream_state
def remove_stream(self, stream_id: int) -> None:
self._pending_streams.pop(stream_id, None)
def _alloc_stream_id(self) -> int:
sid = self._next_stream_id
self._next_stream_id = sid + 2 if sid < 0xFFFF_FFFE else 2
return sid
async def shutdown(self) -> None:
self._closing = True
if self._reconnect_task is not None:
self._reconnect_task.cancel()
if self._reader_task is not None:
self._reader_task.cancel()
if self._ping_task is not None:
self._ping_task.cancel()
tasks = list(self._background_tasks)
for task in tasks:
task.cancel()
if tasks:
await asyncio.gather(*tasks, return_exceptions=True)
if self._ws is not None and not self._ws.closed:
try:
await self._ws.close()
except Exception:
pass
self._ws = None
if self._session is not None and not self._session.closed:
try:
await self._session.close()
except Exception:
pass
self._session = None
if self._pending_streams:
for state in self._pending_streams.values():
state.set_error("hub connection manager shutdown")
self._pending_streams.clear()
logger.info("Hub connection manager shutdown completed")
class HubTunnelTransport(httpx.AsyncBaseTransport):
"""通过 aether-hub 转发请求的 httpx transport。"""
def __init__(self, node_id: str, timeout: float = 60.0) -> None:
self._node_id = node_id
self._timeout = timeout
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
manager = get_hub_connection_manager()
headers: dict[str, str] = {}
for key, value in request.headers.raw:
if key not in _HOP_BY_HOP_HEADERS_BYTES:
headers[key.decode("latin-1")] = value.decode("latin-1")
body = request.content or await request.aread() or None
stream_state: _StreamState | None = None
try:
stream_state = await manager.send_request(
self._node_id,
method=request.method,
url=str(request.url),
headers=headers,
body=body,
timeout=self._timeout,
)
await stream_state.wait_headers(timeout=self._timeout)
return httpx.Response(
status_code=stream_state.status,
headers=httpx.Headers(stream_state.headers),
stream=HubResponseStream(manager, stream_state, timeout=self._timeout),
)
except TunnelStreamError as e:
self._cleanup_stream(manager, stream_state)
if stream_state and stream_state.status > 0:
raise httpx.ReadError(str(e)) from e
raise httpx.ConnectError(str(e)) from e
except asyncio.TimeoutError:
self._cleanup_stream(manager, stream_state)
raise httpx.ReadTimeout("hub tunnel request timeout") from None
def _cleanup_stream(
self,
manager: HubConnectionManager,
stream_state: _StreamState | None,
) -> None:
if stream_state is None:
return
manager.remove_stream(stream_state.stream_id)
class HubResponseStream(httpx.AsyncByteStream):
def __init__(
self,
manager: HubConnectionManager,
stream_state: _StreamState,
timeout: float = 60.0,
) -> None:
self._manager = manager
self._stream_state = stream_state
self._timeout = timeout
async def __aiter__(self) -> AsyncGenerator[bytes, None]:
async for chunk in self._stream_state.iter_body(chunk_timeout=self._timeout):
yield chunk
async def aclose(self) -> None:
self._manager.remove_stream(self._stream_state.stream_id)
def _compress_frame_payload(data: bytes) -> tuple[bytes, int]:
if len(data) >= _TUNNEL_COMPRESS_MIN_SIZE:
compressed = gzip.compress(data, compresslevel=6)
if len(compressed) < len(data):
return compressed, FrameFlags.GZIP_COMPRESSED
return data, 0
def _decompress_frame_payload(frame: Frame) -> bytes:
if frame.is_gzip:
return gzip.decompress(frame.payload)
return frame.payload
_hub_connection_manager: HubConnectionManager | None = None
def get_hub_connection_manager() -> HubConnectionManager:
global _hub_connection_manager
if _hub_connection_manager is None:
_hub_connection_manager = HubConnectionManager()
return _hub_connection_manager
async def shutdown_hub_connection_manager() -> None:
global _hub_connection_manager
if _hub_connection_manager is None:
return
await _hub_connection_manager.shutdown()
_hub_connection_manager = None

View File

@@ -33,6 +33,51 @@ _TUNNEL_LOCAL_MISS_LOG_COOLDOWN_SECONDS = 30.0
_tunnel_local_miss_log_next_at: dict[str, float] = {} _tunnel_local_miss_log_next_at: dict[str, float] = {}
def _is_hub_mode_enabled() -> bool:
from src.services.proxy_node.hub_config import get_hub_config
return get_hub_config().enabled
def _build_tunnel_local_miss_message(node_id: str) -> str:
"""构建“当前 worker 无本地 tunnel”的用户可读错误信息。"""
return (
f"代理节点 {node_id} 当前 worker 无 tunnel 连接pid={os.getpid()})。"
"这通常发生在多 worker 部署WEB_CONCURRENCY/GUNICORN_WORKERS > 1时。"
"请设置 GUNICORN_WORKERS=1或 WEB_CONCURRENCY=1"
"或确保每个 worker 都建立该节点的 tunnel 连接。"
)
def _is_tunnel_local_miss(node_id: str) -> bool:
"""判断节点是否“全局在线但当前 worker 无本地 tunnel 连接”。
仅在错误路径调用build_proxy_url 解析失败时),用于提供更准确的报错信息。
"""
if _is_hub_mode_enabled():
return False
from src.database import create_session
from src.models.database import ProxyNode, ProxyNodeStatus
from src.services.proxy_node.tunnel_manager import get_tunnel_manager
db = create_session()
try:
node = db.query(ProxyNode).filter(ProxyNode.id == node_id).first()
if not node:
return False
if node.is_manual or not node.tunnel_mode:
return False
if node.status != ProxyNodeStatus.ONLINE:
return False
manager = get_tunnel_manager()
return not manager.has_tunnel(node_id)
except Exception:
return False
finally:
db.close()
def _get_proxy_node_info(node_id: str) -> dict[str, Any] | None: def _get_proxy_node_info(node_id: str) -> dict[str, Any] | None:
""" """
读取 ProxyNode 信息(带内存 TTL 缓存) 读取 ProxyNode 信息(带内存 TTL 缓存)
@@ -71,32 +116,41 @@ def _get_proxy_node_info(node_id: str) -> dict[str, Any] | None:
_proxy_node_cache[node_id] = (None, now + _PROXY_NODE_CACHE_NEGATIVE_TTL_SECONDS) _proxy_node_cache[node_id] = (None, now + _PROXY_NODE_CACHE_NEGATIVE_TTL_SECONDS)
return None return None
# tunnel 模式节点:以 TunnelManager 内存中的实际连接状态为准, # tunnel 模式节点:
# 而非依赖 DB 的 status/tunnel_connected 字段(可能因竞态不同步)。 # - Hub 模式:信任 DB 的 tunnel_connected/status由 Hub 广播统一维护)
# - 非 Hub 模式:以当前 worker 本地 TunnelManager 状态为准
if node.tunnel_mode and not node.is_manual: if node.tunnel_mode and not node.is_manual:
from src.services.proxy_node.tunnel_manager import get_tunnel_manager if _is_hub_mode_enabled():
if node.status != ProxyNodeStatus.ONLINE or not bool(node.tunnel_connected):
_proxy_node_cache[node_id] = (
None,
now + _PROXY_NODE_CACHE_NEGATIVE_TTL_SECONDS,
)
return None
else:
from src.services.proxy_node.tunnel_manager import get_tunnel_manager
manager = get_tunnel_manager() manager = get_tunnel_manager()
if not manager.has_tunnel(node_id): if not manager.has_tunnel(node_id):
# 多 worker 部署时DB 可能显示 ONLINE其他 worker 有 tunnel # 多 worker 部署时DB 可能显示 ONLINE其他 worker 有 tunnel
# 但当前 worker 无本地连接,请求仍不可用。记录限频告警便于定位。 # 但当前 worker 无本地连接,请求仍不可用。记录限频告警便于定位。
if now >= _tunnel_local_miss_log_next_at.get(node_id, 0.0): if now >= _tunnel_local_miss_log_next_at.get(node_id, 0.0):
_tunnel_local_miss_log_next_at[node_id] = ( _tunnel_local_miss_log_next_at[node_id] = (
now + _TUNNEL_LOCAL_MISS_LOG_COOLDOWN_SECONDS now + _TUNNEL_LOCAL_MISS_LOG_COOLDOWN_SECONDS
)
logger.warning(
"tunnel node {} has no local connection on pid={} "
"(db_status={}, db_tunnel_connected={}), request may fail on this worker",
node_id,
os.getpid(),
str(getattr(node, "status", "unknown")),
bool(getattr(node, "tunnel_connected", False)),
)
_proxy_node_cache[node_id] = (
None,
now + _PROXY_NODE_CACHE_TUNNEL_LOCAL_MISS_TTL_SECONDS,
) )
logger.warning( return None
"tunnel node {} has no local connection on pid={} "
"(db_status={}, db_tunnel_connected={}), request may fail on this worker",
node_id,
os.getpid(),
str(getattr(node, "status", "unknown")),
bool(getattr(node, "tunnel_connected", False)),
)
_proxy_node_cache[node_id] = (
None,
now + _PROXY_NODE_CACHE_TUNNEL_LOCAL_MISS_TTL_SECONDS,
)
return None
value: dict[str, Any] = { value: dict[str, Any] = {
"name": node.name, "name": node.name,
"ip": node.ip, "ip": node.ip,
@@ -408,13 +462,13 @@ def build_proxy_client_kwargs(
kwargs: dict[str, Any] = {"timeout": timeout, "verify": verify, **extra} kwargs: dict[str, Any] = {"timeout": timeout, "verify": verify, **extra}
# tunnel 模式优先:当代理节点为 tunnel 模式时,使用 TunnelTransport # tunnel 模式优先:当代理节点为 tunnel 模式时,使用 tunnel transport 工厂
delegate_cfg = resolve_delegate_config(proxy_config) delegate_cfg = resolve_delegate_config(proxy_config)
if delegate_cfg and delegate_cfg.get("tunnel"): if delegate_cfg and delegate_cfg.get("tunnel"):
from src.services.proxy_node.tunnel_transport import TunnelTransport from src.services.proxy_node.tunnel_transport import create_tunnel_transport
timeout_secs = timeout if isinstance(timeout, (int, float)) else 60.0 timeout_secs = timeout if isinstance(timeout, (int, float)) else 60.0
kwargs["transport"] = TunnelTransport(delegate_cfg["node_id"], timeout=timeout_secs) kwargs["transport"] = create_tunnel_transport(delegate_cfg["node_id"], timeout=timeout_secs)
return kwargs return kwargs
proxy_param = resolve_proxy_param(proxy_config) proxy_param = resolve_proxy_param(proxy_config)
@@ -455,7 +509,12 @@ def build_proxy_url(proxy_config: dict[str, Any]) -> str | None:
node_info = _get_proxy_node_info(node_id) node_info = _get_proxy_node_info(node_id)
if not node_info: if not node_info:
logger.warning("代理节点不可用(离线或不存在): node_id={}", node_id) logger.warning("代理节点不可用(离线或不存在): node_id={}", node_id)
raise ProxyNodeUnavailableError(f"代理节点 {node_id} 不可用", node_id=node_id) message = (
_build_tunnel_local_miss_message(node_id)
if _is_tunnel_local_miss(node_id)
else f"代理节点 {node_id} 不可用"
)
raise ProxyNodeUnavailableError(message, node_id=node_id)
# 手动节点:直接使用存储的代理 URL含认证信息 # 手动节点:直接使用存储的代理 URL含认证信息
if node_info.get("is_manual"): if node_info.get("is_manual"):

View File

@@ -177,10 +177,10 @@ async def _test_tunnel_connectivity(node_id: str) -> dict[str, Any]:
"""通过 WebSocket tunnel 测试连通性,返回标准化结果 dict""" """通过 WebSocket tunnel 测试连通性,返回标准化结果 dict"""
import time as _time import time as _time
from .tunnel_transport import TunnelTransport from .tunnel_transport import create_tunnel_transport
test_url = "https://1.1.1.1/cdn-cgi/trace" test_url = "https://1.1.1.1/cdn-cgi/trace"
transport = TunnelTransport(node_id, timeout=15.0) transport = create_tunnel_transport(node_id, timeout=15.0)
start = _time.monotonic() start = _time.monotonic()
try: try:
@@ -556,17 +556,38 @@ class ProxyNodeService:
# tunnel 节点:通过 WebSocket tunnel 测试 # tunnel 节点:通过 WebSocket tunnel 测试
if not node.is_manual: if not node.is_manual:
# 以 TunnelManager 内存中的实际连接状态为准(与 health_scheduler 一致), from src.services.proxy_node.hub_config import get_hub_config
# 而非仅依赖 DB 的 tunnel_connected 字段,避免竞态导致误判。
from src.services.proxy_node.tunnel_manager import get_tunnel_manager
manager = get_tunnel_manager() hub_enabled = get_hub_config().enabled
if not manager.has_tunnel(node.id): connected = False
if hub_enabled:
connected = bool(node.tunnel_connected) and node.status == ProxyNodeStatus.ONLINE
else:
# 非 Hub 模式:以当前 worker 本地 TunnelManager 状态为准,
# 避免多 worker 场景的跨进程状态误判。
from src.services.proxy_node.tunnel_manager import get_tunnel_manager
manager = get_tunnel_manager()
connected = manager.has_tunnel(node.id)
if not connected:
hint = ""
try:
from src.config import config
if config.worker_processes > 1 and not hub_enabled:
hint = (
"(当前 worker 无 tunnel 连接;检测到多 worker 部署,"
"建议设置 GUNICORN_WORKERS/WEB_CONCURRENCY=1"
)
except Exception:
# 配置读取失败时保持原始错误,避免影响主流程
hint = ""
return { return {
"success": False, "success": False,
"latency_ms": None, "latency_ms": None,
"exit_ip": None, "exit_ip": None,
"error": "tunnel 未连接", "error": f"tunnel 未连接{hint}",
} }
result = await _test_tunnel_connectivity(node.id) result = await _test_tunnel_connectivity(node.id)

View File

@@ -29,6 +29,7 @@ class MsgType(IntEnum):
GOAWAY = 0x12 # \u53cc\u5411: \u4f18\u96c5\u5173\u95ed (stream_id=0) GOAWAY = 0x12 # \u53cc\u5411: \u4f18\u96c5\u5173\u95ed (stream_id=0)
HEARTBEAT_DATA = 0x13 # Proxy -> Aether: \u6307\u6807\u4e0a\u62a5 HEARTBEAT_DATA = 0x13 # Proxy -> Aether: \u6307\u6807\u4e0a\u62a5
HEARTBEAT_ACK = 0x14 # Aether -> Proxy: \u5fc3\u8df3\u786e\u8ba4 + \u8fdc\u7a0b\u914d\u7f6e HEARTBEAT_ACK = 0x14 # Aether -> Proxy: \u5fc3\u8df3\u786e\u8ba4 + \u8fdc\u7a0b\u914d\u7f6e
NODE_STATUS = 0x15 # Hub -> Worker: \u8282\u70b9\u8fde\u63a5\u72b6\u6001\u5e7f\u64ad
class FrameFlags: class FrameFlags:

View File

@@ -136,3 +136,15 @@ def is_tunnel_node(node_info: dict[str, Any] | None) -> bool:
if not node_info: if not node_info:
return False return False
return bool(node_info.get("tunnel_mode")) and bool(node_info.get("tunnel_connected")) return bool(node_info.get("tunnel_mode")) and bool(node_info.get("tunnel_connected"))
def create_tunnel_transport(node_id: str, timeout: float = 60.0) -> httpx.AsyncBaseTransport:
"""根据配置创建 tunnel transportHub 模式或直连 tunnel 模式)。"""
from .hub_config import get_hub_config
hub_cfg = get_hub_config()
if hub_cfg.enabled:
from .hub_transport import HubTunnelTransport
return HubTunnelTransport(node_id, timeout=timeout)
return TunnelTransport(node_id, timeout=timeout)

View File

@@ -0,0 +1,34 @@
import pytest
from src.services.proxy_node.hub_config import reset_hub_config_cache
from src.services.proxy_node.hub_transport import HubTunnelTransport
from src.services.proxy_node.tunnel_protocol import Frame, MsgType
from src.services.proxy_node.tunnel_transport import TunnelTransport, create_tunnel_transport
def test_create_tunnel_transport_uses_legacy_transport_when_hub_disabled(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setenv("DOCKER_CONTAINER", "false")
monkeypatch.setattr("src.services.proxy_node.hub_config.os.path.exists", lambda _: False)
reset_hub_config_cache()
transport = create_tunnel_transport("node-1", timeout=12.0)
assert isinstance(transport, TunnelTransport)
def test_create_tunnel_transport_uses_hub_transport_when_hub_enabled(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setenv("DOCKER_CONTAINER", "true")
monkeypatch.setattr("src.services.proxy_node.hub_config.os.path.exists", lambda _: False)
reset_hub_config_cache()
transport = create_tunnel_transport("node-1", timeout=12.0)
assert isinstance(transport, HubTunnelTransport)
def test_tunnel_protocol_supports_node_status_msg_type() -> None:
raw = Frame(0, MsgType.NODE_STATUS, 0, b'{"node_id":"n1","connected":true}').encode()
decoded = Frame.decode(raw)
assert decoded.msg_type == MsgType.NODE_STATUS