mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 01:40:21 +08:00
refactor(proxy): 将 aether-proxy 从 HMAC 正向代理迁移到 WebSocket 隧道模式
移除 HMAC 认证、TLS 自签名证书、HTTP CONNECT 代理和代发(delegate)模式, 改为 aether-proxy 主动通过 WebSocket 连接 Aether 服务端建立隧道。 Aether 服务端新增: - WebSocket 隧道端点 (proxy_tunnel.py) - TunnelManager 管理隧道连接和请求分发 - TunnelTransport 作为 httpx 自定义 transport 层 - 基于二进制帧的隧道协议 (tunnel_protocol.py) aether-proxy (Rust) 重构: - 新增 tunnel 模块 (client/dispatcher/stream_handler/protocol) - 支持多 Aether 服务端连接 ([[servers]] 配置) - 移除 proxy/auth/delegate 模块和 hyper 依赖 - 改用 tokio-tungstenite 实现 WebSocket 客户端 同时: - 添加浏览器指纹 Headers 绕过 Cloudflare 防护 - 删除节点时自动清理 Provider/Endpoint 的代理引用 - 数据库迁移: 新增 tunnel_mode/tunnel_connected/tunnel_connected_at 字段
This commit is contained in:
@@ -32,7 +32,7 @@ pipeline = ApiRequestPipeline()
|
||||
class ProxyNodeRegisterRequest(BaseModel):
|
||||
name: str = Field(..., min_length=1, max_length=100, description="节点名")
|
||||
ip: str = Field(..., description="公网 IP(IPv4/IPv6)")
|
||||
port: int = Field(..., ge=1, le=65535, description="代理端口")
|
||||
port: int = Field(0, ge=0, le=65535, description="代理端口(tunnel 模式下为 0)")
|
||||
region: str | None = Field(None, max_length=100, description="区域标签")
|
||||
heartbeat_interval: int = Field(30, ge=5, le=600, description="心跳间隔(秒)")
|
||||
|
||||
@@ -41,16 +41,13 @@ class ProxyNodeRegisterRequest(BaseModel):
|
||||
total_requests: int | None = Field(None, ge=0, description="累计请求数")
|
||||
avg_latency_ms: float | None = Field(None, ge=0, description="平均延迟(毫秒)")
|
||||
|
||||
# TLS
|
||||
tls_enabled: bool = Field(False, description="是否启用 TLS 加密")
|
||||
tls_cert_fingerprint: str | None = Field(
|
||||
None, max_length=128, description="TLS 证书 SHA-256 指纹"
|
||||
)
|
||||
|
||||
# 硬件信息
|
||||
hardware_info: dict | None = Field(None, description="硬件信息 JSON")
|
||||
estimated_max_concurrency: int | None = Field(None, ge=0, description="估算最大并发连接数")
|
||||
|
||||
# Tunnel 模式
|
||||
tunnel_mode: bool = Field(False, description="是否使用 tunnel 模式连接")
|
||||
|
||||
@field_validator("ip")
|
||||
@classmethod
|
||||
def validate_ip(cls, v: str) -> str:
|
||||
@@ -82,9 +79,6 @@ class ProxyNodeRemoteConfigRequest(BaseModel):
|
||||
allowed_ports: list[int] | None = Field(None, description="允许代理的目标端口")
|
||||
log_level: str | None = Field(None, description="日志级别 (trace/debug/info/warn/error)")
|
||||
heartbeat_interval: int | None = Field(None, ge=5, le=600, description="心跳间隔(秒)")
|
||||
timestamp_tolerance: int | None = Field(
|
||||
None, ge=10, le=3600, description="HMAC 时间戳容差(秒)"
|
||||
)
|
||||
|
||||
@field_validator("allowed_ports")
|
||||
@classmethod
|
||||
@@ -216,12 +210,6 @@ async def test_proxy_node(node_id: str, request: Request, db: Session = Depends(
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.get("/hmac-key")
|
||||
async def get_proxy_hmac_key(request: Request, db: Session = Depends(get_db)) -> Any:
|
||||
adapter = AdminGetProxyHmacKeyAdapter()
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.post("/test-url")
|
||||
async def test_proxy_url(request: Request, db: Session = Depends(get_db)) -> Any:
|
||||
adapter = AdminTestProxyUrlAdapter()
|
||||
@@ -273,14 +261,13 @@ class AdminRegisterProxyNodeAdapter(AdminApiAdapter):
|
||||
port=req.port,
|
||||
region=req.region,
|
||||
heartbeat_interval=req.heartbeat_interval,
|
||||
tls_enabled=req.tls_enabled,
|
||||
tls_cert_fingerprint=req.tls_cert_fingerprint,
|
||||
hardware_info=req.hardware_info,
|
||||
estimated_max_concurrency=req.estimated_max_concurrency,
|
||||
active_connections=req.active_connections,
|
||||
total_requests=req.total_requests,
|
||||
avg_latency_ms=req.avg_latency_ms,
|
||||
registered_by=context.user.id if context.user else None,
|
||||
tunnel_mode=req.tunnel_mode,
|
||||
)
|
||||
|
||||
context.add_audit_metadata(
|
||||
@@ -376,11 +363,24 @@ class AdminDeleteProxyNodeAdapter(AdminApiAdapter):
|
||||
)
|
||||
|
||||
was_system_proxy = result["cleared_system_proxy"]
|
||||
msg = "deleted, system default proxy cleared" if was_system_proxy else "deleted"
|
||||
cleared_providers = result.get("cleared_providers", 0)
|
||||
cleared_endpoints = result.get("cleared_endpoints", 0)
|
||||
|
||||
parts = ["deleted"]
|
||||
if was_system_proxy:
|
||||
parts.append("system default proxy cleared")
|
||||
if cleared_providers or cleared_endpoints:
|
||||
parts.append(
|
||||
f"cleared proxy from {cleared_providers} provider(s) "
|
||||
f"and {cleared_endpoints} endpoint(s)"
|
||||
)
|
||||
|
||||
return {
|
||||
"message": msg,
|
||||
"message": ", ".join(parts),
|
||||
"node_id": self.node_id,
|
||||
"cleared_system_proxy": was_system_proxy,
|
||||
"cleared_providers": cleared_providers,
|
||||
"cleared_endpoints": cleared_endpoints,
|
||||
}
|
||||
|
||||
|
||||
@@ -478,8 +478,6 @@ class AdminUpdateProxyNodeConfigAdapter(AdminApiAdapter):
|
||||
config_updates["log_level"] = req.log_level
|
||||
if req.heartbeat_interval is not None:
|
||||
config_updates["heartbeat_interval"] = req.heartbeat_interval
|
||||
if req.timestamp_tolerance is not None:
|
||||
config_updates["timestamp_tolerance"] = req.timestamp_tolerance
|
||||
|
||||
node = ProxyNodeService.update_node_config(
|
||||
context.db, node_id=self.node_id, config_updates=config_updates
|
||||
@@ -499,23 +497,6 @@ class AdminUpdateProxyNodeConfigAdapter(AdminApiAdapter):
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdminGetProxyHmacKeyAdapter(AdminApiAdapter):
|
||||
"""获取 proxy_hmac_key 供管理员复制到 aether-proxy 部署"""
|
||||
|
||||
name: str = "admin_get_proxy_hmac_key"
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any:
|
||||
from src.config.settings import config
|
||||
|
||||
key = config.proxy_hmac_key
|
||||
if not key:
|
||||
raise InvalidRequestException(
|
||||
"PROXY_HMAC_KEY 未配置(也未设置 ENCRYPTION_KEY 用于自动派生)"
|
||||
)
|
||||
return {"proxy_hmac_key": key}
|
||||
|
||||
|
||||
class TestProxyUrlRequest(BaseModel):
|
||||
proxy_url: str = Field(..., min_length=1, max_length=500)
|
||||
username: str | None = Field(None, max_length=255)
|
||||
|
||||
171
src/api/admin/proxy_tunnel.py
Normal file
171
src/api/admin/proxy_tunnel.py
Normal file
@@ -0,0 +1,171 @@
|
||||
"""
|
||||
WebSocket 隧道端点
|
||||
|
||||
aether-proxy 通过此端点建立 tunnel 连接。
|
||||
路径: /api/internal/proxy-tunnel
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
|
||||
from fastapi import APIRouter, WebSocket, WebSocketDisconnect
|
||||
|
||||
from src.core.logger import logger
|
||||
from src.services.proxy_node.tunnel_manager import (
|
||||
TunnelConnection,
|
||||
get_tunnel_manager,
|
||||
)
|
||||
from src.services.proxy_node.tunnel_protocol import Frame
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
# 单帧最大 64 MB -- AI API 请求体可能包含多张 base64 图片,需要足够余量
|
||||
_MAX_FRAME_SIZE = 64 * 1024 * 1024
|
||||
|
||||
# WebSocket 空闲超时(秒)-- proxy 端 ping 间隔默认 15s,3 倍余量
|
||||
_IDLE_TIMEOUT = 90.0
|
||||
|
||||
|
||||
async def _authenticate(ws: WebSocket) -> tuple[str, str] | None:
|
||||
"""验证 WebSocket 连接的认证信息,返回 (node_id, node_name) 或 None
|
||||
|
||||
认证方式:Bearer <management_token>,通过 Management Token 系统验证。
|
||||
authenticate_management_token 是 async 方法(内部有 Redis 速率限制),
|
||||
因此直接 await 调用。节点存在性检查复用同一 session。
|
||||
"""
|
||||
auth = ws.headers.get("authorization", "")
|
||||
if not auth.startswith("Bearer "):
|
||||
return None
|
||||
|
||||
token = auth[7:]
|
||||
if not token or not token.startswith("ae_"):
|
||||
return None
|
||||
|
||||
client_ip = getattr(ws.client, "host", "unknown") if ws.client else "unknown"
|
||||
node_id_header = ws.headers.get("x-node-id", "").strip()
|
||||
node_name_header = ws.headers.get("x-node-name", "").strip()
|
||||
|
||||
if not node_id_header:
|
||||
return None
|
||||
|
||||
from src.database import create_session
|
||||
from src.models.database import ProxyNode
|
||||
from src.services.auth.service import AuthService
|
||||
|
||||
db = create_session()
|
||||
try:
|
||||
result = await AuthService.authenticate_management_token(db, token, client_ip)
|
||||
if not result:
|
||||
return None
|
||||
|
||||
# 节点存在性检查(复用同一 session,避免额外连接开销)
|
||||
exists = db.query(
|
||||
db.query(ProxyNode).filter(ProxyNode.id == node_id_header).exists()
|
||||
).scalar()
|
||||
if not exists:
|
||||
logger.warning("tunnel auth: node_id={} not found in DB", node_id_header)
|
||||
return None
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
return node_id_header, node_name_header or node_id_header
|
||||
|
||||
|
||||
@router.websocket("/api/internal/proxy-tunnel")
|
||||
async def proxy_tunnel_ws(ws: WebSocket) -> None:
|
||||
"""aether-proxy tunnel WebSocket 端点"""
|
||||
try:
|
||||
auth = await _authenticate(ws)
|
||||
except Exception as e:
|
||||
logger.warning("tunnel auth error: {}", e)
|
||||
await ws.accept()
|
||||
await ws.close(code=4002, reason="authentication error")
|
||||
return
|
||||
|
||||
if not auth:
|
||||
await ws.accept()
|
||||
await ws.close(code=4001, reason="unauthorized")
|
||||
return
|
||||
|
||||
node_id: str = auth[0]
|
||||
node_name: str = auth[1]
|
||||
await ws.accept()
|
||||
|
||||
manager = get_tunnel_manager()
|
||||
conn = TunnelConnection(node_id, node_name, ws)
|
||||
manager.register(conn)
|
||||
|
||||
# 更新 DB: tunnel_connected = True
|
||||
await _update_tunnel_status(node_id, connected=True)
|
||||
|
||||
try:
|
||||
oversized_count = 0
|
||||
while True:
|
||||
try:
|
||||
data = await asyncio.wait_for(ws.receive_bytes(), timeout=_IDLE_TIMEOUT)
|
||||
except asyncio.TimeoutError:
|
||||
logger.warning("tunnel idle timeout for node_id={}", node_id)
|
||||
await ws.close(code=4004, reason="idle timeout")
|
||||
break
|
||||
if len(data) > _MAX_FRAME_SIZE:
|
||||
oversized_count += 1
|
||||
logger.warning("tunnel frame too large from {}: {} bytes", node_id, len(data))
|
||||
if oversized_count >= 5:
|
||||
logger.warning("too many oversized frames from {}, closing", node_id)
|
||||
await ws.close(code=4003, reason="too many oversized frames")
|
||||
break
|
||||
continue
|
||||
oversized_count = 0 # 正常帧重置计数
|
||||
try:
|
||||
frame = Frame.decode(data)
|
||||
except ValueError as e:
|
||||
logger.warning("tunnel frame decode error from {}: {}", node_id, e)
|
||||
continue
|
||||
|
||||
await manager.handle_incoming_frame(node_id, frame)
|
||||
|
||||
except WebSocketDisconnect:
|
||||
logger.info("tunnel WebSocket disconnected: node_id={}", node_id)
|
||||
except Exception as e:
|
||||
logger.error("tunnel WebSocket error for node_id={}: {}", node_id, e)
|
||||
finally:
|
||||
manager.unregister(node_id)
|
||||
await _update_tunnel_status(node_id, connected=False)
|
||||
|
||||
|
||||
async def _update_tunnel_status(node_id: str, *, connected: bool) -> None:
|
||||
"""更新 ProxyNode 的 tunnel 连接状态(在线程池中执行,避免阻塞 event loop)"""
|
||||
|
||||
def _sync_update() -> None:
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from src.database import create_session
|
||||
from src.models.database import ProxyNode, ProxyNodeStatus
|
||||
|
||||
db = create_session()
|
||||
try:
|
||||
node = db.query(ProxyNode).filter(ProxyNode.id == node_id).first()
|
||||
if node:
|
||||
node.tunnel_connected = connected
|
||||
now = datetime.now(timezone.utc)
|
||||
if connected:
|
||||
node.tunnel_connected_at = now
|
||||
node.status = ProxyNodeStatus.ONLINE
|
||||
else:
|
||||
# 记录断开时刻,供 health_scheduler 计算 UNHEALTHY 缓冲期
|
||||
node.tunnel_connected_at = now
|
||||
node.status = ProxyNodeStatus.UNHEALTHY
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
try:
|
||||
await asyncio.to_thread(_sync_update)
|
||||
except Exception as e:
|
||||
logger.warning("failed to update tunnel status for {}: {}", node_id, e)
|
||||
|
||||
# 清除节点信息缓存,确保后续请求能立即感知连接状态变化
|
||||
from src.services.proxy_node.resolver import invalidate_proxy_node_cache
|
||||
|
||||
invalidate_proxy_node_cache(node_id)
|
||||
Reference in New Issue
Block a user