mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 01:40:21 +08:00
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:
12
.env.example
12
.env.example
@@ -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 处理指定数量请求后自动重启,防止内存泄漏
|
||||||
|
|||||||
59
.github/workflows/docker-publish.yml
vendored
59
.github/workflows/docker-publish.yml
vendored
@@ -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
|
||||||
|
|||||||
@@ -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 \
|
||||||
|
|||||||
@@ -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
3
aether-hub/.dockerignore
Normal file
@@ -0,0 +1,3 @@
|
|||||||
|
target/
|
||||||
|
.git/
|
||||||
|
.DS_Store
|
||||||
1229
aether-hub/Cargo.lock
generated
Normal file
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
23
aether-hub/Cargo.toml
Normal 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
27
aether-hub/Dockerfile
Normal 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
695
aether-hub/src/hub.rs
Normal 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
159
aether-hub/src/main.rs
Normal 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
232
aether-hub/src/protocol.rs
Normal 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))
|
||||||
|
}
|
||||||
|
}
|
||||||
144
aether-hub/src/proxy_conn.rs
Normal file
144
aether-hub/src/proxy_conn.rs
Normal 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
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
130
aether-hub/src/worker_conn.rs
Normal file
130
aether-hub/src/worker_conn.rs
Normal 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
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
84
deploy.sh
84
deploy.sh
@@ -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
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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}")
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
70
src/services/proxy_node/hub_config.py
Normal file
70
src/services/proxy_node/hub_config.py
Normal 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
|
||||||
639
src/services/proxy_node/hub_transport.py
Normal file
639
src/services/proxy_node/hub_transport.py
Normal 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
|
||||||
@@ -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"):
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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 transport(Hub 模式或直连 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)
|
||||||
|
|||||||
34
tests/unit/test_tunnel_hub_factory.py
Normal file
34
tests/unit/test_tunnel_hub_factory.py
Normal 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
|
||||||
Reference in New Issue
Block a user