refactor(tunnel): rename aether-proxy to aether-tunnel

This commit is contained in:
fawney19
2026-05-20 01:02:01 +08:00
parent f4d0d5904a
commit 94760dbc14
57 changed files with 939 additions and 889 deletions

View File

@@ -0,0 +1,19 @@
# Aether server URL
AETHER_TUNNEL_AETHER_URL=https://aether.example.com
# Management Token (ae_xxx, must belong to an ADMIN user)
AETHER_TUNNEL_MANAGEMENT_TOKEN=ae_xxxxx
# Node identification
AETHER_TUNNEL_NODE_NAME=jp-proxy-01
# Maximum request body buffered for 307/308 replay (supports K/M/G, 0 disables body replay buffering)
AETHER_TUNNEL_REDIRECT_REPLAY_BUDGET_BYTES=5M
# Logging
AETHER_TUNNEL_LOG_LEVEL=info
AETHER_TUNNEL_LOG_DESTINATION=stdout
AETHER_TUNNEL_LOG_DIR=/var/log/aether-tunnel
AETHER_TUNNEL_LOG_ROTATION=daily
AETHER_TUNNEL_LOG_RETENTION_DAYS=7
AETHER_TUNNEL_LOG_MAX_FILES=30

3403
apps/aether-tunnel/Cargo.lock generated Normal file

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,46 @@
[package]
name = "aether-tunnel"
version = "0.3.12"
edition = "2021"
description = "Tunnel agent for Aether"
[dependencies]
aether-contracts.workspace = true
aether-http.workspace = true
aether-runtime.workspace = true
aether-runtime-state.workspace = true
axum.workspace = true
tokio = { version = "1", features = ["full"] }
reqwest.workspace = true
hyper = { version = "1", features = ["client", "http1", "http2"] }
hyper-util = { version = "0.1", features = ["client", "client-legacy", "http1", "http2", "tokio"] }
http-body-util = "0.1"
tokio-tungstenite = { version = "0.24", features = ["rustls-tls-webpki-roots"] }
tokio-rustls = "0.26"
futures-util = "0.3"
base64 = "0.22"
clap = { version = "4", features = ["derive", "env"] }
tracing = "0.1"
serde = { version = "1", features = ["derive"] }
serde_json.workspace = true
thiserror = "2"
bytes = "1"
sha2 = "0.10"
hex = "0.4"
anyhow = "1"
arc-swap = "1"
toml = "0.8"
rustls = { version = "0.23", features = ["ring"] }
ratatui = "0.30"
crossterm = "0.28"
url = "2"
sysinfo = "0.32"
libc = "0.2"
flate2 = "1"
tar = "0.4"
socket2 = { version = "0.5", features = ["all"] }
tower-service = "0.3"
webpki-roots = "0.26"
[dev-dependencies]
aether-gateway.workspace = true

View File

@@ -0,0 +1,231 @@
# aether-tunnel
Aether Tunnel 代理节点,部署在海外 VPS 上,通过 WebSocket 隧道为 Aether 实例中转 API 流量。
Tunnel 模式下代理节点**无需对外监听端口**,仅需出站连接到 Aether 服务器。
## 安装
`aether-tunnel` 会根据宿主机自动选择服务管理器:
- 常规 Linux 发行版:`systemd`
- Alpine Linux`OpenRC`
### 下载预编译二进制
<!-- DOWNLOAD_TABLE_START -->
| Platform | Download |
|----------|----------|
| Linux x86_64 (GNU) | [aether-tunnel-linux-amd64.tar.gz](https://github.com/fawney19/Aether/releases/download/tunnel-v0.3.12/aether-tunnel-linux-amd64.tar.gz) |
| Linux ARM64 (GNU) | [aether-tunnel-linux-arm64.tar.gz](https://github.com/fawney19/Aether/releases/download/tunnel-v0.3.12/aether-tunnel-linux-arm64.tar.gz) |
| Linux x86_64 (musl) | [aether-tunnel-linux-musl-amd64.tar.gz](https://github.com/fawney19/Aether/releases/download/tunnel-v0.3.12/aether-tunnel-linux-musl-amd64.tar.gz) |
| Linux ARM64 (musl) | [aether-tunnel-linux-musl-arm64.tar.gz](https://github.com/fawney19/Aether/releases/download/tunnel-v0.3.12/aether-tunnel-linux-musl-arm64.tar.gz) |
| macOS x86_64 | [aether-tunnel-macos-amd64.tar.gz](https://github.com/fawney19/Aether/releases/download/tunnel-v0.3.12/aether-tunnel-macos-amd64.tar.gz) |
| macOS ARM64 | [aether-tunnel-macos-arm64.tar.gz](https://github.com/fawney19/Aether/releases/download/tunnel-v0.3.12/aether-tunnel-macos-arm64.tar.gz) |
| Windows x86_64 | [aether-tunnel-windows-amd64.zip](https://github.com/fawney19/Aether/releases/download/tunnel-v0.3.12/aether-tunnel-windows-amd64.zip) |
<!-- DOWNLOAD_TABLE_END -->
上表展示的是最新已发布版本的下载链接。从下一次 `tunnel-v*` 发布开始,表格会自动补上 `Linux x86_64 (musl)` / `Linux ARM64 (musl)` 包,供 Alpine 等 musl 系统直接使用。
## 快速开始
### 一键安装 / 添加节点
一键脚本会自动从 GitHub Releases 中筛选最新的 `tunnel-v*` tag并按当前系统下载对应制品Linux x86_64/ARM64GNU 或 musl、macOS x86_64/ARM64、Windows x86_64。仓库的通用 `latest` release 可能不是 tunnel 版本,因此脚本不会使用 `/releases/latest`
脚本会安装/更新 `aether-tunnel` 二进制,并把新的服务器配置追加到 `aether-tunnel.toml``[[servers]]` 数组中;如果配置文件已存在,不会覆盖原有内容。检测到相同 `aether_url + node_name` 时会跳过追加。
macOS / Linux:
```bash
curl -fsSL https://raw.githubusercontent.com/fawney19/Aether/main/apps/aether-tunnel/install.sh | sh
```
Windows PowerShell:
```powershell
irm https://raw.githubusercontent.com/fawney19/Aether/main/apps/aether-tunnel/install.ps1 | iex
```
也可以用环境变量非交互式执行,适合在控制台“添加隧道节点”时生成命令:
```bash
curl -fsSL https://raw.githubusercontent.com/fawney19/Aether/main/apps/aether-tunnel/install.sh | \
AETHER_TUNNEL_AETHER_URL="https://aether.example.com" \
AETHER_TUNNEL_MANAGEMENT_TOKEN="ae_xxx" \
AETHER_TUNNEL_NODE_NAME="jp-proxy-01" \
sh
```
```powershell
$env:AETHER_TUNNEL_AETHER_URL = "https://aether.example.com"
$env:AETHER_TUNNEL_MANAGEMENT_TOKEN = "ae_xxx"
$env:AETHER_TUNNEL_NODE_NAME = "jp-proxy-01"
irm https://raw.githubusercontent.com/fawney19/Aether/main/apps/aether-tunnel/install.ps1 | iex
```
可选变量:`AETHER_TUNNEL_RELEASE_TAG` 固定安装某个 `tunnel-v*` tag`AETHER_TUNNEL_CONFIG` 指定配置文件路径,`AETHER_TUNNEL_INSTALL_DIR` 指定二进制安装目录。
```bash
# 1. 首次安装配置TUI 向导,勾选 Install Service 随系统启动服务)
sudo ./aether-tunnel setup
# 2. 日常管理 (勾选 Install Service 作为系统服务的情况下)
aether-tunnel status # 看状态
sudo aether-tunnel logs # 看日志
sudo aether-tunnel start # 启动服务
sudo aether-tunnel stop # 停止服务
sudo aether-tunnel restart # 重启服务
# 3. 重新配置(改完自动重启服务)
sudo aether-tunnel setup
# 4. 彻底卸载
sudo aether-tunnel uninstall
```
完成向导后, 配置自动保存到 `aether-tunnel.toml`,如果启用了 Install Service将自动注册并启动当前系统支持的服务`systemd``OpenRC`)。
### 直接运行
如果不需要安装为系统服务,可以直接运行。缺少必填参数时会自动进入 setup 向导:
```bash
./aether-tunnel
```
## 配置
配置按以下优先级加载(高优先级覆盖低优先级):
1. CLI 参数
2. 环境变量(`AETHER_TUNNEL_*`
3. 配置文件(`aether-tunnel.toml`,或通过 `AETHER_TUNNEL_CONFIG` 指定路径)
### 参数一览
#### 基础配置
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--aether-url` | `AETHER_TUNNEL_AETHER_URL` | **必填** | Aether 服务器地址 |
| `--management-token` | `AETHER_TUNNEL_MANAGEMENT_TOKEN` | **必填** | 管理员 Token`ae_xxx` 格式) |
| `--node-name` | `AETHER_TUNNEL_NODE_NAME` | **必填** | 节点名称标识 |
| `--public-ip` | `AETHER_TUNNEL_PUBLIC_IP` | 自动检测 | 公网 IP |
| `--node-region` | `AETHER_TUNNEL_NODE_REGION` | 自动检测 | 地区标识 |
| `--heartbeat-interval` | `AETHER_TUNNEL_HEARTBEAT_INTERVAL` | `5` | 心跳间隔(秒) |
| `--allowed-ports` | `AETHER_TUNNEL_ALLOWED_PORTS` | `80,443,8080,8443` | 允许代理的目标端口 |
| `--allow-private-targets` | `AETHER_TUNNEL_ALLOW_PRIVATE_TARGETS` | `true` | 允许 private/reserved 目标地址,通过后仍受 `allowed_ports` 限制;设为 `false` 可恢复严格拦截 |
#### Tunnel 连接
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--tunnel-connections` | `AETHER_TUNNEL_CONNECTIONS` | 自动(硬件估算) | 最小连接池大小;显式设置后默认固定为该值 |
| `--tunnel-connections-max` | `AETHER_TUNNEL_CONNECTIONS_MAX` | 自动(硬件估算) | 连接池自动扩容上限;大于 `tunnel_connections` 时启用 autoscale |
| `--tunnel-max-streams` | `AETHER_TUNNEL_MAX_STREAMS` | 自动(硬件估算) | 单连接最大并发 stream 数 |
| `--tunnel-ping-interval-ms` | `AETHER_TUNNEL_PING_INTERVAL_MS` | `10000` | WebSocket ping 周期(毫秒) |
| `--tunnel-connect-timeout-ms` | `AETHER_TUNNEL_CONNECT_TIMEOUT_MS` | `3000` | tunnel 建连超时(毫秒) |
| `--tunnel-stale-timeout-ms` | `AETHER_TUNNEL_STALE_TIMEOUT_MS` | `30000` | 无入站数据断连阈值(毫秒) |
| `--tunnel-scale-check-interval-ms` | `AETHER_TUNNEL_SCALE_CHECK_INTERVAL_MS` | `1000` | autoscale 采样周期(毫秒) |
| `--tunnel-scale-up-threshold-percent` | `AETHER_TUNNEL_SCALE_UP_THRESHOLD_PERCENT` | `50` | 单 tunnel 占用率超过该值时扩容 |
| `--tunnel-scale-down-threshold-percent` | `AETHER_TUNNEL_SCALE_DOWN_THRESHOLD_PERCENT` | `35` | 单 tunnel 占用率持续低于该值时允许缩容 |
| `--tunnel-scale-down-grace-secs` | `AETHER_TUNNEL_SCALE_DOWN_GRACE_SECS` | `15` | 低负载持续时间达到该值后才回收次级 tunnel |
| `--tunnel-tcp-keepalive-secs` | `AETHER_TUNNEL_TCP_KEEPALIVE_SECS` | `30` | TCP keepalive 初始延迟(秒) |
| `--tunnel-tcp-nodelay` | `AETHER_TUNNEL_TCP_NODELAY` | `true` | 禁用 Nagle 算法 |
| `--tunnel-reconnect-base-ms` | `AETHER_TUNNEL_RECONNECT_BASE_MS` | `50` | 指数退避基础延迟(毫秒) |
| `--tunnel-reconnect-max-ms` | `AETHER_TUNNEL_RECONNECT_MAX_MS` | `250` | 指数退避上限(毫秒) |
省略 `tunnel_connections`tunnel 会按设备能力自动计算一个基线值和偏单机上限的扩容上限:默认至少保留 2 条常驻 tunnel并会更早触发扩容如果显式设置了 `tunnel_connections` 但没有设置 `tunnel_connections_max`,则保持固定连接池,不自动扩缩。
#### 上游 HTTP 请求
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--upstream-connect-timeout-secs` | `AETHER_TUNNEL_UPSTREAM_CONNECT_TIMEOUT_SECS` | `30` | 上游建连超时(秒) |
| `--upstream-pool-max-idle-per-host` | `AETHER_TUNNEL_UPSTREAM_POOL_MAX_IDLE_PER_HOST` | `64` | 每 Host 最大空闲连接数 |
| `--upstream-pool-idle-timeout-secs` | `AETHER_TUNNEL_UPSTREAM_POOL_IDLE_TIMEOUT_SECS` | `300` | 连接池空闲超时(秒) |
| `--upstream-tcp-keepalive-secs` | `AETHER_TUNNEL_UPSTREAM_TCP_KEEPALIVE_SECS` | `60` | TCP keepalive0 关闭) |
| `--upstream-tcp-nodelay` | `AETHER_TUNNEL_UPSTREAM_TCP_NODELAY` | `true` | 启用 TCP_NODELAY |
| `--upstream-proxy-url` | `AETHER_TUNNEL_UPSTREAM_PROXY_URL` | 空 | 仅 provider 上游请求使用的出口代理 |
| `--redirect-replay-budget-bytes` | `AETHER_TUNNEL_REDIRECT_REPLAY_BUDGET_BYTES` | `5M` | 307/308 请求体重放的预读预算,支持 `K/M/G``0` 表示禁用 body replay buffering |
出口代理支持 `http://``socks5://``socks5h://`。配合 WARP sidecar 时可填写:
```toml
upstream_proxy_url = "socks5h://microwarp:1080"
```
如果需要让 Aether 管理 API 和 WebSocket tunnel 也走代理,使用 `aether_outbound_proxy_url`
#### Aether API 客户端
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--aether-request-timeout-secs` | `AETHER_TUNNEL_AETHER_REQUEST_TIMEOUT_SECS` | `10` | 请求总超时(秒) |
| `--aether-connect-timeout-secs` | `AETHER_TUNNEL_AETHER_CONNECT_TIMEOUT_SECS` | `10` | 建连超时(秒) |
| `--aether-outbound-proxy-url` | `AETHER_TUNNEL_AETHER_OUTBOUND_PROXY_URL` | 空 | Aether 注册、心跳和 WebSocket tunnel 回连使用的出口代理(默认不走代理) |
| `--aether-retry-max-attempts` | `AETHER_TUNNEL_AETHER_RETRY_MAX_ATTEMPTS` | `3` | 最大重试次数 |
#### DNS 与安全
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--allow-private-targets` | `AETHER_TUNNEL_ALLOW_PRIVATE_TARGETS` | `true` | 默认允许 private/reserved 目标地址;设为 `false` 可恢复拦截,且仅影响重启后的进程 |
| `--dns-cache-ttl-secs` | `AETHER_TUNNEL_DNS_CACHE_TTL_SECS` | `60` | DNS 缓存 TTL |
| `--dns-cache-capacity` | `AETHER_TUNNEL_DNS_CACHE_CAPACITY` | `1024` | DNS 缓存容量(条目数) |
#### 日志
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--log-level` | `AETHER_TUNNEL_LOG_LEVEL` | `info` | 日志级别 |
| `--log-destination` | `AETHER_TUNNEL_LOG_DESTINATION` | `both` | 输出到 `stdout`、文件或两者同时输出 |
| `--log-dir` | `AETHER_TUNNEL_LOG_DIR` | `logs` | 文件日志目录,`file/both` 时必填 |
| `--log-rotation` | `AETHER_TUNNEL_LOG_ROTATION` | `daily` | 文件日志按小时或按天轮转 |
| `--log-retention-days` | `AETHER_TUNNEL_LOG_RETENTION_DAYS` | `7` | 文件日志保留天数 |
| `--log-max-files` | `AETHER_TUNNEL_LOG_MAX_FILES` | `30` | 文件日志最多保留文件数 |
### 日志落点
- 默认 `AETHER_TUNNEL_LOG_DESTINATION=both`,同时输出到 stdout 和 `logs/` 文件目录
- 需要只交给容器日志驱动或宿主机服务管理器时,可改成 `stdout`setup TUI 里可用 `Save Logs to File` 开关关闭文件日志
- 文件日志固定写普通文本,并支持 `hourly/daily` 轮转;默认按天轮换、保留 7 天,最多保留 30 个文件
-`systemd``OpenRC` 安装时默认会额外打开文件日志到 `/var/log/aether-tunnel`
- OpenRC 安装时,`aether-tunnel logs` 实际读取 `/var/log/aether-tunnel/current.log``/var/log/aether-tunnel/error.log`;这些文件通常需要用 `sudo aether-tunnel logs` 查看
### 隧道健康上报Heartbeat
tunnel 会在心跳兼容字段 `proxy_metadata` 中主动上报隧道稳定性指标,便于后端直接入库/告警:
- `proxy_metadata.tunnel_metrics`:建连尝试/成功/失败、断开次数、累计在线时长、心跳 RTT、WebSocket 收发帧与字节等。
- `proxy_metadata.recent_tunnel_errors`:最近隧道异常事件(时间戳、类别、错误摘要,环形缓冲)。
说明:仅主连接(`conn=0`)发送 heartbeat避免多条 tunnel 重复上报同一份全局指标。
### 多服务器配置
`aether-tunnel.toml` 中使用 `[[servers]]` 配置 Aether 服务器。即使只有一个服务器,也必须写成一个 `[[servers]]` 条目;旧的顶层单服务器写法已不再支持。
```toml
[[servers]]
aether_url = "https://aether-1.example.com"
management_token = "ae_xxx"
node_name = "jp-proxy-01"
[[servers]]
aether_url = "https://aether-2.example.com"
management_token = "ae_yyy"
node_name = "jp-proxy-02"
```
## 发布新版本
推送 `tunnel-v*` 格式的 tagGitHub Actions 会自动:
- 编译所有平台二进制并发布到 Releases
- 更新 README 中的下载链接表格
```bash
git tag tunnel-v0.2.0
git push origin tunnel-v0.2.0
```

View File

@@ -0,0 +1,161 @@
$ErrorActionPreference = 'Stop'
$Repo = if ($env:AETHER_TUNNEL_RELEASE_REPO) { $env:AETHER_TUNNEL_RELEASE_REPO } else { 'fawney19/Aether' }
$ReleaseTag = $env:AETHER_TUNNEL_RELEASE_TAG
$InstallDir = $env:AETHER_TUNNEL_INSTALL_DIR
$ConfigPath = $env:AETHER_TUNNEL_CONFIG
function Say([string]$Message) { Write-Host "[Aether Tunnel] $Message" }
function Fail([string]$Message) { throw "[Aether Tunnel] $Message" }
function Prompt-IfEmpty([string]$Name, [string]$Value, [string]$Prompt) {
if (-not [string]::IsNullOrWhiteSpace($Value)) { return $Value }
$Read = Read-Host $Prompt
if ([string]::IsNullOrWhiteSpace($Read)) { Fail "$Name cannot be empty" }
return $Read
}
function ConvertTo-TomlQuotedString([string]$Value) {
return ($Value | ConvertTo-Json -Compress)
}
function Resolve-LatestTunnelTag {
if (-not [string]::IsNullOrWhiteSpace($ReleaseTag)) { return $ReleaseTag }
$Uri = "https://api.github.com/repos/$Repo/releases?per_page=100"
$Releases = Invoke-RestMethod -Uri $Uri -Headers @{ 'User-Agent' = 'aether-tunnel-installer' }
$TunnelReleases = @($Releases | Where-Object { -not $_.draft -and $_.tag_name -like 'tunnel-v*' } | Sort-Object published_at -Descending)
if ($TunnelReleases.Count -eq 0) { Fail "No tunnel-v* release found in $Repo" }
return $TunnelReleases[0].tag_name
}
function Test-IsAdministrator {
$Identity = [Security.Principal.WindowsIdentity]::GetCurrent()
$Principal = [Security.Principal.WindowsPrincipal]::new($Identity)
return $Principal.IsInRole([Security.Principal.WindowsBuiltInRole]::Administrator)
}
function Initialize-Paths {
if ([string]::IsNullOrWhiteSpace($script:InstallDir)) {
if (Test-IsAdministrator) {
$script:InstallDir = Join-Path $env:ProgramFiles 'AetherTunnel'
} else {
$script:InstallDir = Join-Path $env:LOCALAPPDATA 'AetherTunnel'
}
}
if ([string]::IsNullOrWhiteSpace($script:ConfigPath)) {
if (Test-IsAdministrator) {
$script:ConfigPath = Join-Path $env:ProgramData 'AetherTunnel\aether-tunnel.toml'
} else {
$script:ConfigPath = Join-Path $env:APPDATA 'AetherTunnel\aether-tunnel.toml'
}
}
}
function Install-AetherTunnelBinary([string]$Tag, [string]$TempDir) {
if (-not [Environment]::Is64BitOperatingSystem) { Fail 'Windows release currently supports amd64 only' }
$Asset = 'aether-tunnel-windows-amd64.zip'
$Base = "https://github.com/$Repo/releases/download/$Tag"
$Archive = Join-Path $TempDir $Asset
$Sums = Join-Path $TempDir 'SHA256SUMS.txt'
Say "Downloading $Tag / $Asset"
Invoke-WebRequest -Uri "$Base/$Asset" -OutFile $Archive
try { Invoke-WebRequest -Uri "$Base/SHA256SUMS.txt" -OutFile $Sums } catch { $Sums = $null }
if ($Sums -and (Test-Path $Sums)) {
$ExpectedLine = Get-Content $Sums | Where-Object { $_ -match "\s$([regex]::Escape($Asset))$" } | Select-Object -First 1
if ($ExpectedLine) {
$Expected = ($ExpectedLine -split '\s+')[0]
$Actual = (Get-FileHash -Algorithm SHA256 $Archive).Hash.ToLowerInvariant()
if ($Actual -ne $Expected.ToLowerInvariant()) { Fail "SHA256 verification failed for $Asset" }
}
}
$ExtractDir = Join-Path $TempDir 'extract'
Expand-Archive -Path $Archive -DestinationPath $ExtractDir -Force
$Binary = Join-Path $ExtractDir 'aether-tunnel.exe'
if (-not (Test-Path $Binary)) { Fail 'aether-tunnel.exe not found in release asset' }
New-Item -ItemType Directory -Force -Path $script:InstallDir | Out-Null
Copy-Item $Binary (Join-Path $script:InstallDir 'aether-tunnel.exe') -Force
Say "Installed binary: $(Join-Path $script:InstallDir 'aether-tunnel.exe')"
}
function Test-LegacySingleServerConfig([string]$Path) {
if (-not (Test-Path $Path)) { return $false }
foreach ($Line in Get-Content $Path) {
if ($Line -match '^\s*\[') { return $false }
if ($Line -match '^\s*(aether_url|management_token)\s*=') { return $true }
}
return $false
}
function Test-ServerExists([string]$Path, [string]$QuotedUrl, [string]$QuotedName) {
if (-not (Test-Path $Path)) { return $false }
$FoundUrl = $false
$FoundName = $false
foreach ($Line in Get-Content $Path) {
if ($Line -match '^\s*\[\[servers\]\]\s*$') {
if ($FoundUrl -and $FoundName) { return $true }
$FoundUrl = $false
$FoundName = $false
}
if ($Line.Trim() -eq "aether_url = $QuotedUrl") { $FoundUrl = $true }
if ($Line.Trim() -eq "node_name = $QuotedName") { $FoundName = $true }
}
return ($FoundUrl -and $FoundName)
}
function Add-ServerConfig([string]$AetherUrl, [string]$ManagementToken, [string]$NodeName) {
$ConfigDir = Split-Path -Parent $script:ConfigPath
New-Item -ItemType Directory -Force -Path $ConfigDir | Out-Null
if (Test-LegacySingleServerConfig $script:ConfigPath) {
Fail "Existing config uses removed top-level aether_url/management_token. Run aether-tunnel setup to migrate to [[servers]] first: $script:ConfigPath"
}
$QuotedUrl = ConvertTo-TomlQuotedString $AetherUrl
$QuotedToken = ConvertTo-TomlQuotedString $ManagementToken
$QuotedName = ConvertTo-TomlQuotedString $NodeName
if (Test-ServerExists $script:ConfigPath $QuotedUrl $QuotedName) {
Say "Same aether_url + node_name already exists, skipping config append: $script:ConfigPath"
return
}
if (Test-Path $script:ConfigPath) {
Copy-Item $script:ConfigPath "$script:ConfigPath.bak.$(Get-Date -Format yyyyMMddHHmmss)" -Force
}
$Prefix = if ((Test-Path $script:ConfigPath) -and ((Get-Item $script:ConfigPath).Length -gt 0)) { "`n" } else { '' }
$Block = @(
"$Prefix# Added by Aether Tunnel one-click installer. Existing config is preserved.",
'[[servers]]',
"aether_url = $QuotedUrl",
"management_token = $QuotedToken",
"node_name = $QuotedName"
) -join "`n"
Add-Content -Path $script:ConfigPath -Value ($Block + "`n") -Encoding UTF8
Say "Appended [[servers]] to: $script:ConfigPath"
}
function Main {
Initialize-Paths
$AetherUrl = Prompt-IfEmpty 'AETHER_TUNNEL_AETHER_URL' $env:AETHER_TUNNEL_AETHER_URL 'Aether URL'
$ManagementToken = Prompt-IfEmpty 'AETHER_TUNNEL_MANAGEMENT_TOKEN' $env:AETHER_TUNNEL_MANAGEMENT_TOKEN 'Management token (ae_xxx)'
$NodeName = Prompt-IfEmpty 'AETHER_TUNNEL_NODE_NAME' $env:AETHER_TUNNEL_NODE_NAME 'Node name'
$TempDir = Join-Path ([IO.Path]::GetTempPath()) ("aether-tunnel-" + [Guid]::NewGuid().ToString('N'))
New-Item -ItemType Directory -Force -Path $TempDir | Out-Null
try {
$Tag = Resolve-LatestTunnelTag
Install-AetherTunnelBinary $Tag $TempDir
Add-ServerConfig $AetherUrl $ManagementToken $NodeName
} finally {
Remove-Item -Recurse -Force $TempDir -ErrorAction SilentlyContinue
}
Say 'Complete. Start or configure the node with:'
Say " & '$(Join-Path $script:InstallDir 'aether-tunnel.exe')' setup '$script:ConfigPath'"
}
Main

241
apps/aether-tunnel/install.sh Executable file
View File

@@ -0,0 +1,241 @@
#!/bin/sh
set -eu
REPO="${AETHER_TUNNEL_RELEASE_REPO:-fawney19/Aether}"
TAG="${AETHER_TUNNEL_RELEASE_TAG:-}"
INSTALL_DIR="${AETHER_TUNNEL_INSTALL_DIR:-}"
CONFIG_PATH="${AETHER_TUNNEL_CONFIG:-}"
TMP_DIR=""
say() { printf '%s\n' "[Aether Tunnel] $1"; }
fail() { printf '%s\n' "[Aether Tunnel] $1" >&2; exit 1; }
cleanup() {
if [ -n "$TMP_DIR" ] && [ -d "$TMP_DIR" ]; then
rm -rf "$TMP_DIR"
fi
}
trap cleanup EXIT INT TERM
need_cmd() {
command -v "$1" >/dev/null 2>&1 || fail "缺少命令:$1"
}
download() {
url="$1"
out="$2"
if command -v curl >/dev/null 2>&1; then
curl -fL --retry 3 --connect-timeout 10 -o "$out" "$url"
elif command -v wget >/dev/null 2>&1; then
wget -O "$out" "$url"
else
fail "需要 curl 或 wget 下载 release 制品"
fi
}
prompt_if_empty() {
name="$1"
value="$2"
prompt="$3"
if [ -n "$value" ]; then
printf '%s' "$value"
return
fi
printf '%s' "$prompt" >&2
if [ -r /dev/tty ]; then
IFS= read -r value < /dev/tty
else
fail "$name 未通过环境变量提供,且当前环境无法交互输入"
fi
[ -n "$value" ] || fail "$name 不能为空"
printf '%s' "$value"
}
toml_quote() {
value="$1"
if command -v python3 >/dev/null 2>&1; then
python3 -c 'import json,sys; print(json.dumps(sys.argv[1], ensure_ascii=False))' "$value"
else
escaped=$(printf '%s' "$value" | sed 's/\\/\\\\/g; s/"/\\"/g')
printf '"%s"\n' "$escaped"
fi
}
resolve_latest_tunnel_tag() {
[ -n "$TAG" ] && { printf '%s\n' "$TAG"; return; }
api_url="https://api.github.com/repos/${REPO}/releases?per_page=100"
releases="$TMP_DIR/releases.json"
download "$api_url" "$releases" >/dev/null 2>&1 || fail "无法读取 GitHub Releases$api_url"
if command -v python3 >/dev/null 2>&1; then
python3 - "$releases" <<'PY'
import json, sys
releases = json.load(open(sys.argv[1], encoding='utf-8'))
tunnel = [r for r in releases if not r.get('draft') and str(r.get('tag_name', '')).startswith('tunnel-v')]
tunnel.sort(key=lambda r: r.get('published_at') or r.get('created_at') or '', reverse=True)
if tunnel:
print(tunnel[0]['tag_name'])
PY
else
grep -o '"tag_name"[[:space:]]*:[[:space:]]*"tunnel-v[^"]*"' "$releases" | head -n 1 | sed 's/.*"\(tunnel-v[^"]*\)".*/\1/'
fi
}
detect_asset() {
os=$(uname -s 2>/dev/null || printf unknown)
arch=$(uname -m 2>/dev/null || printf unknown)
case "$os" in
Linux) platform=linux ;;
Darwin) platform=macos ;;
MINGW*|MSYS*|CYGWIN*) fail "检测到 Windows shell请使用 PowerShellirm <install.ps1-url> | iex" ;;
*) fail "不支持的系统:$os" ;;
esac
case "$arch" in
x86_64|amd64) cpu=amd64 ;;
aarch64|arm64) cpu=arm64 ;;
*) fail "不支持的 CPU 架构:$arch" ;;
esac
if [ "$platform" = "linux" ] && command -v ldd >/dev/null 2>&1 && ldd --version 2>&1 | grep -qi musl; then
printf 'aether-tunnel-linux-musl-%s.tar.gz\n' "$cpu"
else
printf 'aether-tunnel-%s-%s.tar.gz\n' "$platform" "$cpu"
fi
}
choose_paths() {
if [ -z "$INSTALL_DIR" ]; then
if [ "$(id -u 2>/dev/null || printf 1)" = "0" ]; then
INSTALL_DIR="/usr/local/bin"
else
INSTALL_DIR="$HOME/.local/bin"
fi
fi
if [ -z "$CONFIG_PATH" ]; then
if [ "$(id -u 2>/dev/null || printf 1)" = "0" ]; then
CONFIG_PATH="/etc/aether-tunnel/aether-tunnel.toml"
else
CONFIG_PATH="$HOME/.aether-tunnel/aether-tunnel.toml"
fi
fi
}
verify_checksum() {
archive="$1"
sums="$2"
asset="$3"
[ -f "$sums" ] || return 0
expected=$(awk -v asset="$asset" '$2 == asset { print $1 }' "$sums" | head -n 1)
[ -n "$expected" ] || return 0
if command -v sha256sum >/dev/null 2>&1; then
actual=$(sha256sum "$archive" | awk '{print $1}')
elif command -v shasum >/dev/null 2>&1; then
actual=$(shasum -a 256 "$archive" | awk '{print $1}')
else
say "未找到 sha256sum/shasum跳过校验"
return 0
fi
[ "$actual" = "$expected" ] || fail "SHA256 校验失败:$asset"
}
install_binary() {
tag="$1"
asset="$2"
base="https://github.com/${REPO}/releases/download/${tag}"
archive="$TMP_DIR/$asset"
say "下载 $tag / $asset"
download "$base/$asset" "$archive"
download "$base/SHA256SUMS.txt" "$TMP_DIR/SHA256SUMS.txt" >/dev/null 2>&1 || true
verify_checksum "$archive" "$TMP_DIR/SHA256SUMS.txt" "$asset"
tar -xzf "$archive" -C "$TMP_DIR"
[ -f "$TMP_DIR/aether-tunnel" ] || fail "制品中未找到 aether-tunnel"
mkdir -p "$INSTALL_DIR"
cp "$TMP_DIR/aether-tunnel" "$INSTALL_DIR/aether-tunnel"
chmod +x "$INSTALL_DIR/aether-tunnel"
say "已安装二进制:$INSTALL_DIR/aether-tunnel"
}
has_legacy_single_server_keys() {
[ -f "$CONFIG_PATH" ] || return 1
awk '
/^[[:space:]]*\[/ { exit }
/^[[:space:]]*(aether_url|management_token)[[:space:]]*=/ { found=1; exit }
END { exit found ? 0 : 1 }
' "$CONFIG_PATH"
}
server_exists() {
[ -f "$CONFIG_PATH" ] || return 1
quoted_url="$1"
quoted_name="$2"
awk -v url="aether_url = $quoted_url" -v name="node_name = $quoted_name" '
BEGIN { found_url=0; found_name=0 }
/^\[\[servers\]\]/ {
if (found_url && found_name) { found=1 }
found_url=0; found_name=0
}
$0 == url { found_url=1 }
$0 == name { found_name=1 }
END { if (found_url && found_name) { found=1 }; exit found ? 0 : 1 }
' "$CONFIG_PATH"
}
append_server_config() {
aether_url="$1"
management_token="$2"
node_name="$3"
mkdir -p "$(dirname "$CONFIG_PATH")"
quoted_url=$(toml_quote "$aether_url")
quoted_token=$(toml_quote "$management_token")
quoted_name=$(toml_quote "$node_name")
if has_legacy_single_server_keys; then
fail "现有配置仍使用旧的顶层 aether_url/management_token请先运行 aether-tunnel setup 迁移为 [[servers]] 后重试:$CONFIG_PATH"
fi
if server_exists "$quoted_url" "$quoted_name"; then
say "配置中已存在相同 aether_url + node_name跳过追加$CONFIG_PATH"
return
fi
if [ -f "$CONFIG_PATH" ]; then
cp "$CONFIG_PATH" "$CONFIG_PATH.bak.$(date +%Y%m%d%H%M%S)"
fi
{
if [ -f "$CONFIG_PATH" ] && [ -s "$CONFIG_PATH" ]; then
printf '\n'
fi
printf '# Added by Aether Tunnel one-click installer. Existing config is preserved.\n'
printf '[[servers]]\n'
printf 'aether_url = %s\n' "$quoted_url"
printf 'management_token = %s\n' "$quoted_token"
printf 'node_name = %s\n' "$quoted_name"
} >> "$CONFIG_PATH"
chmod 600 "$CONFIG_PATH" 2>/dev/null || true
say "已追加 [[servers]] 到:$CONFIG_PATH"
}
main() {
TMP_DIR=$(mktemp -d 2>/dev/null || mktemp -d -t aether-tunnel)
need_cmd tar
choose_paths
aether_url=$(prompt_if_empty AETHER_TUNNEL_AETHER_URL "${AETHER_TUNNEL_AETHER_URL:-}" "Aether URL: ")
management_token=$(prompt_if_empty AETHER_TUNNEL_MANAGEMENT_TOKEN "${AETHER_TUNNEL_MANAGEMENT_TOKEN:-}" "Management token (ae_xxx): ")
node_name=$(prompt_if_empty AETHER_TUNNEL_NODE_NAME "${AETHER_TUNNEL_NODE_NAME:-}" "Node name: ")
tag=$(resolve_latest_tunnel_tag)
[ -n "$tag" ] || fail "没有找到可用的 tunnel-v* release"
asset=$(detect_asset)
install_binary "$tag" "$asset"
append_server_config "$aether_url" "$management_token" "$node_name"
say "完成。运行以下命令启动/配置服务:"
say " $INSTALL_DIR/aether-tunnel setup $CONFIG_PATH"
}
main "$@"

File diff suppressed because it is too large Load Diff

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,465 @@
use std::io;
use std::net::{IpAddr, SocketAddr};
use std::time::Duration;
use base64::Engine;
use socket2::{SockRef, TcpKeepalive};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpStream;
use url::Url;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum UpstreamProxyScheme {
Http,
Socks5,
Socks5h,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct UpstreamProxyConfig {
raw: String,
scheme: UpstreamProxyScheme,
host: String,
port: u16,
username: Option<String>,
password: Option<String>,
}
impl UpstreamProxyConfig {
pub(crate) fn parse(raw: &str) -> Result<Self, String> {
let trimmed = raw.trim();
if trimmed.is_empty() {
return Err("upstream proxy URL must not be empty".to_string());
}
let parsed =
Url::parse(trimmed).map_err(|err| format!("invalid upstream proxy URL: {err}"))?;
let scheme = match parsed.scheme().to_ascii_lowercase().as_str() {
"http" => UpstreamProxyScheme::Http,
"socks5" => UpstreamProxyScheme::Socks5,
"socks5h" => UpstreamProxyScheme::Socks5h,
other => {
return Err(format!(
"unsupported upstream proxy scheme `{other}`; use http, socks5, or socks5h"
))
}
};
let host = parsed
.host_str()
.map(str::trim)
.filter(|value| !value.is_empty())
.ok_or_else(|| "upstream proxy URL must include a host".to_string())?
.to_string();
let port = parsed.port().unwrap_or(match scheme {
UpstreamProxyScheme::Http => 80,
UpstreamProxyScheme::Socks5 | UpstreamProxyScheme::Socks5h => 1080,
});
let username = non_empty_url_part(parsed.username());
let password = parsed.password().and_then(non_empty_url_part);
Ok(Self {
raw: trimmed.to_string(),
scheme,
host,
port,
username,
password,
})
}
pub(crate) fn scheme(&self) -> UpstreamProxyScheme {
self.scheme
}
pub(crate) fn host(&self) -> &str {
&self.host
}
pub(crate) fn port(&self) -> u16 {
self.port
}
pub(crate) fn username(&self) -> Option<&str> {
self.username.as_deref()
}
pub(crate) fn password(&self) -> Option<&str> {
self.password.as_deref()
}
pub(crate) fn uses_remote_dns(&self) -> bool {
self.scheme == UpstreamProxyScheme::Socks5h
}
pub(crate) fn basic_auth_header(&self) -> Option<String> {
let username = self.username()?;
let mut credentials = String::with_capacity(
username.len() + self.password.as_ref().map(|value| value.len()).unwrap_or(0) + 1,
);
credentials.push_str(username);
credentials.push(':');
if let Some(password) = self.password() {
credentials.push_str(password);
}
Some(format!(
"Basic {}",
base64::engine::general_purpose::STANDARD.encode(credentials)
))
}
pub(crate) fn redacted_url(&self) -> String {
let Ok(mut parsed) = Url::parse(&self.raw) else {
return "<invalid>".to_string();
};
if !parsed.username().is_empty() {
let _ = parsed.set_username("****");
}
if parsed.password().is_some() {
let _ = parsed.set_password(Some("****"));
}
parsed.to_string()
}
}
#[derive(Debug, Clone, Copy)]
pub(crate) struct ProxyConnectOptions {
pub connect_timeout: Duration,
pub tcp_nodelay: bool,
pub tcp_keepalive: Option<Duration>,
}
pub(crate) async fn connect_target_via_proxy(
proxy: &UpstreamProxyConfig,
target_host: &str,
target_port: u16,
options: ProxyConnectOptions,
) -> io::Result<TcpStream> {
let mut tcp = connect_proxy_tcp(
proxy,
options.connect_timeout,
options.tcp_nodelay,
options.tcp_keepalive,
)
.await?;
match proxy.scheme() {
UpstreamProxyScheme::Http => {
http_connect(&mut tcp, &target_authority(target_host, target_port), proxy).await?;
}
UpstreamProxyScheme::Socks5 | UpstreamProxyScheme::Socks5h => {
socks5_connect(&mut tcp, proxy, target_host, target_port).await?;
}
}
Ok(tcp)
}
pub(crate) async fn connect_proxy_tcp(
proxy: &UpstreamProxyConfig,
connect_timeout: Duration,
tcp_nodelay: bool,
tcp_keepalive: Option<Duration>,
) -> io::Result<TcpStream> {
let resolved = tokio::time::timeout(
connect_timeout,
tokio::net::lookup_host((proxy.host(), proxy.port())),
)
.await
.map_err(|_| io::Error::new(io::ErrorKind::TimedOut, "proxy DNS timeout"))?
.map_err(|err| io::Error::other(format!("proxy DNS failed: {err}")))?;
let mut last_error = None;
for addr in resolved {
match tokio::time::timeout(connect_timeout, TcpStream::connect(addr)).await {
Ok(Ok(stream)) => {
configure_tcp_stream(&stream, tcp_nodelay, tcp_keepalive)?;
return Ok(stream);
}
Ok(Err(error)) => last_error = Some(error),
Err(_) => {
last_error = Some(io::Error::new(
io::ErrorKind::TimedOut,
format!("proxy connect timeout: {addr}"),
));
}
}
}
Err(last_error.unwrap_or_else(|| io::Error::other("proxy DNS returned no addresses")))
}
fn configure_tcp_stream(
stream: &TcpStream,
tcp_nodelay: bool,
tcp_keepalive: Option<Duration>,
) -> io::Result<()> {
stream.set_nodelay(tcp_nodelay)?;
if let Some(keepalive) = tcp_keepalive {
let keepalive = TcpKeepalive::new().with_time(keepalive);
SockRef::from(stream).set_tcp_keepalive(&keepalive)?;
}
Ok(())
}
pub(crate) async fn http_connect(
stream: &mut TcpStream,
target_authority: &str,
proxy: &UpstreamProxyConfig,
) -> io::Result<()> {
let mut request = format!(
"CONNECT {target_authority} HTTP/1.1\r\nHost: {target_authority}\r\nProxy-Connection: Keep-Alive\r\n"
);
if let Some(auth) = proxy.basic_auth_header() {
request.push_str("Proxy-Authorization: ");
request.push_str(&auth);
request.push_str("\r\n");
}
request.push_str("\r\n");
stream.write_all(request.as_bytes()).await?;
stream.flush().await?;
let mut response = Vec::with_capacity(1024);
let mut chunk = [0u8; 1024];
loop {
if response.len() >= 16 * 1024 {
return Err(io::Error::other("proxy CONNECT response too large"));
}
let n = stream.read(&mut chunk).await?;
if n == 0 {
return Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
"proxy closed during CONNECT",
));
}
response.extend_from_slice(&chunk[..n]);
if response.windows(4).any(|window| window == b"\r\n\r\n") {
break;
}
}
let status_line_end = response
.windows(2)
.position(|window| window == b"\r\n")
.ok_or_else(|| io::Error::other("proxy CONNECT response missing status line"))?;
let status_line = std::str::from_utf8(&response[..status_line_end])
.map_err(|_| io::Error::other("proxy CONNECT status line is not UTF-8"))?;
let status = status_line.split_whitespace().nth(1).unwrap_or_default();
if status == "200" {
Ok(())
} else {
Err(io::Error::other(format!(
"proxy CONNECT failed: {status_line}"
)))
}
}
pub(crate) async fn socks5_connect(
stream: &mut TcpStream,
proxy: &UpstreamProxyConfig,
target_host: &str,
target_port: u16,
) -> io::Result<()> {
let requires_auth = proxy.username().is_some();
if requires_auth {
stream.write_all(&[0x05, 0x02, 0x00, 0x02]).await?;
} else {
stream.write_all(&[0x05, 0x01, 0x00]).await?;
}
let mut method_response = [0u8; 2];
stream.read_exact(&mut method_response).await?;
if method_response[0] != 0x05 {
return Err(io::Error::other("invalid SOCKS5 method response"));
}
match method_response[1] {
0x00 => {}
0x02 => socks5_authenticate(stream, proxy).await?,
0xff => return Err(io::Error::other("SOCKS5 proxy rejected all auth methods")),
method => {
return Err(io::Error::other(format!(
"SOCKS5 proxy selected unsupported auth method 0x{method:02x}"
)))
}
}
let address = socks5_target_address(target_host, target_port, proxy.uses_remote_dns()).await?;
stream.write_all(&address).await?;
let mut response = [0u8; 4];
stream.read_exact(&mut response).await?;
if response[0] != 0x05 {
return Err(io::Error::other("invalid SOCKS5 connect response"));
}
if response[1] != 0x00 {
return Err(io::Error::other(format!(
"SOCKS5 connect failed: {}",
socks5_reply_message(response[1])
)));
}
match response[3] {
0x01 => {
let mut ignored = [0u8; 4 + 2];
stream.read_exact(&mut ignored).await?;
}
0x03 => {
let mut len = [0u8; 1];
stream.read_exact(&mut len).await?;
let mut ignored = vec![0u8; len[0] as usize + 2];
stream.read_exact(&mut ignored).await?;
}
0x04 => {
let mut ignored = [0u8; 16 + 2];
stream.read_exact(&mut ignored).await?;
}
atyp => {
return Err(io::Error::other(format!(
"SOCKS5 proxy returned unsupported address type 0x{atyp:02x}"
)))
}
}
Ok(())
}
async fn socks5_authenticate(
stream: &mut TcpStream,
proxy: &UpstreamProxyConfig,
) -> io::Result<()> {
let username = proxy.username().unwrap_or_default().as_bytes();
let password = proxy.password().unwrap_or_default().as_bytes();
if username.len() > u8::MAX as usize || password.len() > u8::MAX as usize {
return Err(io::Error::other(
"SOCKS5 username/password must be at most 255 bytes",
));
}
let mut request = Vec::with_capacity(username.len() + password.len() + 3);
request.push(0x01);
request.push(username.len() as u8);
request.extend_from_slice(username);
request.push(password.len() as u8);
request.extend_from_slice(password);
stream.write_all(&request).await?;
let mut response = [0u8; 2];
stream.read_exact(&mut response).await?;
if response[0] != 0x01 || response[1] != 0x00 {
return Err(io::Error::other("SOCKS5 username/password auth failed"));
}
Ok(())
}
pub(crate) async fn socks5_target_address(
target_host: &str,
target_port: u16,
remote_dns: bool,
) -> io::Result<Vec<u8>> {
let mut request = vec![0x05, 0x01, 0x00];
if let Ok(ip) = target_host.parse::<IpAddr>() {
push_socks5_ip_address(&mut request, ip);
} else if remote_dns {
let host = target_host.as_bytes();
if host.len() > u8::MAX as usize {
return Err(io::Error::other("SOCKS5 target hostname is too long"));
}
request.push(0x03);
request.push(host.len() as u8);
request.extend_from_slice(host);
} else {
let mut resolved = tokio::net::lookup_host((target_host, target_port))
.await
.map_err(|err| io::Error::other(format!("SOCKS5 target DNS failed: {err}")))?;
let addr = resolved
.next()
.ok_or_else(|| io::Error::other("SOCKS5 target DNS returned no addresses"))?;
push_socks5_socket_address(&mut request, addr);
}
request.extend_from_slice(&target_port.to_be_bytes());
Ok(request)
}
fn push_socks5_socket_address(request: &mut Vec<u8>, addr: SocketAddr) {
push_socks5_ip_address(request, addr.ip());
}
fn push_socks5_ip_address(request: &mut Vec<u8>, ip: IpAddr) {
match ip {
IpAddr::V4(ip) => {
request.push(0x01);
request.extend_from_slice(&ip.octets());
}
IpAddr::V6(ip) => {
request.push(0x04);
request.extend_from_slice(&ip.octets());
}
}
}
fn socks5_reply_message(reply: u8) -> &'static str {
match reply {
0x01 => "general failure",
0x02 => "connection not allowed",
0x03 => "network unreachable",
0x04 => "host unreachable",
0x05 => "connection refused",
0x06 => "TTL expired",
0x07 => "command not supported",
0x08 => "address type not supported",
_ => "unknown error",
}
}
pub(crate) fn target_authority(host: &str, port: u16) -> String {
if host.contains(':') && !host.starts_with('[') {
format!("[{host}]:{port}")
} else {
format!("{host}:{port}")
}
}
fn non_empty_url_part(value: &str) -> Option<String> {
let value = value.trim();
if value.is_empty() {
None
} else {
Some(value.to_string())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parses_http_proxy_with_default_port() {
let proxy = UpstreamProxyConfig::parse("http://proxy.example").expect("proxy should parse");
assert_eq!(proxy.scheme(), UpstreamProxyScheme::Http);
assert_eq!(proxy.host(), "proxy.example");
assert_eq!(proxy.port(), 80);
}
#[test]
fn parses_socks5h_proxy_with_auth() {
let proxy = UpstreamProxyConfig::parse("socks5h://user:pass@127.0.0.1:1080")
.expect("proxy should parse");
assert_eq!(proxy.scheme(), UpstreamProxyScheme::Socks5h);
assert_eq!(proxy.username(), Some("user"));
assert_eq!(proxy.password(), Some("pass"));
assert!(proxy.uses_remote_dns());
assert_eq!(
proxy.basic_auth_header().as_deref(),
Some("Basic dXNlcjpwYXNz")
);
}
#[test]
fn rejects_unsupported_proxy_scheme() {
let error = UpstreamProxyConfig::parse("https://proxy.example:8443")
.expect_err("https proxy scheme should be rejected");
assert!(error.contains("unsupported upstream proxy scheme"));
}
}

View File

@@ -0,0 +1,193 @@
use std::sync::Mutex;
use std::time::{SystemTime, UNIX_EPOCH};
use serde::Serialize;
use sysinfo::{get_current_pid, Pid, ProcessesToUpdate, System};
use tracing::info;
/// Hardware information collected at startup.
///
/// The struct is `Serialize`-able so it can be sent directly as the
/// `hardware_info` JSON bag in the registration request. New fields
/// can be added without database schema migrations.
#[derive(Debug, Clone, Serialize)]
pub struct HardwareInfo {
pub cpu_cores: u32,
pub total_memory_mb: u64,
pub os_info: String,
pub fd_limit: u64,
#[serde(skip)]
pub estimated_max_concurrency: u64,
}
/// Runtime resource usage sampled during heartbeat reporting.
#[derive(Debug, Clone, Serialize)]
pub struct RuntimeResourceSnapshot {
pub sampled_at_unix_secs: u64,
pub system_cpu_usage_percent: f64,
pub process_cpu_usage_percent: f64,
pub memory_total_bytes: u64,
pub memory_used_bytes: u64,
pub memory_available_bytes: u64,
pub memory_used_percent: f64,
pub process_memory_bytes: u64,
pub process_virtual_memory_bytes: u64,
pub process_memory_percent: f64,
pub load_average_1m: f64,
pub load_average_5m: f64,
pub load_average_15m: f64,
pub system_uptime_secs: u64,
pub process_uptime_secs: Option<u64>,
}
/// Small, reusable sysinfo monitor. Keeping it alive between samples makes CPU
/// usage deltas meaningful without re-enumerating the whole machine every time.
pub struct RuntimeResourceMonitor {
system: Mutex<System>,
current_pid: Option<Pid>,
}
impl RuntimeResourceMonitor {
pub fn new() -> Self {
let mut system = System::new_all();
let current_pid = get_current_pid().ok();
if let Some(pid) = current_pid {
system.refresh_processes(ProcessesToUpdate::Some(&[pid]), true);
}
system.refresh_cpu_usage();
system.refresh_memory();
Self {
system: Mutex::new(system),
current_pid,
}
}
pub fn snapshot(&self) -> RuntimeResourceSnapshot {
let mut system = match self.system.lock() {
Ok(guard) => guard,
Err(poisoned) => poisoned.into_inner(),
};
system.refresh_cpu_usage();
system.refresh_memory();
if let Some(pid) = self.current_pid {
system.refresh_processes(ProcessesToUpdate::Some(&[pid]), true);
}
let memory_total_bytes = system.total_memory();
let memory_used_bytes = system.used_memory();
let memory_available_bytes = system.available_memory();
let (
process_cpu_usage_percent,
process_memory_bytes,
process_virtual_memory_bytes,
process_uptime_secs,
) = self
.current_pid
.and_then(|pid| system.process(pid))
.map(|process| {
(
process.cpu_usage() as f64,
process.memory(),
process.virtual_memory(),
Some(process.run_time()),
)
})
.unwrap_or((0.0, 0, 0, None));
let load = System::load_average();
RuntimeResourceSnapshot {
sampled_at_unix_secs: current_unix_secs(),
system_cpu_usage_percent: system.global_cpu_usage() as f64,
process_cpu_usage_percent,
memory_total_bytes,
memory_used_bytes,
memory_available_bytes,
memory_used_percent: ratio_percent(memory_used_bytes, memory_total_bytes),
process_memory_bytes,
process_virtual_memory_bytes,
process_memory_percent: ratio_percent(process_memory_bytes, memory_total_bytes),
load_average_1m: load.one,
load_average_5m: load.five,
load_average_15m: load.fifteen,
system_uptime_secs: System::uptime(),
process_uptime_secs,
}
}
}
/// Collect hardware information and estimate max concurrency.
///
/// Should be called once at startup -- hardware does not change at runtime.
pub fn collect() -> HardwareInfo {
let sys = System::new_all();
let cpu_cores = sys.cpus().len() as u32;
let total_memory_mb = sys.total_memory() / (1024 * 1024);
let os_info = format!(
"{} {}",
System::name().unwrap_or_else(|| "Unknown".into()),
System::os_version().unwrap_or_default(),
)
.trim()
.to_string();
// Estimate max concurrent connections:
// - Each tokio async task uses ~8-16 KB stack + heap buffers
// - OS file descriptor limit is often the real bottleneck
// - Conservative formula: min(fd_limit - 100, ram_mb * 40, cpu_cores * 2000)
let fd_limit = get_fd_limit();
let by_fd = fd_limit.saturating_sub(100);
let by_ram = total_memory_mb.saturating_mul(40);
let by_cpu = (cpu_cores as u64).saturating_mul(2000);
let estimated_max_concurrency = by_fd.min(by_ram).min(by_cpu);
info!(
cpu_cores,
total_memory_mb,
os_info = %os_info,
fd_limit,
estimated_max_concurrency,
"hardware info collected"
);
HardwareInfo {
cpu_cores,
total_memory_mb,
os_info,
fd_limit,
estimated_max_concurrency,
}
}
/// Read the soft file-descriptor limit (RLIMIT_NOFILE).
fn get_fd_limit() -> u64 {
#[cfg(unix)]
{
let mut rlim = libc::rlimit {
rlim_cur: 0,
rlim_max: 0,
};
let ret = unsafe { libc::getrlimit(libc::RLIMIT_NOFILE, &mut rlim) };
if ret == 0 {
return rlim.rlim_cur;
}
}
// Fallback for non-unix or error
1024
}
fn ratio_percent(value: u64, total: u64) -> f64 {
if total == 0 {
0.0
} else {
value as f64 * 100.0 / total as f64
}
}
fn current_unix_secs() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|duration| duration.as_secs())
.unwrap_or(0)
}

View File

@@ -0,0 +1,189 @@
#![allow(clippy::large_enum_variant)]
mod app;
mod config;
mod egress_proxy;
mod hardware;
mod net;
mod registration;
mod runtime;
mod setup;
mod state;
mod target_filter;
mod tunnel;
mod upstream_client;
use std::path::PathBuf;
use clap::{CommandFactory, FromArgMatches, Parser};
use config::Config;
/// Default config file name.
const DEFAULT_CONFIG: &str = "aether-tunnel.toml";
const OUTBOUND_PROXY_ENV: &str = "AETHER_TUNNEL_AETHER_OUTBOUND_PROXY_URL";
const LEGACY_OUTBOUND_PROXY_ENV: &str = concat!("AETHER_TUNNEL_AETHER_", "PROXY_URL");
/// Build the full clap command: Config args + discoverable subcommands.
///
/// `subcommand_negates_reqs` lets subcommands bypass the required Config
/// flags so that e.g. `aether-tunnel setup` doesn't demand `--aether-url`.
fn build_command() -> clap::Command {
Config::command()
.subcommand(
clap::Command::new("setup")
.about("Interactive setup wizard (TUI)")
.arg(
clap::Arg::new("config_path")
.help("Path to config file")
.default_value(DEFAULT_CONFIG),
),
)
.subcommand(clap::Command::new("start").about("Start the installed service"))
.subcommand(clap::Command::new("status").about("Show service status"))
.subcommand(clap::Command::new("logs").about("Tail service logs"))
.subcommand(clap::Command::new("restart").about("Restart the installed service"))
.subcommand(clap::Command::new("stop").about("Stop the installed service"))
.subcommand(clap::Command::new("uninstall").about("Uninstall the installed service"))
.subcommand(
clap::Command::new("upgrade")
.about("Self-upgrade from GitHub releases")
.arg(clap::Arg::new("version").help("Target version (e.g. 0.2.0)")),
)
.subcommand_negates_reqs(true)
}
#[tokio::main]
async fn main() -> anyhow::Result<()> {
rustls::crypto::ring::default_provider()
.install_default()
.map_err(|_| anyhow::anyhow!("Failed to install rustls CryptoProvider"))?;
promote_legacy_env_overrides();
// Load config file as env-var defaults (before clap parsing)
let config_file_path =
std::env::var("AETHER_TUNNEL_CONFIG").unwrap_or_else(|_| DEFAULT_CONFIG.to_string());
let config_path = std::path::Path::new(&config_file_path);
if config_path.exists() {
match config::ConfigFile::load(config_path) {
Ok(file_cfg) => file_cfg.inject_env(),
Err(error) => {
eprintln!(
" WARNING: failed to load config {}: {}",
config_path.display(),
error
);
}
}
}
// Parse CLI (subcommands + config args in one pass)
match build_command().try_get_matches() {
Ok(matches) => match matches.subcommand() {
Some(("setup", sub_m)) => {
let path = sub_m
.get_one::<String>("config_path")
.map(PathBuf::from)
.unwrap_or_else(|| PathBuf::from(DEFAULT_CONFIG));
handle_setup_result(setup::run(path)?).await
}
Some(("start", _)) => setup::service::cmd_start(),
Some(("status", _)) => setup::service::cmd_status(),
Some(("logs", _)) => setup::service::cmd_logs(),
Some(("restart", _)) => setup::service::cmd_restart(),
Some(("stop", _)) => setup::service::cmd_stop(),
Some(("uninstall", _)) => setup::service::cmd_uninstall(),
Some(("upgrade", sub_m)) => {
let version = sub_m.get_one::<String>("version").cloned();
setup::upgrade::cmd_upgrade(version).await
}
Some(_) => unreachable!(),
None => {
// No subcommand: run the tunnel with parsed config.
let config = Config::from_arg_matches(&matches)?;
run_tunnel(config).await
}
},
Err(e) => {
if e.kind() == clap::error::ErrorKind::MissingRequiredArgument {
eprintln!("Missing required config, launching setup wizard...\n");
handle_setup_result(setup::run(PathBuf::from(&config_file_path))?).await
} else {
e.exit();
}
}
}
}
fn promote_legacy_env_overrides() {
if std::env::var_os(OUTBOUND_PROXY_ENV).is_none() {
if let Some(value) = std::env::var_os(LEGACY_OUTBOUND_PROXY_ENV) {
std::env::set_var(OUTBOUND_PROXY_ENV, value);
}
}
}
/// Decide what to do after the setup wizard completes.
async fn handle_setup_result(outcome: setup::SetupOutcome) -> anyhow::Result<()> {
match outcome {
setup::SetupOutcome::ServiceInstalled => Ok(()),
setup::SetupOutcome::ReadyToRun(config_path) => {
// Reload config from the file that setup just wrote, overriding
// any stale env vars from a previous config.
match config::ConfigFile::load(&config_path) {
Ok(file_cfg) => file_cfg.inject_env_override(),
Err(e) => anyhow::bail!("failed to reload config after setup: {}", e),
}
// Parse from env-only (argv may still contain "setup" etc.)
let config = Config::try_parse_from(["aether-tunnel"])
.map_err(|e| anyhow::anyhow!("config invalid after setup: {}", e))?;
eprintln!(" Starting tunnel...\n");
run_tunnel(config).await
}
setup::SetupOutcome::Cancelled => {
eprintln!(" Setup cancelled.");
Ok(())
}
}
}
/// Start the tunnel agent, checking for managed-service conflicts first.
async fn run_tunnel(config: Config) -> anyhow::Result<()> {
// Warn if a managed service is already running (would cause conflicts).
if std::env::var_os("AETHER_TUNNEL_SERVICE_MANAGER").is_none()
&& std::env::var_os("INVOCATION_ID").is_none()
&& setup::service::is_service_active()
{
eprintln!(
"Warning: {} service is already running.",
setup::service::preferred_manager_name()
);
eprintln!("Use `./aether-tunnel stop` to stop it first, or manage via subcommands:");
eprintln!(" ./aether-tunnel status / logs / restart / stop");
std::process::exit(1);
}
// Resolve server list: if a config file exists, it must use [[servers]].
// Otherwise fall back to CLI/env single-server mode.
let config_path =
std::env::var("AETHER_TUNNEL_CONFIG").unwrap_or_else(|_| DEFAULT_CONFIG.to_string());
let servers = if std::path::Path::new(&config_path).exists() {
let file_cfg = config::ConfigFile::load(std::path::Path::new(&config_path))?;
if file_cfg.servers.is_empty() {
anyhow::bail!(
"config file {} must contain at least one [[servers]] entry",
config_path
);
}
file_cfg.servers.clone()
} else {
vec![config::ServerEntry {
aether_url: config.aether_url.clone(),
management_token: config.management_token.clone(),
node_name: None,
}]
};
app::run(config, servers).await
}

View File

@@ -0,0 +1,90 @@
//! Network utility functions (public IP detection, region detection).
//!
//! These are standalone helpers not tied to any specific client or service.
use aether_http::{build_http_client, HttpClientConfig};
use tracing::{debug, info};
/// Auto-detect public IP by querying external services.
pub async fn detect_public_ip() -> anyhow::Result<String> {
let endpoints = [
"https://api.ipify.org",
"https://ifconfig.me/ip",
"https://icanhazip.com",
];
let client = build_http_client(&HttpClientConfig {
request_timeout_ms: Some(5_000),
user_agent: Some("aether-tunnel/net".to_string()),
..HttpClientConfig::default()
})?;
for endpoint in &endpoints {
match client.get(*endpoint).send().await {
Ok(resp) if resp.status().is_success() => {
let ip = resp.text().await?.trim().to_string();
if !ip.is_empty() {
info!(ip = %ip, source = %endpoint, "detected public IP");
return Ok(ip);
}
}
Ok(resp) => {
debug!(endpoint = %endpoint, status = %resp.status(), "IP detection failed");
}
Err(e) => {
debug!(endpoint = %endpoint, error = %e, "IP detection failed");
}
}
}
anyhow::bail!("failed to detect public IP from any source; use --public-ip")
}
/// Auto-detect geographic region from a public IP address.
///
/// Uses multiple providers with HTTPS preferred. Falls back to ip-api.com
/// over plain HTTP (their free tier doesn't support HTTPS).
/// This is best-effort and non-sensitive -- region detection should never
/// block startup.
pub async fn detect_region(ip: &str) -> Option<String> {
// Try HTTPS provider first
let https_url = format!("https://ipinfo.io/{}/country", ip);
let client = build_http_client(&HttpClientConfig {
request_timeout_ms: Some(5_000),
user_agent: Some("aether-tunnel/net".to_string()),
..HttpClientConfig::default()
})
.ok()?;
// Try ipinfo.io (HTTPS, returns plain text country code)
if let Ok(resp) = client.get(&https_url).send().await {
if resp.status().is_success() {
if let Ok(text) = resp.text().await {
let code = text.trim();
if !code.is_empty() && code.len() <= 3 {
info!(region = %code, ip = %ip, source = "ipinfo.io", "detected region");
return Some(code.to_string());
}
}
}
}
// Fallback: ip-api.com (HTTP only on free tier, non-sensitive data)
let http_url = format!("http://ip-api.com/json/{}?fields=countryCode", ip);
match client.get(&http_url).send().await {
Ok(resp) if resp.status().is_success() => {
let body: serde_json::Value = resp.json().await.ok()?;
let code = body.get("countryCode")?.as_str()?;
if code.is_empty() {
return None;
}
info!(region = %code, ip = %ip, source = "ip-api.com", "detected region");
Some(code.to_string())
}
_ => {
debug!(ip = %ip, "region detection failed");
None
}
}
}

View File

@@ -0,0 +1,257 @@
use aether_http::{build_http_client, jittered_delay_for_retry, HttpClientConfig, HttpRetryConfig};
use aether_runtime::summarize_text_payload;
use reqwest::{Client, StatusCode};
use serde::{Deserialize, Serialize};
use tokio::time::sleep;
use tracing::{debug, error, info};
use crate::config::Config;
use crate::hardware::HardwareInfo;
#[derive(Debug, Serialize)]
struct RegisterRequest {
name: String,
ip: String,
port: u16,
#[serde(skip_serializing_if = "Option::is_none")]
region: Option<String>,
heartbeat_interval: u64,
#[serde(skip_serializing_if = "Option::is_none")]
hardware_info: Option<serde_json::Value>,
#[serde(skip_serializing_if = "Option::is_none")]
estimated_max_concurrency: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
proxy_metadata: Option<serde_json::Value>,
tunnel_mode: bool,
}
#[derive(Debug, Deserialize)]
pub struct RegisterResponse {
pub node_id: String,
}
/// Remote configuration pushed by the Aether management backend.
#[derive(Debug, Clone, Deserialize)]
pub struct RemoteConfig {
pub node_name: Option<String>,
pub allowed_ports: Option<Vec<u16>>,
pub log_level: Option<String>,
pub heartbeat_interval: Option<u64>,
}
#[derive(Debug, Serialize)]
struct UnregisterRequest {
node_id: String,
}
/// Aether API client for tunnel node lifecycle management.
pub struct AetherClient {
http: Client,
base_url: String,
token: String,
retry: HttpRetryConfig,
}
impl AetherClient {
pub fn new(config: &Config, aether_url: &str, management_token: &str) -> Self {
let http = build_http_client(&HttpClientConfig {
connect_timeout_ms: Some(config.aether_connect_timeout_secs.saturating_mul(1_000)),
request_timeout_ms: Some(config.aether_request_timeout_secs.saturating_mul(1_000)),
pool_idle_timeout_ms: Some(config.aether_pool_idle_timeout_secs.saturating_mul(1_000)),
pool_max_idle_per_host: Some(config.aether_pool_max_idle_per_host),
tcp_keepalive_ms: if config.aether_tcp_keepalive_secs > 0 {
Some(config.aether_tcp_keepalive_secs.saturating_mul(1_000))
} else {
None
},
tcp_nodelay: config.aether_tcp_nodelay,
http2_adaptive_window: config.aether_http2,
user_agent: Some(format!("aether-tunnel/{}", env!("CARGO_PKG_VERSION"))),
proxy_url: config
.effective_aether_outbound_proxy_url()
.map(str::to_string),
..HttpClientConfig::default()
})
.expect("failed to create HTTP client");
let retry = HttpRetryConfig {
max_attempts: config.aether_retry_max_attempts,
base_delay_ms: config.aether_retry_base_delay_ms,
max_delay_ms: config.aether_retry_max_delay_ms,
}
.normalized();
Self {
http,
base_url: aether_url.trim_end_matches('/').to_string(),
token: management_token.to_string(),
retry,
}
}
/// Register this node with Aether (idempotent upsert by ip:port).
///
/// Returns the stable node_id assigned by Aether.
pub async fn register(
&self,
config: &Config,
node_name: &str,
public_ip: &str,
hw: Option<&HardwareInfo>,
) -> anyhow::Result<String> {
let url = format!("{}/api/admin/proxy-nodes/register", self.base_url);
let body = RegisterRequest {
name: node_name.to_string(),
ip: public_ip.to_string(),
port: 0,
region: config.node_region.clone(),
heartbeat_interval: config.heartbeat_interval,
hardware_info: hw.and_then(|h| serde_json::to_value(h).ok()),
estimated_max_concurrency: hw.map(|h| h.estimated_max_concurrency),
proxy_metadata: Some(serde_json::json!({
"version": env!("CARGO_PKG_VERSION"),
})),
tunnel_mode: true,
};
info!(
url = %url,
name = %body.name,
ip = %body.ip,
"registering with Aether"
);
let resp = self
.send_with_retry(
|| {
self.http
.post(&url)
.header("Authorization", format!("Bearer {}", self.token))
.json(&body)
},
"register",
)
.await?;
let status = resp.status();
if !status.is_success() {
let text = resp.text().await.unwrap_or_default();
let summary = summarize_text_payload(&text);
anyhow::bail!(
"register failed (HTTP {}): response body redacted (bytes={}, sha256={})",
status,
summary.bytes,
summary.sha256
);
}
let data: RegisterResponse = resp.json().await?;
info!(node_id = %data.node_id, "registered successfully");
Ok(data.node_id)
}
/// Unregister this node from Aether (graceful shutdown).
pub async fn unregister(&self, node_id: &str) -> anyhow::Result<()> {
let url = format!("{}/api/admin/proxy-nodes/unregister", self.base_url);
let body = UnregisterRequest {
node_id: node_id.to_string(),
};
info!(node_id = %node_id, "unregistering from Aether");
let resp = self
.send_with_retry(
|| {
self.http
.post(&url)
.header("Authorization", format!("Bearer {}", self.token))
.json(&body)
},
"unregister",
)
.await;
match resp {
Ok(r) if r.status().is_success() => {
info!(node_id = %node_id, "unregistered successfully");
Ok(())
}
Ok(r) => {
let status = r.status();
let text = r.text().await.unwrap_or_default();
let summary = summarize_text_payload(&text);
error!(
status = %status,
body_bytes = summary.bytes,
body_sha256 = %summary.sha256,
"unregister failed"
);
anyhow::bail!(
"unregister failed (HTTP {}): response body redacted (bytes={}, sha256={})",
status,
summary.bytes,
summary.sha256
);
}
Err(e) => {
// Best-effort during shutdown
error!(error = %e, "unregister request failed");
anyhow::bail!("unregister request failed: {}", e);
}
}
}
async fn send_with_retry<F>(
&self,
mut make_req: F,
label: &str,
) -> Result<reqwest::Response, reqwest::Error>
where
F: FnMut() -> reqwest::RequestBuilder,
{
let mut attempt: u32 = 0;
loop {
attempt = attempt.saturating_add(1);
let resp = make_req().send().await;
match resp {
Ok(resp) => {
if should_retry_status(resp.status()) && attempt < self.retry.max_attempts {
let sleep_for = jittered_delay_for_retry(self.retry, attempt - 1);
debug!(
attempt,
status = %resp.status(),
sleep_ms = sleep_for.as_millis(),
label,
"Aether request retrying"
);
sleep(sleep_for).await;
continue;
}
return Ok(resp);
}
Err(e) => {
if attempt < self.retry.max_attempts {
let sleep_for = jittered_delay_for_retry(self.retry, attempt - 1);
debug!(
attempt,
error = %e,
sleep_ms = sleep_for.as_millis(),
label,
"Aether request retrying"
);
sleep(sleep_for).await;
continue;
}
return Err(e);
}
}
}
}
}
fn should_retry_status(status: StatusCode) -> bool {
status.is_server_error()
|| status == StatusCode::TOO_MANY_REQUESTS
|| status == StatusCode::REQUEST_TIMEOUT
}

View File

@@ -0,0 +1 @@
pub mod client;

View File

@@ -0,0 +1,121 @@
//! Runtime-mutable configuration that can be updated remotely via heartbeat.
//!
//! Fields in [`DynamicConfig`] are initially populated from the static
//! [`Config`](crate::config::Config) and may be overridden by the Aether
//! management backend through the heartbeat response.
use std::collections::HashSet;
use std::sync::{Arc, OnceLock};
use arc_swap::ArcSwap;
use tracing::info;
use crate::config::Config;
/// Configuration that can be changed at runtime without restart.
#[derive(Debug, Clone)]
pub struct DynamicConfig {
pub node_name: String,
pub allowed_ports: Arc<HashSet<u16>>,
pub log_level: String,
pub heartbeat_interval: u64,
/// Monotonically increasing version from the backend.
/// `0` means no remote config has ever been applied.
pub config_version: u64,
}
impl DynamicConfig {
/// Initialize from static config (startup defaults).
pub fn from_config(config: &Config) -> Self {
Self {
node_name: config.node_name.clone(),
allowed_ports: Arc::new(config.allowed_ports.iter().copied().collect()),
log_level: config.log_level.clone(),
heartbeat_interval: config.heartbeat_interval,
config_version: 0,
}
}
}
/// Shared dynamic config handle (lock-free reads via ArcSwap).
pub type SharedDynamicConfig = Arc<ArcSwap<DynamicConfig>>;
// -- Log-level hot-reload -----
/// Global log-level reloader function, set during tracing init.
type LogReloader = Box<dyn Fn(&str) + Send + Sync>;
static LOG_RELOADER: OnceLock<LogReloader> = OnceLock::new();
/// Register the log-level reload function (called once from `init_tracing`).
pub fn set_log_reloader(f: LogReloader) {
let _ = LOG_RELOADER.set(f);
}
/// Apply a remote config update to the dynamic config.
///
/// Uses copy-on-write: loads the current snapshot, clones it, applies changes,
/// and stores the new Arc. Reads are always lock-free.
///
/// Returns `true` if the config was actually changed.
pub fn apply_remote_config(
dynamic: &SharedDynamicConfig,
remote: &crate::registration::client::RemoteConfig,
version: u64,
) -> bool {
let current = dynamic.load();
if version <= current.config_version {
return false;
}
let mut new_cfg = (**current).clone();
let mut changed = Vec::new();
if let Some(ref name) = remote.node_name {
if *name != new_cfg.node_name {
changed.push(format!("node_name -> {}", name));
new_cfg.node_name = name.clone();
}
}
if let Some(ref ports) = remote.allowed_ports {
let new_set: HashSet<u16> = ports.iter().copied().collect();
if new_set != *new_cfg.allowed_ports {
changed.push(format!("allowed_ports -> {:?}", ports));
new_cfg.allowed_ports = Arc::new(new_set);
}
}
if let Some(interval) = remote.heartbeat_interval {
if interval != new_cfg.heartbeat_interval {
changed.push(format!("heartbeat_interval -> {}s", interval));
new_cfg.heartbeat_interval = interval;
}
}
if let Some(ref level) = remote.log_level {
if *level != new_cfg.log_level {
changed.push(format!("log_level -> {}", level));
new_cfg.log_level = level.clone();
// Hot-reload tracing filter
if let Some(reloader) = LOG_RELOADER.get() {
reloader(level);
}
}
}
let has_changes = !changed.is_empty();
if has_changes {
new_cfg.config_version = version;
info!(
version,
changes = %changed.join(", "),
"remote config applied"
);
dynamic.store(Arc::new(new_cfg));
}
has_changes
}

View File

@@ -0,0 +1,66 @@
//! Safe DNS resolver for reqwest that reuses validated addresses from DnsCache.
//!
//! This resolver ensures reqwest connects only to addresses that have been
//! previously validated by `target_filter::validate_target()`, eliminating
//! the TOCTTOU gap where DNS rebinding could redirect traffic to private IPs.
use std::net::SocketAddr;
use std::sync::Arc;
use reqwest::dns::{Addrs, Name, Resolve, Resolving};
use crate::target_filter::{self, DnsCache};
/// A DNS resolver that serves validated public addresses from the shared DnsCache.
///
/// When reqwest needs to resolve a hostname, this resolver returns addresses
/// from the cache (populated by `validate_target()` during request validation).
/// If the hostname is not in cache (shouldn't happen in normal flow), it
/// performs a fresh resolution with private-IP filtering.
pub struct SafeDnsResolver {
dns_cache: Arc<DnsCache>,
}
impl SafeDnsResolver {
pub fn new(dns_cache: Arc<DnsCache>) -> Self {
Self { dns_cache }
}
}
impl Resolve for SafeDnsResolver {
fn resolve(&self, name: Name) -> Resolving {
let dns_cache = Arc::clone(&self.dns_cache);
Box::pin(async move {
let host = name.as_str();
// Try cache first (should be populated by validate_target).
// reqwest resolves by hostname only (no port), so use host-only lookup.
if let Some(addrs) = dns_cache.get_by_host(host).await {
let socket_addrs: Vec<SocketAddr> = (*addrs).clone();
return Ok(Box::new(socket_addrs.into_iter()) as Addrs);
}
// Fallback: resolve with private-IP filtering (defensive).
// This path should rarely be hit since validate_target() runs first.
// We don't know the real port here (reqwest Resolve only gives hostname),
// so resolve directly without caching to avoid polluting the cache with
// an incorrect port-based key.
let addr_str = format!("{}:0", host);
let resolved: Vec<SocketAddr> = tokio::net::lookup_host(&addr_str)
.await
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> { Box::new(e) })?
.filter(|addr| !target_filter::is_private_ip(&addr.ip()))
.collect();
if resolved.is_empty() {
return Err(Box::new(std::io::Error::other(format!(
"all resolved addresses for {} are private/reserved",
host
)))
as Box<dyn std::error::Error + Send + Sync>);
}
Ok(Box::new(resolved.into_iter()) as Addrs)
})
}
}

View File

@@ -0,0 +1,5 @@
pub(crate) mod service;
mod tui;
pub(crate) mod upgrade;
pub use self::tui::{run, SetupOutcome};

View File

@@ -0,0 +1,637 @@
//! Service installation and management for `aether-tunnel`.
//!
//! Supports the host-native service manager we currently target:
//! `systemd` on most Linux distributions and `OpenRC` on Alpine.
use std::fs::OpenOptions;
use std::io::ErrorKind;
use std::path::Path;
use std::process::{Command, ExitStatus, Stdio};
const SERVICE_NAME: &str = "aether-tunnel";
const SYSTEMD_UNIT_PATH: &str = "/etc/systemd/system/aether-tunnel.service";
const OPENRC_INIT_PATH: &str = "/etc/init.d/aether-tunnel";
const OPENRC_PID_PATH: &str = "/run/aether-tunnel.pid";
const OPENRC_LOG_DIR: &str = "/var/log/aether-tunnel";
const OPENRC_STDOUT_LOG: &str = "/var/log/aether-tunnel/current.log";
const OPENRC_STDERR_LOG: &str = "/var/log/aether-tunnel/error.log";
const OPENRC_RUN_BINS: &[&str] = &["/sbin/openrc-run", "/usr/sbin/openrc-run", "openrc-run"];
const OPENRC_SERVICE_BINS: &[&str] = &["/sbin/rc-service", "/usr/sbin/rc-service", "rc-service"];
const OPENRC_UPDATE_BINS: &[&str] = &["/sbin/rc-update", "/usr/sbin/rc-update", "rc-update"];
const OPENRC_SUPERVISE_BINS: &[&str] = &[
"/sbin/supervise-daemon",
"/usr/sbin/supervise-daemon",
"supervise-daemon",
];
const TAIL_BINS: &[&str] = &["/usr/bin/tail", "/bin/tail", "tail"];
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum ServiceManager {
Systemd,
OpenRc,
}
impl ServiceManager {
fn display_name(self) -> &'static str {
match self {
Self::Systemd => "systemd",
Self::OpenRc => "OpenRC",
}
}
fn unit_path(self) -> &'static str {
match self {
Self::Systemd => SYSTEMD_UNIT_PATH,
Self::OpenRc => OPENRC_INIT_PATH,
}
}
fn is_installed(self) -> bool {
Path::new(self.unit_path()).exists()
}
}
pub fn is_available() -> bool {
detect_service_manager().is_some() && is_root()
}
pub fn preferred_manager_name() -> &'static str {
installed_manager()
.or_else(detect_service_manager)
.map(ServiceManager::display_name)
.unwrap_or("service")
}
pub fn unavailable_hint() -> String {
match detect_service_manager() {
Some(manager) if !is_root() => {
format!(
"requires root with {}, use: sudo aether-tunnel setup",
manager.display_name()
)
}
Some(manager) => format!(
"{} is available but service setup is not ready",
manager.display_name()
),
None => "no supported service manager detected (systemd/OpenRC)".into(),
}
}
pub fn install_service(config_path: &Path) -> anyhow::Result<()> {
let manager = detect_service_manager()
.ok_or_else(|| anyhow::anyhow!("no supported service manager detected (systemd/OpenRC)"))?;
if !is_root() {
anyhow::bail!("root required, use: sudo ./aether-tunnel setup");
}
match manager {
ServiceManager::Systemd => install_systemd_service(config_path),
ServiceManager::OpenRc => install_openrc_service(config_path),
}
}
pub(crate) fn is_root() -> bool {
#[cfg(unix)]
{
unsafe { libc::geteuid() == 0 }
}
#[cfg(not(unix))]
{
false
}
}
pub fn is_installed() -> bool {
installed_manager().is_some()
}
pub fn is_service_active() -> bool {
active_service_manager().is_some()
}
pub fn restart_active_service() -> anyhow::Result<()> {
let manager =
active_service_manager().ok_or_else(|| anyhow::anyhow!("no active service detected"))?;
restart_manager(manager)
}
pub fn uninstall_service() -> anyhow::Result<()> {
let Some(manager) = installed_manager() else {
return Ok(());
};
match manager {
ServiceManager::Systemd => uninstall_systemd_service(),
ServiceManager::OpenRc => uninstall_openrc_service(),
}
}
pub fn cmd_status() -> anyhow::Result<()> {
let manager = ensure_service_installed()?;
let status = manager_status(manager)?;
std::process::exit(status.code().unwrap_or(1));
}
pub fn cmd_logs() -> anyhow::Result<()> {
let manager = ensure_service_installed()?;
if manager == ServiceManager::OpenRc {
ensure_openrc_logs_readable()?;
}
let status = match manager {
ServiceManager::Systemd => Command::new("journalctl")
.args(["-u", SERVICE_NAME, "-f", "--no-pager", "-n", "100"])
.status()?,
ServiceManager::OpenRc => Command::new(tail_bin())
.args(["-n", "100", "-f", OPENRC_STDOUT_LOG, OPENRC_STDERR_LOG])
.status()?,
};
std::process::exit(status.code().unwrap_or(1));
}
pub fn cmd_start() -> anyhow::Result<()> {
let manager = ensure_root_and_service()?;
start_manager(manager)?;
eprintln!(" Service started.");
Ok(())
}
pub fn cmd_restart() -> anyhow::Result<()> {
let manager = ensure_root_and_service()?;
restart_manager(manager)?;
eprintln!(" Service restarted.");
Ok(())
}
pub fn cmd_stop() -> anyhow::Result<()> {
let manager = ensure_root_and_service()?;
stop_manager(manager)?;
eprintln!(" Service stopped.");
Ok(())
}
pub fn cmd_uninstall() -> anyhow::Result<()> {
ensure_root_and_service()?;
uninstall_service()?;
eprintln!();
eprintln!(" Config file, TLS certs, and logs are preserved. Remove manually if needed.");
Ok(())
}
pub(crate) fn run_cmd(program: &str, args: &[&str]) -> anyhow::Result<()> {
let display = format!("{} {}", program, args.join(" "));
eprintln!(" > {}", display);
let status = Command::new(program).args(args).status()?;
if !status.success() {
anyhow::bail!("command failed: {}", display);
}
Ok(())
}
fn detect_service_manager() -> Option<ServiceManager> {
if is_systemd_available() {
Some(ServiceManager::Systemd)
} else if is_openrc_available() {
Some(ServiceManager::OpenRc)
} else {
None
}
}
fn installed_manager() -> Option<ServiceManager> {
if let Some(manager) = detect_service_manager() {
if manager.is_installed() {
return Some(manager);
}
}
[ServiceManager::Systemd, ServiceManager::OpenRc]
.into_iter()
.find(|manager| manager.is_installed())
}
fn ensure_openrc_logs_readable() -> anyhow::Result<()> {
for path in [OPENRC_STDOUT_LOG, OPENRC_STDERR_LOG] {
match std::fs::File::open(path) {
Ok(_) => {}
Err(err) if err.kind() == ErrorKind::PermissionDenied => {
anyhow::bail!(
"OpenRC logs are stored under {} and usually require root access. Try `sudo ./aether-tunnel logs`.",
OPENRC_LOG_DIR
);
}
Err(err) if err.kind() == ErrorKind::NotFound => {
anyhow::bail!(
"OpenRC log file not found at {}. Start the service first or check `./aether-tunnel status`.",
path
);
}
Err(err) => return Err(err.into()),
}
}
Ok(())
}
fn active_service_manager() -> Option<ServiceManager> {
if let Some(manager) = installed_manager() {
if manager_is_active(manager) {
return Some(manager);
}
}
[ServiceManager::Systemd, ServiceManager::OpenRc]
.into_iter()
.find(|manager| manager_is_active(*manager))
}
fn ensure_service_installed() -> anyhow::Result<ServiceManager> {
installed_manager().ok_or_else(|| {
anyhow::anyhow!("service not installed, run `sudo ./aether-tunnel setup` first")
})
}
fn ensure_root_and_service() -> anyhow::Result<ServiceManager> {
let manager = ensure_service_installed()?;
if !is_root() {
anyhow::bail!("root required, use: sudo ./aether-tunnel <command>");
}
Ok(manager)
}
fn install_systemd_service(config_path: &Path) -> anyhow::Result<()> {
let exe_path = std::env::current_exe()?.canonicalize()?;
let exe_str = exe_path
.to_str()
.ok_or_else(|| anyhow::anyhow!("binary path contains invalid UTF-8"))?;
let config_abs = std::fs::canonicalize(config_path)?;
let config_str = config_abs
.to_str()
.ok_or_else(|| anyhow::anyhow!("config path contains invalid UTF-8"))?;
let working_dir = config_abs
.parent()
.unwrap_or_else(|| Path::new("/"))
.to_str()
.unwrap_or("/");
if Path::new(SYSTEMD_UNIT_PATH).exists() {
eprintln!(" Stopping existing service...");
let _ = Command::new("systemctl")
.args(["stop", SERVICE_NAME])
.status();
}
eprintln!(" Generating systemd unit file...");
eprintln!(" Binary: {}", exe_str);
eprintln!(" Config: {}", config_str);
eprintln!(" WorkDir: {}", working_dir);
let unit_content = format!(
"[Unit]\n\
Description=Aether Tunnel\n\
After=network.target\n\
\n\
[Service]\n\
Type=simple\n\
WorkingDirectory={working_dir}\n\
Environment=AETHER_TUNNEL_CONFIG={config_str}\n\
Environment=AETHER_TUNNEL_SERVICE_MANAGER=systemd\n\
Environment=AETHER_TUNNEL_LOG_DESTINATION=both\n\
Environment=AETHER_TUNNEL_LOG_DIR=/var/log/aether-tunnel\n\
ExecStart={exe_str}\n\
Restart=on-failure\n\
RestartSec=5\n\
LimitNOFILE=65535\n\
UMask=0077\n\
LogsDirectory=aether-tunnel\n\
LogsDirectoryMode=0750\n\
\n\
[Install]\n\
WantedBy=multi-user.target\n",
);
std::fs::write(SYSTEMD_UNIT_PATH, &unit_content)?;
eprintln!(" Enabling and starting service...");
run_cmd("systemctl", &["daemon-reload"])?;
run_cmd("systemctl", &["enable", "--now", SERVICE_NAME])?;
eprintln!();
if manager_is_active(ServiceManager::Systemd) {
eprintln!(" Service started successfully!");
} else {
eprintln!(" Service state is not active yet. Check `sudo ./aether-tunnel logs`.");
}
print_post_install_commands();
Ok(())
}
fn install_openrc_service(config_path: &Path) -> anyhow::Result<()> {
let exe_path = std::env::current_exe()?.canonicalize()?;
let exe_str = exe_path
.to_str()
.ok_or_else(|| anyhow::anyhow!("binary path contains invalid UTF-8"))?;
let config_abs = std::fs::canonicalize(config_path)?;
let config_str = config_abs
.to_str()
.ok_or_else(|| anyhow::anyhow!("config path contains invalid UTF-8"))?;
let working_dir = config_abs
.parent()
.unwrap_or_else(|| Path::new("/"))
.to_str()
.unwrap_or("/");
if Path::new(OPENRC_INIT_PATH).exists() {
eprintln!(" Stopping existing service...");
let _ = Command::new(openrc_service_bin())
.args([SERVICE_NAME, "stop"])
.status();
}
std::fs::create_dir_all(OPENRC_LOG_DIR)?;
touch_log(OPENRC_STDOUT_LOG)?;
touch_log(OPENRC_STDERR_LOG)?;
set_mode(OPENRC_LOG_DIR, 0o750)?;
set_mode(OPENRC_STDOUT_LOG, 0o640)?;
set_mode(OPENRC_STDERR_LOG, 0o640)?;
eprintln!(" Generating OpenRC init script...");
eprintln!(" Binary: {}", exe_str);
eprintln!(" Config: {}", config_str);
eprintln!(" WorkDir: {}", working_dir);
let init_content = format!(
r#"#!{}
name={}
description={}
supervisor=supervise-daemon
command={}
directory={}
pidfile={}
output_log_dir={}
output_log={}
error_log={}
supervise_daemon={}
config_env={}
service_manager_env={}
log_destination_env={}
log_dir_env={}
respawn_delay=5
respawn_max=10
respawn_period=60
depend() {{
after net
}}
start_pre() {{
checkpath --directory --mode 0750 "$output_log_dir"
checkpath --file --mode 0640 "$output_log"
checkpath --file --mode 0640 "$error_log"
}}
start() {{
ebegin "Starting ${{RC_SVCNAME}}"
"$supervise_daemon" "${{RC_SVCNAME}}" \
--start "$command" \
--pidfile "$pidfile" \
--chdir "$directory" \
--stdout "$output_log" \
--stderr "$error_log" \
--respawn-delay "$respawn_delay" \
--respawn-max "$respawn_max" \
--respawn-period "$respawn_period" \
--umask 0077 \
--env "$config_env" \
--env "$service_manager_env" \
--env "$log_destination_env" \
--env "$log_dir_env"
eend $?
}}
stop() {{
ebegin "Stopping ${{RC_SVCNAME}}"
"$supervise_daemon" "${{RC_SVCNAME}}" --stop "$command" --pidfile "$pidfile"
eend $?
}}
"#,
openrc_run_bin(),
shell_quote(SERVICE_NAME),
shell_quote("Aether Tunnel"),
shell_quote(exe_str),
shell_quote(working_dir),
shell_quote(OPENRC_PID_PATH),
shell_quote(OPENRC_LOG_DIR),
shell_quote(OPENRC_STDOUT_LOG),
shell_quote(OPENRC_STDERR_LOG),
shell_quote(supervise_daemon_bin()),
shell_quote(&format!("AETHER_TUNNEL_CONFIG={config_str}")),
shell_quote("AETHER_TUNNEL_SERVICE_MANAGER=openrc"),
shell_quote("AETHER_TUNNEL_LOG_DESTINATION=both"),
shell_quote(&format!("AETHER_TUNNEL_LOG_DIR={OPENRC_LOG_DIR}")),
);
std::fs::write(OPENRC_INIT_PATH, &init_content)?;
set_mode(OPENRC_INIT_PATH, 0o755)?;
eprintln!(" Enabling and starting service...");
run_cmd(openrc_update_bin(), &["add", SERVICE_NAME, "default"])?;
run_cmd(openrc_service_bin(), &[SERVICE_NAME, "start"])?;
eprintln!();
if manager_is_active(ServiceManager::OpenRc) {
eprintln!(" Service started successfully!");
} else {
eprintln!(" Service state is not active yet. Check `sudo ./aether-tunnel logs`.");
}
print_post_install_commands();
Ok(())
}
fn uninstall_systemd_service() -> anyhow::Result<()> {
eprintln!(" Stopping and removing existing service...");
let _ = Command::new("systemctl")
.args(["disable", "--now", SERVICE_NAME])
.status();
if Path::new(SYSTEMD_UNIT_PATH).exists() {
std::fs::remove_file(SYSTEMD_UNIT_PATH)?;
eprintln!(" Removed {}", SYSTEMD_UNIT_PATH);
}
run_cmd("systemctl", &["daemon-reload"])?;
eprintln!(" Service uninstalled.");
Ok(())
}
fn uninstall_openrc_service() -> anyhow::Result<()> {
eprintln!(" Stopping and removing existing service...");
let _ = Command::new(openrc_service_bin())
.args([SERVICE_NAME, "stop"])
.status();
let _ = Command::new(openrc_update_bin())
.args(["del", SERVICE_NAME, "default"])
.status();
if Path::new(OPENRC_INIT_PATH).exists() {
std::fs::remove_file(OPENRC_INIT_PATH)?;
eprintln!(" Removed {}", OPENRC_INIT_PATH);
}
eprintln!(" Service uninstalled.");
Ok(())
}
fn start_manager(manager: ServiceManager) -> anyhow::Result<()> {
match manager {
ServiceManager::Systemd => run_cmd("systemctl", &["start", SERVICE_NAME]),
ServiceManager::OpenRc => run_cmd(openrc_service_bin(), &[SERVICE_NAME, "start"]),
}
}
fn stop_manager(manager: ServiceManager) -> anyhow::Result<()> {
match manager {
ServiceManager::Systemd => run_cmd("systemctl", &["stop", SERVICE_NAME]),
ServiceManager::OpenRc => run_cmd(openrc_service_bin(), &[SERVICE_NAME, "stop"]),
}
}
fn restart_manager(manager: ServiceManager) -> anyhow::Result<()> {
match manager {
ServiceManager::Systemd => run_cmd("systemctl", &["restart", SERVICE_NAME]),
ServiceManager::OpenRc => run_cmd(openrc_service_bin(), &[SERVICE_NAME, "restart"]),
}
}
fn manager_status(manager: ServiceManager) -> anyhow::Result<ExitStatus> {
let status = match manager {
ServiceManager::Systemd => Command::new("systemctl")
.args(["status", SERVICE_NAME])
.status()?,
ServiceManager::OpenRc => Command::new(openrc_service_bin())
.args([SERVICE_NAME, "status"])
.status()?,
};
Ok(status)
}
fn manager_is_active(manager: ServiceManager) -> bool {
match manager {
ServiceManager::Systemd => {
Path::new(SYSTEMD_UNIT_PATH).exists()
&& Command::new("systemctl")
.args(["is-active", "--quiet", SERVICE_NAME])
.stdout(Stdio::null())
.stderr(Stdio::null())
.status()
.map(|status| status.success())
.unwrap_or(false)
}
ServiceManager::OpenRc => {
Path::new(OPENRC_INIT_PATH).exists()
&& Command::new(openrc_service_bin())
.args([SERVICE_NAME, "status"])
.stdout(Stdio::null())
.stderr(Stdio::null())
.status()
.map(|status| status.success())
.unwrap_or(false)
}
}
}
fn print_post_install_commands() {
eprintln!();
eprintln!(" Commands:");
eprintln!(" ./aether-tunnel status # service status");
eprintln!(" sudo ./aether-tunnel logs # tail logs");
eprintln!(" sudo ./aether-tunnel restart # restart");
eprintln!(" sudo ./aether-tunnel stop # stop");
eprintln!(" sudo ./aether-tunnel uninstall # remove service");
eprintln!();
}
fn is_systemd_available() -> bool {
Path::new("/run/systemd/system").exists()
&& Command::new("systemctl")
.arg("--version")
.stdout(Stdio::null())
.stderr(Stdio::null())
.status()
.map(|status| status.success())
.unwrap_or(false)
}
fn is_openrc_available() -> bool {
(Path::new("/run/openrc").exists() || Path::new("/run/openrc/softlevel").exists())
&& has_absolute_candidate(OPENRC_RUN_BINS)
&& has_absolute_candidate(OPENRC_SERVICE_BINS)
&& has_absolute_candidate(OPENRC_UPDATE_BINS)
&& has_absolute_candidate(OPENRC_SUPERVISE_BINS)
}
fn has_absolute_candidate(candidates: &[&str]) -> bool {
candidates
.iter()
.any(|candidate| candidate.starts_with('/') && Path::new(candidate).exists())
}
fn openrc_run_bin() -> &'static str {
pick_bin(OPENRC_RUN_BINS)
}
fn openrc_service_bin() -> &'static str {
pick_bin(OPENRC_SERVICE_BINS)
}
fn openrc_update_bin() -> &'static str {
pick_bin(OPENRC_UPDATE_BINS)
}
fn supervise_daemon_bin() -> &'static str {
pick_bin(OPENRC_SUPERVISE_BINS)
}
fn tail_bin() -> &'static str {
pick_bin(TAIL_BINS)
}
fn pick_bin(candidates: &[&'static str]) -> &'static str {
candidates
.iter()
.copied()
.find(|candidate| candidate.starts_with('/') && Path::new(candidate).exists())
.unwrap_or_else(|| candidates[candidates.len() - 1])
}
fn shell_quote(value: &str) -> String {
format!("'{}'", value.replace('\'', "'\"'\"'"))
}
fn touch_log(path: &str) -> anyhow::Result<()> {
OpenOptions::new().create(true).append(true).open(path)?;
Ok(())
}
fn set_mode(path: &str, mode: u32) -> anyhow::Result<()> {
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
let perms = std::fs::Permissions::from_mode(mode);
std::fs::set_permissions(path, perms)?;
}
#[cfg(not(unix))]
let _ = (path, mode);
Ok(())
}

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,426 @@
//! Self-upgrade support for `aether-tunnel`.
//!
//! Downloads a release from GitHub, verifies the SHA256 checksum, replaces the
//! running binary atomically, and restarts the active managed service when
//! applicable.
use std::path::{Path, PathBuf};
use aether_http::{apply_http_client_config, HttpClientConfig};
use sha2::{Digest, Sha256};
const GITHUB_API_BASE: &str = "https://api.github.com";
const GITHUB_REPO: &str = "fawney19/Aether";
const CURRENT_VERSION: &str = env!("CARGO_PKG_VERSION");
// ── GitHub API types ─────────────────────────────────────────────────────────
#[derive(serde::Deserialize)]
struct GithubRelease {
tag_name: String,
name: String,
}
// ── Platform detection ───────────────────────────────────────────────────────
fn detect_platform() -> &'static str {
if cfg!(target_os = "linux") && cfg!(target_arch = "x86_64") && cfg!(target_env = "musl") {
"linux-musl-amd64"
} else if cfg!(target_os = "linux")
&& cfg!(target_arch = "aarch64")
&& cfg!(target_env = "musl")
{
"linux-musl-arm64"
} else if cfg!(target_os = "linux") && cfg!(target_arch = "x86_64") {
"linux-amd64"
} else if cfg!(target_os = "linux") && cfg!(target_arch = "aarch64") {
"linux-arm64"
} else if cfg!(target_os = "macos") && cfg!(target_arch = "x86_64") {
"macos-amd64"
} else if cfg!(target_os = "macos") && cfg!(target_arch = "aarch64") {
"macos-arm64"
} else if cfg!(target_os = "windows") && cfg!(target_arch = "x86_64") {
"windows-amd64"
} else {
// All supported targets are covered above; this is unreachable for
// any platform we actually build for.
panic!("unsupported platform: compile-time target not in the supported matrix")
}
}
// ── GitHub HTTP client ───────────────────────────────────────────────────────
fn build_github_client() -> anyhow::Result<reqwest::Client> {
let mut headers = reqwest::header::HeaderMap::new();
if let Ok(token) = std::env::var("GITHUB_TOKEN") {
headers.insert(
reqwest::header::AUTHORIZATION,
reqwest::header::HeaderValue::from_str(&format!("Bearer {}", token))?,
);
}
headers.insert(
reqwest::header::ACCEPT,
reqwest::header::HeaderValue::from_static("application/vnd.github+json"),
);
Ok(apply_http_client_config(
reqwest::Client::builder().default_headers(headers),
&HttpClientConfig {
request_timeout_ms: Some(300_000),
user_agent: Some(format!("aether-tunnel/{}", CURRENT_VERSION)),
..HttpClientConfig::default()
},
)
.build()?)
}
// ── Release fetching ─────────────────────────────────────────────────────────
async fn fetch_release(
client: &reqwest::Client,
version: Option<&str>,
) -> anyhow::Result<GithubRelease> {
match version {
Some(ver) => {
// Accept both "tunnel-v0.2.0" and the legacy "proxy-v0.2.0".
let tag = if ver.starts_with("tunnel-v") || ver.starts_with("proxy-v") {
ver.to_string()
} else {
format!("tunnel-v{}", ver)
};
let url = format!(
"{}/repos/{}/releases/tags/{}",
GITHUB_API_BASE, GITHUB_REPO, tag
);
let resp = client.get(&url).send().await?;
if !resp.status().is_success() {
let status = resp.status();
let body = resp.text().await.unwrap_or_default();
anyhow::bail!("release '{}' not found (HTTP {}): {}", tag, status, body);
}
Ok(resp.json().await?)
}
None => {
// List releases and find the latest tunnel-v* tag
let url = format!(
"{}/repos/{}/releases?per_page=20",
GITHUB_API_BASE, GITHUB_REPO
);
let resp = client.get(&url).send().await?;
if !resp.status().is_success() {
let status = resp.status();
let body = resp.text().await.unwrap_or_default();
anyhow::bail!("failed to list releases (HTTP {}): {}", status, body);
}
let releases: Vec<GithubRelease> = resp.json().await?;
releases
.into_iter()
.find(|r| r.tag_name.starts_with("tunnel-v") || r.tag_name.starts_with("proxy-v"))
.ok_or_else(|| anyhow::anyhow!("no tunnel-v* release found"))
}
}
}
// ── Download via GitHub release direct links ─────────────────────────────────
/// Download a release asset via the public direct download URL:
/// `https://github.com/{repo}/releases/download/{tag}/{filename}`
async fn download_release_file(
client: &reqwest::Client,
tag: &str,
filename: &str,
) -> anyhow::Result<Vec<u8>> {
let url = format!(
"https://github.com/{}/releases/download/{}/{}",
GITHUB_REPO, tag, filename
);
let resp = client
.get(&url)
.header(reqwest::header::ACCEPT, "application/octet-stream")
.send()
.await?;
if !resp.status().is_success() {
anyhow::bail!(
"download failed for '{}' (HTTP {})",
filename,
resp.status(),
);
}
Ok(resp.bytes().await?.to_vec())
}
fn parse_checksum(sums_text: &str, filename: &str) -> anyhow::Result<String> {
for line in sums_text.lines() {
// Format: "<hash> <filename>" (GNU coreutils convention)
let mut parts = line.split_ascii_whitespace();
let (Some(hash), Some(name)) = (parts.next(), parts.next()) else {
continue;
};
if name == filename || name.ends_with(filename) {
return Ok(hash.to_lowercase());
}
}
anyhow::bail!("checksum for '{}' not found in SHA256SUMS.txt", filename);
}
async fn download_and_verify(
client: &reqwest::Client,
tag: &str,
platform: &str,
dest: &Path,
) -> anyhow::Result<()> {
let archive_name = format!("aether-tunnel-{}.tar.gz", platform);
eprintln!(" Downloading {}...", archive_name);
let (archive_bytes, checksum_bytes) = tokio::try_join!(
download_release_file(client, tag, &archive_name),
download_release_file(client, tag, "SHA256SUMS.txt"),
)?;
let checksum_text = String::from_utf8(checksum_bytes)?;
eprintln!(
" Downloaded {} ({} bytes)",
archive_name,
archive_bytes.len()
);
// Verify SHA256
let expected_hash = parse_checksum(&checksum_text, &archive_name)?;
let mut hasher = Sha256::new();
hasher.update(&archive_bytes);
let actual_hash = hex::encode(hasher.finalize());
if actual_hash != expected_hash {
anyhow::bail!(
"SHA256 mismatch for {}:\n expected: {}\n actual: {}",
archive_name,
expected_hash,
actual_hash
);
}
eprintln!(" SHA256 verified: {}", &actual_hash[..16]);
extract_binary(&archive_bytes, dest)?;
Ok(())
}
// ── Archive extraction ───────────────────────────────────────────────────────
fn extract_binary(archive_bytes: &[u8], dest: &Path) -> anyhow::Result<()> {
use flate2::read::GzDecoder;
use tar::Archive;
// Guard against decompression bombs
const MAX_BINARY_SIZE: u64 = 100 * 1024 * 1024; // 100 MB
let decoder = GzDecoder::new(archive_bytes);
let mut archive = Archive::new(decoder);
let binary_name = if cfg!(target_os = "windows") {
"aether-tunnel.exe"
} else {
"aether-tunnel"
};
for entry in archive.entries()? {
let mut entry = entry?;
// Only accept regular files -- reject symlinks to prevent write-through attacks
if entry.header().entry_type() != tar::EntryType::Regular {
continue;
}
let path = entry.path()?;
if path.file_name().and_then(|n| n.to_str()) == Some(binary_name) {
let size = entry.header().size()?;
if size > MAX_BINARY_SIZE {
anyhow::bail!(
"binary too large ({} bytes, max {} bytes)",
size,
MAX_BINARY_SIZE
);
}
let mut file = std::fs::File::create(dest)?;
std::io::copy(&mut entry, &mut file)?;
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
std::fs::set_permissions(dest, std::fs::Permissions::from_mode(0o755))?;
}
return Ok(());
}
}
anyhow::bail!("'{}' not found in archive", binary_name);
}
// ── Atomic binary replacement ────────────────────────────────────────────────
fn atomic_replace(new_binary: &Path) -> anyhow::Result<PathBuf> {
let current_exe = std::env::current_exe()?.canonicalize()?;
let backup_path = current_exe.with_extension("bak");
// Remove stale backup
let _ = std::fs::remove_file(&backup_path);
// current -> .bak
std::fs::rename(&current_exe, &backup_path).map_err(|e| {
anyhow::anyhow!(
"failed to backup current binary '{}' -> '{}': {}",
current_exe.display(),
backup_path.display(),
e
)
})?;
// new -> current
if let Err(e) = std::fs::rename(new_binary, &current_exe) {
eprintln!(" ERROR: failed to place new binary, rolling back...");
let _ = std::fs::rename(&backup_path, &current_exe);
anyhow::bail!(
"failed to install new binary '{}' -> '{}': {}",
new_binary.display(),
current_exe.display(),
e
);
}
eprintln!(" Binary replaced: {}", current_exe.display());
Ok(backup_path)
}
// ── Public entry point ───────────────────────────────────────────────────────
#[derive(Clone, Copy)]
enum RestartMode {
BestEffort,
Required,
}
async fn execute_upgrade(
version: Option<&str>,
require_root: bool,
restart_mode: RestartMode,
) -> anyhow::Result<()> {
// Resolve exe path once; reuse throughout the function
let current_exe = std::env::current_exe()?.canonicalize()?;
let exe_dir = current_exe
.parent()
.ok_or_else(|| anyhow::anyhow!("cannot determine binary directory"))?;
let temp_path = exe_dir.join(".aether-tunnel.upgrade.tmp");
if require_root {
if !super::service::is_root() {
anyhow::bail!("automatic upgrade requires root privileges");
}
} else if !super::service::is_root() {
// Check write permission to binary directory for manual upgrade mode.
let test_path = exe_dir.join(".aether-tunnel.write-test");
match std::fs::File::create(&test_path) {
Ok(_) => {
let _ = std::fs::remove_file(&test_path);
}
Err(_) => {
anyhow::bail!(
"no write access to {}. Use: sudo aether-tunnel upgrade",
exe_dir.display()
);
}
}
}
let platform = detect_platform();
eprintln!(" Platform: {}", platform);
eprintln!(" Current version: {}", CURRENT_VERSION);
let client = build_github_client()?;
let release = fetch_release(&client, version).await?;
let target_tag = &release.tag_name;
let target_semver = target_tag
.strip_prefix("tunnel-v")
.or_else(|| target_tag.strip_prefix("proxy-v"))
.unwrap_or(target_tag);
eprintln!(" Target version: {} ({})", target_tag, release.name);
if target_semver == CURRENT_VERSION {
eprintln!(
" Already running version {}, nothing to do.",
CURRENT_VERSION
);
return Ok(());
}
eprintln!();
eprintln!(" Upgrading: {} -> {}", CURRENT_VERSION, target_semver);
eprintln!();
if let Err(e) = download_and_verify(&client, target_tag, platform, &temp_path).await {
let _ = std::fs::remove_file(&temp_path);
return Err(e);
}
let backup_path = match atomic_replace(&temp_path) {
Ok(backup) => backup,
Err(e) => {
let _ = std::fs::remove_file(&temp_path);
return Err(e);
}
};
match restart_mode {
RestartMode::BestEffort => {
// Use best-effort: binary is already replaced, so a restart failure should
// not abort the whole upgrade -- the user can restart manually.
if super::service::is_service_active() {
if super::service::is_root() {
eprintln!(" Restarting managed service...");
match super::service::restart_active_service() {
Ok(()) => eprintln!(" Service restarted."),
Err(e) => {
eprintln!(" WARNING: failed to restart service: {}", e);
eprintln!(" Run manually: sudo aether-tunnel restart");
}
}
} else {
eprintln!(" Managed service is active, but restart requires root.");
eprintln!(" Run: sudo aether-tunnel restart");
eprintln!(" Skipping restart.");
}
} else {
eprintln!(" No active service detected, skipping restart.");
}
}
RestartMode::Required => {
if !super::service::is_root() {
anyhow::bail!("automatic upgrade requires root privileges");
}
eprintln!(" Restarting managed service...");
super::service::restart_active_service()?;
eprintln!(" Service restarted.");
}
}
eprintln!();
eprintln!(" Upgrade complete!");
eprintln!(
" Backup kept at: {} (will be cleaned up on next upgrade)",
backup_path.display()
);
Ok(())
}
/// `aether-tunnel upgrade [version]` -- self-upgrade from GitHub releases.
pub async fn cmd_upgrade(version: Option<String>) -> anyhow::Result<()> {
execute_upgrade(version.as_deref(), false, RestartMode::BestEffort).await
}
/// Perform automatic upgrade to a specific version.
///
/// This path is used for server-pushed upgrades: it requires root and expects
/// the currently active managed service to restart successfully.
pub async fn perform_upgrade(version: &str) -> anyhow::Result<()> {
execute_upgrade(Some(version), true, RestartMode::Required).await
}

View File

@@ -0,0 +1,758 @@
//! Shared application state passed to all subsystems.
use std::collections::VecDeque;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex, RwLock};
use std::time::Duration;
use std::time::{SystemTime, UNIX_EPOCH};
use aether_runtime::{
service_up_sample, AdmissionPermit, ConcurrencyError, ConcurrencyGate, ConcurrencySnapshot,
MetricKind, MetricLabel, MetricSample,
};
use aether_runtime_state::{RuntimeSemaphore, RuntimeSemaphoreError, RuntimeSemaphoreSnapshot};
use crate::config::Config;
use crate::hardware::RuntimeResourceMonitor;
use crate::registration::client::AetherClient;
use crate::runtime::SharedDynamicConfig;
use crate::target_filter::DnsCache;
use crate::upstream_client::UpstreamClientPool;
/// Central application state shared across all servers/tunnels.
pub struct AppState {
pub config: Arc<Config>,
/// DNS cache for upstream target resolution (shared).
pub dns_cache: Arc<DnsCache>,
/// Profile-keyed upstream client pool used by tunnel requests.
pub upstream_client_pool: UpstreamClientPool,
/// Shared TLS config for tunnel WebSocket connections (avoids re-parsing root CAs on each reconnect).
pub tunnel_tls_config: Arc<rustls::ClientConfig>,
/// Runtime CPU/memory monitor sampled by heartbeat payloads.
pub resource_monitor: Arc<RuntimeResourceMonitor>,
/// Optional per-process stream admission gate.
pub stream_gate: Option<Arc<ConcurrencyGate>>,
/// Optional cross-instance stream admission gate.
pub distributed_stream_gate: Option<Arc<RuntimeSemaphore>>,
}
/// Per-server state: one instance per Aether server connection.
pub struct ServerContext {
/// Human-readable label for logging (e.g. "server-0").
pub server_label: String,
/// Aether server URL for this connection.
pub aether_url: String,
/// Management token for this server.
pub management_token: String,
/// Resolved node name at registration time (per-server override or global fallback).
/// After startup, the active node_name is read from `dynamic` (may be updated remotely).
#[allow(dead_code)]
pub node_name: String,
/// Node ID assigned by this Aether server.
pub node_id: Arc<RwLock<String>>,
/// API client for this server.
pub aether_client: Arc<AetherClient>,
/// Dynamic config from this server's heartbeat ACKs.
pub dynamic: SharedDynamicConfig,
/// Per-server active connection count.
pub active_connections: Arc<AtomicU64>,
/// Per-server request/latency metrics.
pub metrics: Arc<TunnelRequestMetrics>,
/// Per-server tunnel stability/traffic metrics.
pub tunnel_metrics: Arc<TunnelMetrics>,
}
impl ServerContext {
pub fn metric_samples(&self) -> Vec<MetricSample> {
let mut samples = self.metrics.to_metric_samples(&self.server_label);
samples.extend(self.tunnel_metrics.to_metric_samples(&self.server_label));
samples.push(
MetricSample::new(
"tunnel_active_connections",
"Current number of active tunneled streams handled by this tunnel server context.",
MetricKind::Gauge,
self.active_connections.load(Ordering::Acquire),
)
.with_labels(vec![MetricLabel::new("server", self.server_label.clone())]),
);
samples
}
}
/// Aggregate metrics for reporting to Aether.
pub struct TunnelRequestMetrics {
pub total_requests: AtomicU64,
/// Cumulative connection-establishment latency in nanoseconds
/// (DNS + TCP/TLS + TTFB, excludes response body streaming).
pub total_latency_ns: AtomicU64,
pub failed_requests: AtomicU64,
pub dns_failures: AtomicU64,
pub stream_errors: AtomicU64,
pub slow_requests: AtomicU64,
}
#[derive(Debug, Clone, Copy, Default, serde::Serialize)]
pub struct TunnelRequestMetricsSnapshot {
pub total_requests: u64,
pub total_latency_ns: u64,
pub failed_requests: u64,
pub dns_failures: u64,
pub stream_errors: u64,
pub slow_requests: u64,
}
impl TunnelRequestMetricsSnapshot {
pub fn average_latency_ns(self) -> Option<u64> {
self.total_latency_ns.checked_div(self.total_requests)
}
pub fn average_latency_ms(self) -> Option<f64> {
self.average_latency_ns()
.map(|value| value as f64 / 1_000_000.0)
}
pub fn delta_since(self, baseline: Self) -> Self {
Self {
total_requests: self.total_requests.saturating_sub(baseline.total_requests),
total_latency_ns: self
.total_latency_ns
.saturating_sub(baseline.total_latency_ns),
failed_requests: self
.failed_requests
.saturating_sub(baseline.failed_requests),
dns_failures: self.dns_failures.saturating_sub(baseline.dns_failures),
stream_errors: self.stream_errors.saturating_sub(baseline.stream_errors),
slow_requests: self.slow_requests.saturating_sub(baseline.slow_requests),
}
}
}
impl TunnelRequestMetrics {
pub fn new() -> Self {
Self {
total_requests: AtomicU64::new(0),
total_latency_ns: AtomicU64::new(0),
failed_requests: AtomicU64::new(0),
dns_failures: AtomicU64::new(0),
stream_errors: AtomicU64::new(0),
slow_requests: AtomicU64::new(0),
}
}
/// Record a completed request with its connection-establishment latency
/// (DNS + TCP/TLS + TTFB, excludes response body streaming).
pub fn record_request(&self, connect_elapsed: Duration) {
let nanos = u64::try_from(connect_elapsed.as_nanos()).unwrap_or(u64::MAX);
self.total_requests.fetch_add(1, Ordering::Release);
self.total_latency_ns.fetch_add(nanos, Ordering::Release);
}
pub fn record_slow_request(&self) {
self.slow_requests.fetch_add(1, Ordering::Release);
}
pub fn snapshot(&self) -> TunnelRequestMetricsSnapshot {
TunnelRequestMetricsSnapshot {
total_requests: self.total_requests.load(Ordering::Acquire),
total_latency_ns: self.total_latency_ns.load(Ordering::Acquire),
failed_requests: self.failed_requests.load(Ordering::Acquire),
dns_failures: self.dns_failures.load(Ordering::Acquire),
stream_errors: self.stream_errors.load(Ordering::Acquire),
slow_requests: self.slow_requests.load(Ordering::Acquire),
}
}
pub fn to_metric_samples(&self, server_label: &str) -> Vec<MetricSample> {
let snapshot = self.snapshot();
let labels = vec![MetricLabel::new("server", server_label)];
vec![
MetricSample::new(
"tunnel_requests_total",
"Total number of tunneled upstream requests completed by the tunnel.",
MetricKind::Counter,
snapshot.total_requests,
)
.with_labels(labels.clone()),
MetricSample::new(
"tunnel_request_latency_total_ns",
"Cumulative tunnel request latency in nanoseconds through upstream response headers.",
MetricKind::Counter,
snapshot.total_latency_ns,
)
.with_labels(labels.clone()),
MetricSample::new(
"tunnel_request_latency_avg_ns",
"Average tunnel request latency in nanoseconds through upstream response headers.",
MetricKind::Gauge,
snapshot.average_latency_ns().unwrap_or(0),
)
.with_labels(labels.clone()),
MetricSample::new(
"tunnel_failed_requests_total",
"Total number of tunneled upstream requests that failed before response headers.",
MetricKind::Counter,
snapshot.failed_requests,
)
.with_labels(labels.clone()),
MetricSample::new(
"tunnel_dns_failures_total",
"Total number of tunneled upstream requests rejected or failed during target validation or DNS.",
MetricKind::Counter,
snapshot.dns_failures,
)
.with_labels(labels.clone()),
MetricSample::new(
"tunnel_stream_errors_total",
"Total number of tunneled response body stream errors.",
MetricKind::Counter,
snapshot.stream_errors,
)
.with_labels(labels.clone()),
MetricSample::new(
"tunnel_slow_requests_total",
"Total number of tunneled requests crossing the tunnel slow-request threshold.",
MetricKind::Counter,
snapshot.slow_requests,
)
.with_labels(labels),
]
}
}
const RECENT_TUNNEL_ERROR_CAPACITY: usize = 64;
const TUNNEL_ERROR_CATEGORY_MAX_CHARS: usize = 48;
const TUNNEL_ERROR_MESSAGE_MAX_CHARS: usize = 320;
#[derive(Debug, Clone, serde::Serialize)]
pub struct TunnelErrorEvent {
pub timestamp_unix_secs: u64,
pub timestamp_unix_ms: u64,
pub category: String,
pub message: String,
pub severity: String,
pub component: String,
pub summary: String,
pub operator_action: String,
}
#[derive(Debug, Clone, Copy, Default, serde::Serialize)]
pub struct TunnelMetricsSnapshot {
pub connect_attempts: u64,
pub connect_successes: u64,
pub connect_errors: u64,
pub disconnects: u64,
pub last_connected_at_unix_secs: u64,
pub last_disconnected_at_unix_secs: u64,
pub last_connected_duration_ms: u64,
pub connected_duration_total_ms: u64,
pub heartbeat_sent: u64,
pub heartbeat_ack: u64,
pub heartbeat_rtt_last_ms: u64,
pub heartbeat_rtt_total_ms: u64,
pub ws_in_frames: u64,
pub ws_in_bytes: u64,
pub ws_out_frames: u64,
pub ws_out_bytes: u64,
pub error_events_total: u64,
}
impl TunnelMetricsSnapshot {
pub fn heartbeat_rtt_avg_ms(self) -> Option<f64> {
if self.heartbeat_ack == 0 {
None
} else {
Some(self.heartbeat_rtt_total_ms as f64 / self.heartbeat_ack as f64)
}
}
}
pub struct TunnelMetrics {
connect_attempts: AtomicU64,
connect_successes: AtomicU64,
connect_errors: AtomicU64,
disconnects: AtomicU64,
last_connected_at_unix_secs: AtomicU64,
last_disconnected_at_unix_secs: AtomicU64,
last_connected_duration_ms: AtomicU64,
connected_duration_total_ms: AtomicU64,
heartbeat_sent: AtomicU64,
heartbeat_ack: AtomicU64,
heartbeat_rtt_last_ms: AtomicU64,
heartbeat_rtt_total_ms: AtomicU64,
ws_in_frames: AtomicU64,
ws_in_bytes: AtomicU64,
ws_out_frames: AtomicU64,
ws_out_bytes: AtomicU64,
error_events_total: AtomicU64,
recent_errors: Mutex<VecDeque<TunnelErrorEvent>>,
}
impl TunnelMetrics {
pub fn new() -> Self {
Self {
connect_attempts: AtomicU64::new(0),
connect_successes: AtomicU64::new(0),
connect_errors: AtomicU64::new(0),
disconnects: AtomicU64::new(0),
last_connected_at_unix_secs: AtomicU64::new(0),
last_disconnected_at_unix_secs: AtomicU64::new(0),
last_connected_duration_ms: AtomicU64::new(0),
connected_duration_total_ms: AtomicU64::new(0),
heartbeat_sent: AtomicU64::new(0),
heartbeat_ack: AtomicU64::new(0),
heartbeat_rtt_last_ms: AtomicU64::new(0),
heartbeat_rtt_total_ms: AtomicU64::new(0),
ws_in_frames: AtomicU64::new(0),
ws_in_bytes: AtomicU64::new(0),
ws_out_frames: AtomicU64::new(0),
ws_out_bytes: AtomicU64::new(0),
error_events_total: AtomicU64::new(0),
recent_errors: Mutex::new(VecDeque::with_capacity(RECENT_TUNNEL_ERROR_CAPACITY)),
}
}
pub fn record_connect_attempt(&self) {
self.connect_attempts.fetch_add(1, Ordering::Release);
}
pub fn record_connect_success(&self) {
self.connect_successes.fetch_add(1, Ordering::Release);
self.last_connected_at_unix_secs
.store(now_unix_secs(), Ordering::Release);
}
pub fn record_connect_error(&self) {
self.connect_errors.fetch_add(1, Ordering::Release);
}
pub fn record_disconnect(&self, connected_for: Duration) {
let duration_ms = duration_to_millis_u64(connected_for);
self.disconnects.fetch_add(1, Ordering::Release);
self.last_disconnected_at_unix_secs
.store(now_unix_secs(), Ordering::Release);
self.last_connected_duration_ms
.store(duration_ms, Ordering::Release);
self.connected_duration_total_ms
.fetch_add(duration_ms, Ordering::Release);
}
pub fn record_heartbeat_sent(&self) {
self.heartbeat_sent.fetch_add(1, Ordering::Release);
}
pub fn record_heartbeat_ack(&self, rtt: Duration) {
let rtt_ms = duration_to_millis_u64(rtt);
self.heartbeat_ack.fetch_add(1, Ordering::Release);
self.heartbeat_rtt_last_ms.store(rtt_ms, Ordering::Release);
self.heartbeat_rtt_total_ms
.fetch_add(rtt_ms, Ordering::Release);
}
pub fn record_ws_incoming_frame(&self, payload_len: usize) {
self.ws_in_frames.fetch_add(1, Ordering::Release);
self.ws_in_bytes.fetch_add(
u64::try_from(payload_len).unwrap_or(u64::MAX),
Ordering::Release,
);
}
pub fn record_ws_outgoing_frame(&self, payload_len: usize) {
self.ws_out_frames.fetch_add(1, Ordering::Release);
self.ws_out_bytes.fetch_add(
u64::try_from(payload_len).unwrap_or(u64::MAX),
Ordering::Release,
);
}
pub fn record_error(&self, category: &str, message: &str) {
self.error_events_total.fetch_add(1, Ordering::Release);
let category = normalize_error_field(category, TUNNEL_ERROR_CATEGORY_MAX_CHARS, "unknown");
let message = normalize_error_field(message, TUNNEL_ERROR_MESSAGE_MAX_CHARS, "n/a");
let diagnostic = classify_tunnel_error(category.as_str(), message.as_str());
let timestamp_unix_ms = now_unix_ms();
let event = TunnelErrorEvent {
timestamp_unix_secs: timestamp_unix_ms / 1_000,
timestamp_unix_ms,
category,
message,
severity: diagnostic.severity.to_string(),
component: diagnostic.component.to_string(),
summary: diagnostic.summary.to_string(),
operator_action: diagnostic.operator_action.to_string(),
};
let mut recent_errors = match self.recent_errors.lock() {
Ok(guard) => guard,
Err(poisoned) => poisoned.into_inner(),
};
if recent_errors.len() >= RECENT_TUNNEL_ERROR_CAPACITY {
recent_errors.pop_front();
}
recent_errors.push_back(event);
}
pub fn recent_errors(&self, limit: usize) -> Vec<TunnelErrorEvent> {
if limit == 0 {
return Vec::new();
}
let recent_errors = match self.recent_errors.lock() {
Ok(guard) => guard,
Err(poisoned) => poisoned.into_inner(),
};
let start = recent_errors.len().saturating_sub(limit);
recent_errors.iter().skip(start).cloned().collect()
}
pub fn snapshot(&self) -> TunnelMetricsSnapshot {
TunnelMetricsSnapshot {
connect_attempts: self.connect_attempts.load(Ordering::Acquire),
connect_successes: self.connect_successes.load(Ordering::Acquire),
connect_errors: self.connect_errors.load(Ordering::Acquire),
disconnects: self.disconnects.load(Ordering::Acquire),
last_connected_at_unix_secs: self.last_connected_at_unix_secs.load(Ordering::Acquire),
last_disconnected_at_unix_secs: self
.last_disconnected_at_unix_secs
.load(Ordering::Acquire),
last_connected_duration_ms: self.last_connected_duration_ms.load(Ordering::Acquire),
connected_duration_total_ms: self.connected_duration_total_ms.load(Ordering::Acquire),
heartbeat_sent: self.heartbeat_sent.load(Ordering::Acquire),
heartbeat_ack: self.heartbeat_ack.load(Ordering::Acquire),
heartbeat_rtt_last_ms: self.heartbeat_rtt_last_ms.load(Ordering::Acquire),
heartbeat_rtt_total_ms: self.heartbeat_rtt_total_ms.load(Ordering::Acquire),
ws_in_frames: self.ws_in_frames.load(Ordering::Acquire),
ws_in_bytes: self.ws_in_bytes.load(Ordering::Acquire),
ws_out_frames: self.ws_out_frames.load(Ordering::Acquire),
ws_out_bytes: self.ws_out_bytes.load(Ordering::Acquire),
error_events_total: self.error_events_total.load(Ordering::Acquire),
}
}
pub fn to_metric_samples(&self, server_label: &str) -> Vec<MetricSample> {
let snapshot = self.snapshot();
let labels = vec![MetricLabel::new("server", server_label)];
vec![
MetricSample::new(
"tunnel_connect_attempts_total",
"Total number of WebSocket tunnel connection attempts.",
MetricKind::Counter,
snapshot.connect_attempts,
)
.with_labels(labels.clone()),
MetricSample::new(
"tunnel_connect_successes_total",
"Total number of successful WebSocket tunnel connections.",
MetricKind::Counter,
snapshot.connect_successes,
)
.with_labels(labels.clone()),
MetricSample::new(
"tunnel_connect_errors_total",
"Total number of WebSocket tunnel connection errors.",
MetricKind::Counter,
snapshot.connect_errors,
)
.with_labels(labels.clone()),
MetricSample::new(
"tunnel_disconnects_total",
"Total number of WebSocket tunnel disconnects.",
MetricKind::Counter,
snapshot.disconnects,
)
.with_labels(labels.clone()),
MetricSample::new(
"tunnel_heartbeat_sent_total",
"Total number of tunnel heartbeats sent.",
MetricKind::Counter,
snapshot.heartbeat_sent,
)
.with_labels(labels.clone()),
MetricSample::new(
"tunnel_heartbeat_ack_total",
"Total number of tunnel heartbeat acknowledgements received.",
MetricKind::Counter,
snapshot.heartbeat_ack,
)
.with_labels(labels.clone()),
MetricSample::new(
"tunnel_heartbeat_rtt_last_ms",
"Last observed tunnel heartbeat round-trip time in milliseconds.",
MetricKind::Gauge,
snapshot.heartbeat_rtt_last_ms,
)
.with_labels(labels.clone()),
MetricSample::new(
"tunnel_heartbeat_rtt_avg_ms",
"Average observed tunnel heartbeat round-trip time in milliseconds.",
MetricKind::Gauge,
snapshot.heartbeat_rtt_avg_ms().unwrap_or(0.0) as u64,
)
.with_labels(labels.clone()),
MetricSample::new(
"tunnel_ws_in_frames_total",
"Total number of WebSocket frames received by the tunnel.",
MetricKind::Counter,
snapshot.ws_in_frames,
)
.with_labels(labels.clone()),
MetricSample::new(
"tunnel_ws_in_bytes_total",
"Total number of WebSocket bytes received by the tunnel.",
MetricKind::Counter,
snapshot.ws_in_bytes,
)
.with_labels(labels.clone()),
MetricSample::new(
"tunnel_ws_out_frames_total",
"Total number of WebSocket frames sent by the tunnel.",
MetricKind::Counter,
snapshot.ws_out_frames,
)
.with_labels(labels.clone()),
MetricSample::new(
"tunnel_ws_out_bytes_total",
"Total number of WebSocket bytes sent by the tunnel.",
MetricKind::Counter,
snapshot.ws_out_bytes,
)
.with_labels(labels.clone()),
MetricSample::new(
"tunnel_error_events_total",
"Total number of classified tunnel error events recorded by the tunnel.",
MetricKind::Counter,
snapshot.error_events_total,
)
.with_labels(labels),
]
}
}
fn now_unix_secs() -> u64 {
now_unix_ms() / 1_000
}
fn now_unix_ms() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| u64::try_from(d.as_millis()).unwrap_or(u64::MAX))
.unwrap_or(0)
}
fn duration_to_millis_u64(duration: Duration) -> u64 {
u64::try_from(duration.as_millis()).unwrap_or(u64::MAX)
}
fn normalize_error_field(value: &str, max_chars: usize, fallback: &str) -> String {
let normalized = value.split_whitespace().collect::<Vec<_>>().join(" ");
if normalized.is_empty() {
return fallback.to_string();
}
normalized.chars().take(max_chars).collect()
}
struct TunnelErrorDiagnostic {
severity: &'static str,
component: &'static str,
summary: &'static str,
operator_action: &'static str,
}
fn classify_tunnel_error(category: &str, _message: &str) -> TunnelErrorDiagnostic {
match category {
"stale_timeout" => TunnelErrorDiagnostic {
severity: "warning",
component: "tunnel_read",
summary: "No inbound tunnel frames before stale timeout",
operator_action:
"Check gateway or reverse-proxy idle timeouts, packet loss, and WebSocket ping/pong reachability. Increase AETHER_TUNNEL_STALE_TIMEOUT_MS if the network is high-latency.",
},
"ws_write_error" => TunnelErrorDiagnostic {
severity: "error",
component: "tunnel_write",
summary: "WebSocket write failed because the peer closed or reset the connection",
operator_action:
"Check gateway restarts, load balancer resets, NAT/firewall connection tracking, and whether the tunnel is reconnecting successfully.",
},
"ws_ping_error" => TunnelErrorDiagnostic {
severity: "error",
component: "tunnel_write",
summary: "WebSocket keepalive ping could not be sent",
operator_action:
"Check whether the peer closed the socket or an intermediary is dropping idle WebSocket connections.",
},
"ws_read_error" => TunnelErrorDiagnostic {
severity: "error",
component: "tunnel_read",
summary: "WebSocket read failed",
operator_action:
"Check gateway logs and network stability around the same timestamp; compare with reconnect and heartbeat ACK counters.",
},
"tunnel_connect_error" => TunnelErrorDiagnostic {
severity: "critical",
component: "tunnel_connect",
summary: "Tunnel connection attempt failed",
operator_action:
"Check Aether URL reachability, DNS, TLS, management token validity, and any configured AETHER_TUNNEL_AETHER_OUTBOUND_PROXY_URL.",
},
"frame_decode_error" => TunnelErrorDiagnostic {
severity: "error",
component: "tunnel_protocol",
summary: "Received tunnel frame could not be decoded",
operator_action:
"Check tunnel and gateway version compatibility and whether traffic is being modified by an intermediary.",
},
"stream_dispatch_timeout" => TunnelErrorDiagnostic {
severity: "warning",
component: "stream_dispatch",
summary: "Request body frame could not be delivered to its stream handler in time",
operator_action:
"Check tunnel CPU, memory, stream concurrency saturation, and slow upstream provider requests.",
},
"heartbeat_ack_empty" | "heartbeat_ack_parse" => TunnelErrorDiagnostic {
severity: "warning",
component: "heartbeat",
summary: "Heartbeat ACK from gateway was missing or invalid",
operator_action:
"Check gateway heartbeat handler logs and tunnel/gateway version compatibility.",
},
"writer_task_panic" | "writer_task_cancelled" => TunnelErrorDiagnostic {
severity: "error",
component: "tunnel_writer",
summary: "Tunnel writer task exited unexpectedly",
operator_action:
"Check tunnel logs for the preceding write or ping error and confirm the tunnel reconnect loop is active.",
},
"dispatcher_error" => TunnelErrorDiagnostic {
severity: "error",
component: "tunnel_dispatcher",
summary: "Tunnel dispatcher exited with an error",
operator_action:
"Check the proxied request stream and gateway tunnel logs around the same timestamp.",
},
_ => TunnelErrorDiagnostic {
severity: "info",
component: "tunnel",
summary: "Tunnel reported an unclassified error",
operator_action:
"Inspect the raw message and compare it with tunnel, gateway, and network logs at the same time.",
},
}
}
#[derive(Debug, thiserror::Error, PartialEq, Eq)]
pub enum TunnelAdmissionError {
#[error("tunnel stream admission saturated at {limit} for gate {gate}")]
Saturated { gate: &'static str, limit: usize },
#[error("tunnel stream admission unavailable for gate {gate}: {message}")]
Unavailable {
gate: &'static str,
limit: usize,
message: String,
},
}
impl AppState {
pub async fn metric_samples(&self) -> Vec<MetricSample> {
let mut samples = vec![service_up_sample("aether-tunnel")];
if let Some(snapshot) = self.stream_concurrency_snapshot() {
samples.extend(snapshot.to_metric_samples("tunnel_streams"));
}
if let Some(gate) = self.distributed_stream_gate.as_ref() {
match gate.snapshot().await {
Ok(snapshot) => {
samples.extend(snapshot.to_metric_samples("tunnel_streams_distributed"));
}
Err(_) => samples.push(
MetricSample::new(
"concurrency_unavailable",
"Whether the distributed concurrency gate is currently unavailable.",
MetricKind::Gauge,
1,
)
.with_labels(vec![MetricLabel::new("gate", "tunnel_streams_distributed")]),
),
}
}
samples
}
pub fn with_stream_concurrency_gate(mut self, gate: Arc<ConcurrencyGate>) -> Self {
self.stream_gate = Some(gate);
self
}
pub fn with_distributed_stream_concurrency_gate(mut self, gate: Arc<RuntimeSemaphore>) -> Self {
self.distributed_stream_gate = Some(gate);
self
}
pub fn stream_concurrency_snapshot(&self) -> Option<ConcurrencySnapshot> {
self.stream_gate.as_ref().map(|gate| gate.snapshot())
}
pub async fn distributed_stream_concurrency_snapshot(
&self,
) -> Result<Option<RuntimeSemaphoreSnapshot>, RuntimeSemaphoreError> {
match &self.distributed_stream_gate {
Some(gate) => gate.snapshot().await.map(Some),
None => Ok(None),
}
}
pub async fn try_acquire_stream_permit(
&self,
) -> Result<Option<AdmissionPermit>, TunnelAdmissionError> {
let local = match &self.stream_gate {
Some(gate) => Some(gate.try_acquire().map_err(|err| {
match err {
ConcurrencyError::Saturated { gate, limit } => {
TunnelAdmissionError::Saturated { gate, limit }
}
ConcurrencyError::Closed { gate } => TunnelAdmissionError::Unavailable {
gate,
limit: self
.stream_gate
.as_ref()
.map(|inner| inner.snapshot().limit)
.unwrap_or(0),
message: "local stream gate is closed".to_string(),
},
}
})?),
None => None,
};
let distributed = match &self.distributed_stream_gate {
Some(gate) => Some(gate.try_acquire().await.map_err(|err| {
match err {
RuntimeSemaphoreError::Saturated { gate, limit } => {
TunnelAdmissionError::Saturated { gate, limit }
}
RuntimeSemaphoreError::Unavailable {
gate,
limit,
message,
} => TunnelAdmissionError::Unavailable {
gate,
limit,
message,
},
RuntimeSemaphoreError::InvalidConfiguration(message) => {
TunnelAdmissionError::Unavailable {
gate: "tunnel_streams_distributed",
limit: self
.distributed_stream_gate
.as_ref()
.map(|inner| inner.limit())
.unwrap_or(0),
message,
}
}
}
})?),
None => None,
};
Ok(AdmissionPermit::from_parts(local, distributed))
}
}

View File

@@ -0,0 +1,415 @@
use std::collections::{HashMap, HashSet};
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::RwLock;
/// Check if an IP address belongs to a private/reserved network.
pub fn is_private_ip(ip: &IpAddr) -> bool {
match ip {
IpAddr::V4(v4) => is_private_ipv4(v4),
IpAddr::V6(v6) => is_private_ipv6(v6),
}
}
fn is_private_ipv4(ip: &Ipv4Addr) -> bool {
let octets = ip.octets();
// 10.0.0.0/8
if octets[0] == 10 {
return true;
}
// 172.16.0.0/12
if octets[0] == 172 && (16..=31).contains(&octets[1]) {
return true;
}
// 192.168.0.0/16
if octets[0] == 192 && octets[1] == 168 {
return true;
}
// 127.0.0.0/8
if octets[0] == 127 {
return true;
}
// 169.254.0.0/16 (link-local)
if octets[0] == 169 && octets[1] == 254 {
return true;
}
// 0.0.0.0/8
if octets[0] == 0 {
return true;
}
// 100.64.0.0/10 (CGNAT / shared address space)
if octets[0] == 100 && (64..=127).contains(&octets[1]) {
return true;
}
// 192.0.0.0/24 (IETF protocol assignments)
if octets[0] == 192 && octets[1] == 0 && octets[2] == 0 {
return true;
}
// 198.18.0.0/15 (benchmark testing)
if octets[0] == 198 && (18..=19).contains(&octets[1]) {
return true;
}
// 240.0.0.0/4 (reserved for future use)
if octets[0] >= 240 {
return true;
}
false
}
fn is_private_ipv6(ip: &Ipv6Addr) -> bool {
// ::1 loopback
if ip.is_loopback() {
return true;
}
// :: unspecified
if ip.is_unspecified() {
return true;
}
let segments = ip.segments();
// fc00::/7 (ULA) - first byte is 0xfc or 0xfd
if segments[0] & 0xfe00 == 0xfc00 {
return true;
}
// fe80::/10 (link-local)
if segments[0] & 0xffc0 == 0xfe80 {
return true;
}
// IPv4-mapped IPv6 (::ffff:x.x.x.x) - check the embedded IPv4
if let Some(v4) = ip.to_ipv4_mapped() {
return is_private_ipv4(&v4);
}
false
}
#[derive(Debug)]
pub enum FilterError {
PrivateIp(IpAddr),
PortNotAllowed(u16),
DnsResolutionFailed(String),
NoPublicAddrs(String),
}
impl std::fmt::Display for FilterError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::PrivateIp(ip) => write!(f, "target IP {} is in private/reserved range", ip),
Self::PortNotAllowed(port) => write!(f, "port {} not in allowed list", port),
Self::DnsResolutionFailed(host) => write!(f, "DNS resolution failed for {}", host),
Self::NoPublicAddrs(host) => {
write!(
f,
"all resolved addresses for {} are private/reserved",
host
)
}
}
}
}
struct DnsCacheEntry {
addrs: Arc<Vec<SocketAddr>>,
expires_at: Instant,
inserted_at: Instant,
}
/// Lightweight DNS cache with TTL + capacity bounds.
/// Stores all public resolved addresses per host (used by SafeDnsResolver
/// to ensure reqwest connects to the same validated addresses).
pub struct DnsCache {
ttl: Duration,
capacity: usize,
entries: RwLock<HashMap<String, DnsCacheEntry>>,
}
impl DnsCache {
pub fn new(ttl: Duration, capacity: usize) -> Self {
Self {
ttl,
capacity,
entries: RwLock::new(HashMap::new()),
}
}
/// Look up cached public addresses for a host (any port).
///
/// Used by `SafeDnsResolver` which only knows the hostname — returns the
/// first unexpired entry whose key starts with `host:`.
pub async fn get_by_host(&self, host: &str) -> Option<Arc<Vec<SocketAddr>>> {
if self.capacity == 0 || self.ttl.is_zero() {
return None;
}
let prefix = format!("{}:", host.to_ascii_lowercase());
let now = Instant::now();
let entries = self.entries.read().await;
for (key, entry) in entries.iter() {
if key.starts_with(&prefix) && entry.expires_at > now {
return Some(Arc::clone(&entry.addrs));
}
}
None
}
/// Look up cached public addresses for a host + port.
pub async fn get(&self, host: &str, port: u16) -> Option<Arc<Vec<SocketAddr>>> {
if self.capacity == 0 || self.ttl.is_zero() {
return None;
}
let key = Self::key(host, port);
let now = Instant::now();
// Fast path: read lock for cache hit
{
let entries = self.entries.read().await;
match entries.get(&key) {
Some(entry) if entry.expires_at > now => return Some(Arc::clone(&entry.addrs)),
None => return None,
Some(_) => {} // expired, fall through to evict
}
}
// Slow path: write lock to remove expired entry
let mut entries = self.entries.write().await;
entries.remove(&key);
None
}
/// Insert resolved public addresses into cache.
pub async fn insert(&self, host: &str, port: u16, addrs: Arc<Vec<SocketAddr>>) {
if self.capacity == 0 || self.ttl.is_zero() || addrs.is_empty() {
return;
}
let key = Self::key(host, port);
let now = Instant::now();
let mut entries = self.entries.write().await;
entries.retain(|_, entry| entry.expires_at > now);
while entries.len() >= self.capacity {
let oldest_key = entries
.iter()
.min_by_key(|(_, entry)| entry.inserted_at)
.map(|(key, _)| key.clone());
if let Some(key) = oldest_key {
entries.remove(&key);
} else {
break;
}
}
entries.insert(
key,
DnsCacheEntry {
addrs,
expires_at: now + self.ttl,
inserted_at: now,
},
);
}
fn key(host: &str, port: u16) -> String {
format!("{}:{}", host.to_ascii_lowercase(), port)
}
}
/// Resolve a hostname to validated socket addresses.
///
/// Results are cached in `dns_cache`. Private/reserved IPs are filtered out
/// unless `allow_private` is enabled. Returns an error if filtering removes
/// every resolved address.
pub async fn resolve_public_addrs(
host: &str,
port: u16,
allow_private: bool,
dns_cache: &DnsCache,
) -> Result<Vec<SocketAddr>, FilterError> {
// Cache hit
if let Some(addrs) = dns_cache.get(host, port).await {
return Ok((*addrs).clone());
}
// Async DNS resolution
let addr_str = format!("{}:{}", host, port);
let resolved: Vec<SocketAddr> = tokio::net::lookup_host(&addr_str)
.await
.map_err(|_| FilterError::DnsResolutionFailed(host.to_string()))?
.collect();
if resolved.is_empty() {
return Err(FilterError::DnsResolutionFailed(host.to_string()));
}
// Filter out private/reserved addresses unless explicitly allowed.
let public: Vec<SocketAddr> = if allow_private {
resolved
} else {
resolved
.into_iter()
.filter(|addr| !is_private_ip(&addr.ip()))
.collect()
};
if public.is_empty() {
return Err(FilterError::NoPublicAddrs(host.to_string()));
}
// Cache the validated public addresses
let arc_addrs = Arc::new(public);
dns_cache.insert(host, port, Arc::clone(&arc_addrs)).await;
Ok((*arc_addrs).clone())
}
/// Validate that the target host:port is allowed.
///
/// Performs port whitelist check, private IP filtering, and DNS resolution
/// with caching. The resolved addresses are stored in the shared DnsCache
/// so that the SafeDnsResolver can reuse them, eliminating the TOCTTOU gap.
pub async fn validate_target(
host: &str,
port: u16,
allowed_ports: &HashSet<u16>,
allow_private: bool,
dns_cache: &DnsCache,
) -> Result<Vec<SocketAddr>, FilterError> {
// Port whitelist check
if !allowed_ports.contains(&port) {
return Err(FilterError::PortNotAllowed(port));
}
// Try parsing as IP directly (no DNS needed)
if let Ok(ip) = host.parse::<IpAddr>() {
if !allow_private && is_private_ip(&ip) {
return Err(FilterError::PrivateIp(ip));
}
return Ok(vec![SocketAddr::new(ip, port)]);
}
// Resolve and validate DNS (populates cache for SafeDnsResolver)
resolve_public_addrs(host, port, allow_private, dns_cache).await
}
#[cfg(test)]
mod tests {
use super::*;
fn ports() -> HashSet<u16> {
[80, 443, 8080, 8443].into_iter().collect()
}
fn cache() -> DnsCache {
DnsCache::new(Duration::from_secs(60), 128)
}
#[test]
fn test_private_ipv4() {
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1))));
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(172, 16, 0, 1))));
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1))));
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1))));
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(169, 254, 1, 1))));
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(0, 0, 0, 0))));
// CGNAT
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(100, 64, 0, 1))));
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(
100, 127, 255, 254
))));
assert!(!is_private_ip(&IpAddr::V4(Ipv4Addr::new(
100, 63, 255, 254
))));
// Benchmark testing
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(198, 18, 0, 1))));
// Reserved
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(240, 0, 0, 1))));
// Public
assert!(!is_private_ip(&IpAddr::V4(Ipv4Addr::new(8, 8, 8, 8))));
assert!(!is_private_ip(&IpAddr::V4(Ipv4Addr::new(203, 0, 113, 1))));
}
#[test]
fn test_private_ipv6() {
assert!(is_private_ip(&IpAddr::V6(Ipv6Addr::LOCALHOST)));
assert!(is_private_ip(&IpAddr::V6(Ipv6Addr::UNSPECIFIED)));
// fc00::1 (ULA)
assert!(is_private_ip(&IpAddr::V6(Ipv6Addr::new(
0xfc00, 0, 0, 0, 0, 0, 0, 1
))));
// fe80::1 (link-local)
assert!(is_private_ip(&IpAddr::V6(Ipv6Addr::new(
0xfe80, 0, 0, 0, 0, 0, 0, 1
))));
}
#[tokio::test]
async fn test_port_not_allowed() {
let cache = cache();
let result = validate_target("8.8.8.8", 22, &ports(), false, &cache).await;
assert!(matches!(result, Err(FilterError::PortNotAllowed(22))));
}
#[tokio::test]
async fn test_private_ip_blocked() {
let cache = cache();
let result = validate_target("127.0.0.1", 80, &ports(), false, &cache).await;
assert!(matches!(result, Err(FilterError::PrivateIp(_))));
}
#[tokio::test]
async fn test_public_ip_allowed() {
let cache = cache();
let result = validate_target("8.8.8.8", 443, &ports(), false, &cache).await;
assert!(result.is_ok());
let addrs = result.unwrap();
assert_eq!(addrs.len(), 1);
assert_eq!(addrs[0].ip(), IpAddr::V4(Ipv4Addr::new(8, 8, 8, 8)));
}
#[tokio::test]
async fn test_private_ip_allowed_when_enabled() {
let cache = cache();
let result = validate_target("127.0.0.1", 80, &ports(), true, &cache).await;
assert!(result.is_ok());
let addrs = result.unwrap();
assert_eq!(
addrs,
vec![SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 80)]
);
}
#[tokio::test]
async fn test_localhost_hostname_blocked_by_default() {
let cache = cache();
let result = validate_target("localhost", 80, &ports(), false, &cache).await;
assert!(matches!(result, Err(FilterError::NoPublicAddrs(_))));
}
#[tokio::test]
async fn test_localhost_hostname_allowed_when_enabled() {
let cache = cache();
let result = validate_target("localhost", 80, &ports(), true, &cache).await;
assert!(result.is_ok());
assert!(!result.unwrap().is_empty());
}
#[tokio::test]
async fn test_cache_stores_multiple_addrs() {
let cache = cache();
let addrs = vec![
SocketAddr::new(IpAddr::V4(Ipv4Addr::new(1, 1, 1, 1)), 443),
SocketAddr::new(IpAddr::V4(Ipv4Addr::new(1, 0, 0, 1)), 443),
];
cache
.insert("example.com", 443, Arc::new(addrs.clone()))
.await;
let cached = cache.get("example.com", 443).await.unwrap();
assert_eq!(*cached, addrs);
}
#[tokio::test]
async fn test_cache_key_case_insensitive() {
let cache = cache();
let addrs = vec![SocketAddr::new(IpAddr::V4(Ipv4Addr::new(1, 1, 1, 1)), 443)];
cache
.insert("Example.COM", 443, Arc::new(addrs.clone()))
.await;
let cached = cache.get("example.com", 443).await.unwrap();
assert_eq!(*cached, addrs);
}
}

View File

@@ -0,0 +1,394 @@
//! WebSocket tunnel client: connect, authenticate, and run the tunnel.
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::net::TcpStream;
use tokio::sync::watch;
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
use tokio_tungstenite::tungstenite::http;
use tokio_tungstenite::tungstenite::protocol::WebSocketConfig;
use tracing::{debug, info, warn};
use crate::egress_proxy::{connect_target_via_proxy, ProxyConnectOptions, UpstreamProxyConfig};
use crate::state::{AppState, ServerContext};
use aether_contracts::tunnel::{CURRENT_TUNNEL_PROTOCOL_VERSION, TUNNEL_PROTOCOL_VERSION_HEADER};
use super::{dispatcher, heartbeat, writer};
/// Outcome of a tunnel session.
pub enum TunnelOutcome {
/// Graceful shutdown requested by the local process.
Shutdown,
/// Remote side disconnected or connection lost — should reconnect.
Disconnected,
}
/// Connect to Aether's WebSocket tunnel endpoint and run until disconnected.
///
/// `conn_idx` identifies which connection in the pool this is (0-based).
/// Only connection 0 sends heartbeats to avoid resetting shared metrics.
pub async fn connect_and_run(
state: &Arc<AppState>,
server: &Arc<ServerContext>,
conn_idx: usize,
shutdown: &mut watch::Receiver<bool>,
drain: watch::Receiver<bool>,
) -> Result<TunnelOutcome, anyhow::Error> {
let ws_url = build_tunnel_url(server);
debug!(url = %ws_url, conn = conn_idx, "connecting tunnel");
// Build WebSocket request with auth headers
let mut request = ws_url.clone().into_client_request()?;
let headers = request.headers_mut();
headers.insert(
"Authorization",
http::HeaderValue::from_str(&format!("Bearer {}", server.management_token))?,
);
headers.insert(
TUNNEL_PROTOCOL_VERSION_HEADER,
http::HeaderValue::from_str(&CURRENT_TUNNEL_PROTOCOL_VERSION.to_string())?,
);
let node_id = server.node_id.read().unwrap().clone();
headers.insert("X-Node-Id", http::HeaderValue::from_str(&node_id)?);
// Use dynamic node_name (may be updated by remote config) instead of
// the static server.node_name, so that remote name changes take effect
// on the next reconnect.
let dynamic_node_name = server.dynamic.load().node_name.clone();
headers.insert(
"X-Node-Name",
http::HeaderValue::from_str(&dynamic_node_name)?,
);
// Advertise per-connection max concurrent streams so the backend can
// respect the proxy's capacity limit.
let max_streams = state.config.tunnel_max_streams.unwrap_or(128);
headers.insert("X-Tunnel-Max-Streams", http::HeaderValue::from(max_streams));
// Parse host:port from URL
let uri: http::Uri = ws_url.parse()?;
let host = uri
.host()
.ok_or_else(|| anyhow::anyhow!("missing host in tunnel URL"))?;
let is_tls = uri.scheme_str() == Some("wss");
let port = uri.port_u16().unwrap_or(if is_tls { 443 } else { 80 });
// TCP connect with timeout
let connect_timeout = state
.config
.tunnel_connect_timeout()
.expect("validated config should resolve tunnel connect timeout");
let tcp_stream = connect_tunnel_tcp(state, host, port, connect_timeout).await?;
// Configure TCP parameters via socket2
configure_tcp_socket(&tcp_stream, state);
// WebSocket upgrade (with TLS if wss://)
let connector = if is_tls {
Some(tokio_tungstenite::Connector::Rustls(Arc::clone(
&state.tunnel_tls_config,
)))
} else {
None
};
// Match Python-side _MAX_FRAME_SIZE (64 MiB) to prevent tungstenite's
// default 16 MiB limit from rejecting large AI API payloads (multi-image
// base64 requests can exceed 16 MiB).
let ws_config = WebSocketConfig {
max_frame_size: Some(64 << 20),
max_message_size: Some(64 << 20),
..Default::default()
};
let handshake_timeout = connect_timeout;
let (ws_stream, _response) = tokio::time::timeout(
handshake_timeout,
tokio_tungstenite::client_async_tls_with_config(
request,
tcp_stream,
Some(ws_config),
connector,
),
)
.await
.map_err(|_| {
anyhow::anyhow!(
"tunnel WebSocket handshake timeout ({}ms)",
handshake_timeout.as_millis()
)
})??;
let stale_timeout = state
.config
.tunnel_stale_timeout()
.expect("validated config should resolve tunnel stale timeout");
let ping_interval = state
.config
.tunnel_ping_interval()
.expect("validated config should resolve tunnel ping interval");
debug!(
conn = conn_idx,
tcp_keepalive_secs = state.config.tunnel_tcp_keepalive_secs,
tcp_nodelay = state.config.tunnel_tcp_nodelay,
connect_timeout_ms = connect_timeout.as_millis(),
stale_timeout_ms = stale_timeout.as_millis(),
ping_interval_ms = ping_interval.as_millis(),
"tunnel connected"
);
server.tunnel_metrics.record_connect_success();
let connected_at = Instant::now();
// NOTE: reconnect_attempts reset is handled by the caller (mod.rs)
// based on how long the connection stayed alive.
// Split into read/write halves
let (ws_sink, ws_read) = futures_util::StreamExt::split(ws_stream);
// Spawn writer task (with WebSocket ping keepalive)
let (frame_tx, mut writer_handle) = writer::spawn_writer_with_metrics(
ws_sink,
ping_interval,
Some(Arc::clone(&server.tunnel_metrics)),
);
let drain_signal = spawn_drain_signal(conn_idx, frame_tx.clone(), drain.clone());
// Spawn heartbeat task (only for primary connection to avoid
// resetting shared atomic metrics via swap(0))
let hb_handle = if conn_idx == 0 {
heartbeat::spawn(
Arc::clone(state),
Arc::clone(server),
frame_tx.clone(),
shutdown.clone(),
)
} else {
heartbeat::spawn_noop()
};
// Run dispatcher (blocks until disconnect or shutdown).
// Also watch for writer exit — if the write half dies (e.g. the peer
// closed the connection) but the read half stays open, dispatcher would
// block forever on `ws_stream.next()`. Monitoring `writer_handle`
// ensures we detect this and trigger a reconnect promptly.
let state_clone = Arc::clone(state);
let server_clone = Arc::clone(server);
let outcome = tokio::select! {
result = dispatcher::run(
state_clone,
server_clone,
ws_read,
frame_tx.clone(),
hb_handle,
drain.clone(),
) => {
match result {
Ok(()) => Ok(TunnelOutcome::Disconnected),
Err(e) => {
server
.tunnel_metrics
.record_error("dispatcher_error", &e.to_string());
Err(e)
}
}
}
writer_result = &mut writer_handle => {
match writer_result {
Ok(()) => warn!("writer task exited normally, triggering reconnect"),
Err(e) => {
if e.is_panic() {
tracing::error!(error = %e, "writer task panicked, triggering reconnect");
server
.tunnel_metrics
.record_error("writer_task_panic", &e.to_string());
} else {
warn!(error = %e, "writer task cancelled, triggering reconnect");
server
.tunnel_metrics
.record_error("writer_task_cancelled", &e.to_string());
}
}
}
Ok(TunnelOutcome::Disconnected)
}
_ = shutdown.changed() => {
debug!("shutdown during tunnel dispatch");
Ok(TunnelOutcome::Shutdown)
}
};
// Drop our sender; the writer will exit once all stream handler clones
// are also dropped (i.e. after they finish their in-flight work).
drop(frame_tx);
if !drain_signal.is_finished() {
drain_signal.abort();
let _ = drain_signal.await;
}
// Wait for the writer task to finish with a generous timeout — the
// dispatcher already waits up to 30s for stream handlers, so 35s here
// covers that plus a small margin.
// Skip if the writer already exited (the select branch that fired).
if !writer_handle.is_finished() {
let _ = tokio::time::timeout(Duration::from_secs(35), writer_handle).await;
}
let connected_for = connected_at.elapsed();
match &outcome {
Ok(TunnelOutcome::Shutdown) => info!(
conn = conn_idx,
connected_duration_ms = connected_for.as_millis() as u64,
close_reason = "shutdown",
"tunnel session ending"
),
Ok(TunnelOutcome::Disconnected) => info!(
conn = conn_idx,
connected_duration_ms = connected_for.as_millis() as u64,
close_reason = "disconnected",
"tunnel session ending"
),
Err(error) => warn!(
conn = conn_idx,
connected_duration_ms = connected_for.as_millis() as u64,
close_reason = "error",
error = %error,
"tunnel session ending"
),
}
server.tunnel_metrics.record_disconnect(connected_for);
debug!("tunnel disconnected");
outcome
}
fn spawn_drain_signal(
conn_idx: usize,
frame_tx: writer::FrameSender,
mut drain: watch::Receiver<bool>,
) -> tokio::task::JoinHandle<()> {
tokio::spawn(async move {
if !*drain.borrow() {
loop {
if drain.changed().await.is_err() {
return;
}
if *drain.borrow() {
break;
}
}
}
debug!(conn = conn_idx, "sending GOAWAY for tunnel drain");
match tokio::time::timeout(
Duration::from_millis(250),
frame_tx.send(super::protocol::Frame::control(
super::protocol::MsgType::GoAway,
bytes::Bytes::new(),
)),
)
.await
{
Ok(Ok(())) => info!(conn = conn_idx, "sent GOAWAY for tunnel drain"),
Ok(Err(error)) => warn!(
conn = conn_idx,
error = ?error,
"failed to queue GOAWAY for tunnel drain"
),
Err(_) => warn!(
conn = conn_idx,
"timed out queueing GOAWAY for tunnel drain"
),
}
})
}
async fn connect_tunnel_tcp(
state: &Arc<AppState>,
host: &str,
port: u16,
connect_timeout: Duration,
) -> Result<TcpStream, anyhow::Error> {
if let Some(proxy_url) = state.config.effective_aether_outbound_proxy_url() {
let proxy = UpstreamProxyConfig::parse(proxy_url)
.map_err(|err| anyhow::anyhow!("Aether outbound proxy URL invalid: {err}"))?;
debug!(
proxy_url = %proxy.redacted_url(),
host = %host,
port = port,
"connecting tunnel via Aether egress proxy"
);
return tokio::time::timeout(
connect_timeout,
connect_target_via_proxy(
&proxy,
host,
port,
ProxyConnectOptions {
connect_timeout,
tcp_nodelay: state.config.tunnel_tcp_nodelay,
tcp_keepalive: (state.config.tunnel_tcp_keepalive_secs > 0)
.then(|| Duration::from_secs(state.config.tunnel_tcp_keepalive_secs)),
},
),
)
.await
.map_err(|_| {
anyhow::anyhow!(
"tunnel outbound proxy TCP connect timeout ({}ms)",
connect_timeout.as_millis()
)
})?
.map_err(anyhow::Error::from);
}
tokio::time::timeout(connect_timeout, TcpStream::connect((host, port)))
.await
.map_err(|_| {
anyhow::anyhow!(
"tunnel TCP connect timeout ({}ms)",
connect_timeout.as_millis()
)
})?
.map_err(anyhow::Error::from)
}
/// Configure TCP keepalive and NODELAY on an established socket.
fn configure_tcp_socket(stream: &TcpStream, state: &Arc<AppState>) {
let sock_ref = socket2::SockRef::from(stream);
if state.config.tunnel_tcp_keepalive_secs > 0 {
let keepalive = socket2::TcpKeepalive::new()
.with_time(Duration::from_secs(state.config.tunnel_tcp_keepalive_secs))
.with_interval(Duration::from_secs(5));
#[cfg(not(target_os = "windows"))]
let keepalive = keepalive.with_retries(3);
if let Err(e) = sock_ref.set_tcp_keepalive(&keepalive) {
warn!(error = %e, "failed to set TCP keepalive on tunnel socket");
}
}
if state.config.tunnel_tcp_nodelay {
if let Err(e) = sock_ref.set_nodelay(true) {
warn!(error = %e, "failed to set TCP_NODELAY on tunnel socket");
}
}
}
/// Build rustls ClientConfig with system root certificates.
pub fn build_tls_config() -> rustls::ClientConfig {
let _ = rustls::crypto::ring::default_provider().install_default();
let root_store =
rustls::RootCertStore::from_iter(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
rustls::ClientConfig::builder()
.with_root_certificates(root_store)
.with_no_client_auth()
}
fn build_tunnel_url(server: &ServerContext) -> String {
let base = server.aether_url.trim_end_matches('/');
let ws_base = if base.starts_with("https://") {
base.replacen("https://", "wss://", 1)
} else if base.starts_with("http://") {
base.replacen("http://", "ws://", 1)
} else {
format!("wss://{}", base)
};
format!("{}/api/internal/proxy-tunnel", ws_base)
}

View File

@@ -0,0 +1,440 @@
//! Frame dispatcher: reads incoming WebSocket frames and routes them.
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use bytes::Bytes;
use futures_util::StreamExt;
use tokio::sync::mpsc;
use tokio::sync::watch;
use tokio::task::JoinHandle;
use tokio_tungstenite::tungstenite::Message;
use tracing::{debug, error, info, warn};
use crate::state::{AppState, ServerContext};
use super::heartbeat::HeartbeatHandle;
use super::protocol::{decompress_if_gzip, Frame, MsgType, RequestMeta};
use super::stream_handler;
use super::writer::FrameSender;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum StreamDispatchStatus {
Delivered,
Closed,
TimedOut,
}
/// Run the dispatcher loop, reading from the WebSocket stream.
pub async fn run<S>(
state: Arc<AppState>,
server: Arc<ServerContext>,
mut ws_stream: S,
frame_tx: FrameSender,
heartbeat: HeartbeatHandle,
mut drain: watch::Receiver<bool>,
) -> Result<(), anyhow::Error>
where
S: StreamExt<Item = Result<Message, tokio_tungstenite::tungstenite::Error>>
+ Unpin
+ Send
+ 'static,
{
// Active streams: stream_id -> body sender
let mut streams: HashMap<u32, mpsc::Sender<Frame>> = HashMap::new();
// Track spawned stream handlers so we can wait for them on shutdown
let mut handler_handles: Vec<JoinHandle<()>> = Vec::new();
let max_streams = state.config.tunnel_max_streams.unwrap_or(128) as usize;
let mut frames_since_cleanup: u32 = 0;
let stale_timeout = state
.config
.tunnel_stale_timeout()
.expect("validated config should resolve tunnel stale timeout");
// Track last time we received any data to detect stale connections
let mut last_data_at = tokio::time::Instant::now();
let mut draining = *drain.borrow();
let read_err = loop {
if draining && streams.is_empty() {
info!("tunnel drained after in-flight streams completed");
break None;
}
let msg_result = tokio::select! {
msg = ws_stream.next() => {
match msg {
Some(r) => r,
None => break None,
}
}
changed = drain.changed() => {
if changed.is_err() {
continue;
}
if *drain.borrow() {
info!("tunnel drain requested, waiting for in-flight streams");
draining = true;
}
continue;
}
_ = tokio::time::sleep_until(last_data_at + stale_timeout) => {
warn!(
stale_ms = stale_timeout.as_millis(),
"tunnel connection stale, no data received"
);
server.tunnel_metrics.record_error(
"stale_timeout",
&format!("no tunnel frame received for {}ms", stale_timeout.as_millis()),
);
break None;
}
};
let msg = match msg_result {
Ok(m) => m,
Err(e) => {
error!(error = %e, "WebSocket read error");
server
.tunnel_metrics
.record_error("ws_read_error", &e.to_string());
break Some(e);
}
};
// Any successfully received message proves the connection is alive
last_data_at = tokio::time::Instant::now();
let data = match msg {
Message::Binary(data) => {
server.tunnel_metrics.record_ws_incoming_frame(data.len());
Bytes::from(data)
}
Message::Ping(_) => continue,
Message::Pong(_) => continue,
Message::Close(_) => {
debug!("received WebSocket close");
break None;
}
_ => continue,
};
let frame = match Frame::decode(data) {
Ok(f) => f,
Err(e) => {
warn!(error = %e, "failed to decode frame");
server
.tunnel_metrics
.record_error("frame_decode_error", &e.to_string());
continue;
}
};
match frame.msg_type {
MsgType::RequestHeaders => {
if draining {
if frame_tx
.try_send(Frame::new(
frame.stream_id,
MsgType::StreamError,
0,
Bytes::from("tunnel draining"),
))
.is_err()
{
warn!(
stream_id = frame.stream_id,
"writer channel full, StreamError dropped during drain"
);
}
continue;
}
// Decompress if the frame is gzip-compressed, then parse metadata
let payload = match decompress_if_gzip(&frame) {
Ok(p) => p,
Err(e) => {
warn!(stream_id = frame.stream_id, error = %e, "frame decompress failed");
continue;
}
};
let meta: RequestMeta = match serde_json::from_slice(&payload) {
Ok(m) => m,
Err(e) => {
warn!(stream_id = frame.stream_id, error = %e, "invalid request metadata");
// Use try_send to avoid blocking the read loop
if frame_tx
.try_send(Frame::new(
frame.stream_id,
MsgType::StreamError,
0,
Bytes::from(format!("invalid request metadata: {e}")),
))
.is_err()
{
warn!(
stream_id = frame.stream_id,
"writer channel full, StreamError dropped"
);
}
continue;
}
};
if streams.len() >= max_streams {
warn!(
stream_id = frame.stream_id,
"max concurrent streams reached"
);
if frame_tx
.try_send(Frame::new(
frame.stream_id,
MsgType::StreamError,
0,
Bytes::from("max concurrent streams reached"),
))
.is_err()
{
warn!(
stream_id = frame.stream_id,
"writer channel full, StreamError dropped"
);
}
continue;
}
// Create body channel and spawn handler
let (body_tx, body_rx) = mpsc::channel::<Frame>(64);
streams.insert(frame.stream_id, body_tx);
let state_clone = Arc::clone(&state);
let server_clone = Arc::clone(&server);
let tx_clone = frame_tx.clone();
let sid = frame.stream_id;
let handle = tokio::spawn(async move {
stream_handler::handle_stream(
state_clone,
server_clone,
sid,
meta,
body_rx,
tx_clone,
)
.await;
});
handler_handles.push(handle);
debug!(stream_id = frame.stream_id, "new stream started");
}
MsgType::RequestBody => {
if let Some(tx) = streams.get(&frame.stream_id).cloned() {
let is_end = frame.is_end_stream();
let sid = frame.stream_id;
let dispatch = dispatch_stream_frame(&tx, frame).await;
if is_end || dispatch != StreamDispatchStatus::Delivered {
streams.remove(&sid);
if dispatch == StreamDispatchStatus::TimedOut {
server.tunnel_metrics.record_error(
"stream_dispatch_timeout",
&format!("request body dispatch timed out for stream {}", sid),
);
try_send_stream_error(
&frame_tx,
sid,
"tunnel request body dispatch stalled",
);
}
if draining && streams.is_empty() {
info!("tunnel drained after request body completion");
break None;
}
}
}
}
MsgType::StreamEnd | MsgType::StreamError => {
// Client-side cancellation or end
if let Some(tx) = streams.remove(&frame.stream_id) {
let _ = dispatch_stream_frame(&tx, frame).await;
if draining && streams.is_empty() {
info!("tunnel drained after stream termination");
break None;
}
}
}
MsgType::Ping => {
// Use try_send to avoid blocking the read loop when writer is congested
if frame_tx
.try_send(Frame::control(MsgType::Pong, frame.payload))
.is_err()
{
warn!("writer channel full, Pong dropped");
}
}
MsgType::HeartbeatAck => {
heartbeat.on_ack(frame.payload).await;
}
MsgType::GoAway => {
info!("received GOAWAY");
break None;
}
_ => {
debug!(msg_type = ?frame.msg_type, "ignoring unexpected frame type");
}
}
// Periodically clean up finished handles to avoid unbounded growth.
// Trigger every 64 frames OR when the count exceeds max_streams.
frames_since_cleanup += 1;
if frames_since_cleanup >= 64 || handler_handles.len() > max_streams {
handler_handles.retain(|h| !h.is_finished());
frames_since_cleanup = 0;
}
};
// Drop body senders so stream handlers waiting on body_rx will unblock
streams.clear();
// Wait for active stream handlers to finish so their frame_tx clones
// are dropped before the writer closes the sink.
drain_handlers(handler_handles).await;
match read_err {
Some(e) => Err(e.into()),
None => Ok(()),
}
}
async fn dispatch_stream_frame(tx: &mpsc::Sender<Frame>, frame: Frame) -> StreamDispatchStatus {
let stream_id = frame.stream_id;
match tokio::time::timeout(stream_frame_dispatch_timeout(), tx.send(frame)).await {
Ok(Ok(())) => StreamDispatchStatus::Delivered,
Ok(Err(_)) => {
warn!(
stream_id,
"stream handler channel closed while dispatching tunnel frame"
);
StreamDispatchStatus::Closed
}
Err(_) => {
warn!(
stream_id,
timeout_ms = stream_frame_dispatch_timeout().as_millis(),
"stream handler channel blocked while dispatching tunnel frame"
);
StreamDispatchStatus::TimedOut
}
}
}
/// Bound how long a single stream handler is allowed to block the shared
/// WebSocket read loop while receiving request-body frames.
fn stream_frame_dispatch_timeout() -> Duration {
#[cfg(test)]
{
Duration::from_millis(25)
}
#[cfg(not(test))]
{
Duration::from_millis(500)
}
}
fn try_send_stream_error(frame_tx: &FrameSender, stream_id: u32, message: &'static str) {
if frame_tx
.try_send(Frame::new(
stream_id,
MsgType::StreamError,
0,
Bytes::from(message),
))
.is_err()
{
warn!(
stream_id,
"writer channel full, StreamError dropped while aborting stalled stream"
);
}
}
/// Wait for all active stream handlers to finish (with a timeout).
async fn drain_handlers(handles: Vec<JoinHandle<()>>) {
if handles.is_empty() {
return;
}
let count = handles.len();
debug!(count, "waiting for active stream handlers to finish");
let _ = tokio::time::timeout(Duration::from_secs(30), async {
for h in handles {
let _ = h.await;
}
})
.await;
}
#[cfg(test)]
mod tests {
use super::*;
use aether_runtime::bounded_queue;
#[tokio::test]
async fn dispatch_stream_frame_times_out_when_handler_stops_draining() {
let (tx, mut rx) = mpsc::channel::<Frame>(1);
tx.send(Frame::new(
7,
MsgType::RequestBody,
0,
Bytes::from_static(b"first"),
))
.await
.expect("first frame should enqueue");
let stalled_send = tokio::spawn({
let tx = tx.clone();
async move {
dispatch_stream_frame(
&tx,
Frame::new(7, MsgType::RequestBody, 0, Bytes::from_static(b"second")),
)
.await
}
});
assert_eq!(
stalled_send.await.expect("dispatch task should join"),
StreamDispatchStatus::TimedOut
);
let retained = rx
.recv()
.await
.expect("queued frame should still be present");
assert_eq!(retained.payload, Bytes::from_static(b"first"));
}
#[tokio::test]
async fn try_send_stream_error_emits_stream_error_frame() {
let (high_tx, mut high_rx) = bounded_queue::<Frame>(4);
let (normal_tx, _normal_rx) = bounded_queue::<Frame>(4);
let frame_tx = FrameSender::from_test_queues(high_tx, normal_tx);
try_send_stream_error(&frame_tx, 9, "tunnel request body dispatch stalled");
let frame = high_rx
.recv()
.await
.expect("stream error frame should enqueue");
assert_eq!(frame.stream_id, 9);
assert_eq!(frame.msg_type, MsgType::StreamError);
assert_eq!(
frame.payload,
Bytes::from_static(b"tunnel request body dispatch stalled")
);
}
}

View File

@@ -0,0 +1,537 @@
//! Tunnel heartbeat: sends metrics over the tunnel, processes ACKs.
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::Duration;
use std::time::SystemTime;
use std::time::UNIX_EPOCH;
use bytes::Bytes;
use tokio::sync::watch;
use tokio::time::Instant;
use tracing::{debug, info, warn};
use crate::registration::client::RemoteConfig;
use crate::runtime;
use crate::state::{AppState, ServerContext, TunnelRequestMetricsSnapshot};
use super::protocol::{Frame, MsgType};
use super::writer::FrameSender;
const CURRENT_VERSION: &str = env!("CARGO_PKG_VERSION");
static UPGRADE_IN_PROGRESS: AtomicBool = AtomicBool::new(false);
static NON_ROOT_UPGRADE_WARNED: AtomicBool = AtomicBool::new(false);
enum AckDecision {
Accept {
heartbeat_id: u64,
upgrade_to: Option<String>,
},
Ignore,
}
/// Handle for the dispatcher to forward HeartbeatAck frames.
#[derive(Clone)]
pub struct HeartbeatHandle {
ack_tx: tokio::sync::mpsc::Sender<Bytes>,
}
impl HeartbeatHandle {
pub async fn on_ack(&self, payload: Bytes) {
let _ = self.ack_tx.send(payload).await;
}
}
/// Create a no-op heartbeat handle that silently discards ACKs.
/// Used for non-primary tunnel connections (conn_idx > 0) to avoid
/// duplicating heartbeat ACK processing.
pub fn spawn_noop() -> HeartbeatHandle {
let (ack_tx, _) = tokio::sync::mpsc::channel::<Bytes>(1);
// receiver is immediately dropped; on_ack() calls will silently fail
HeartbeatHandle { ack_tx }
}
#[derive(Debug, Clone, Copy, Default)]
struct HeartbeatSnapshot {
cumulative: TunnelRequestMetricsSnapshot,
window: TunnelRequestMetricsSnapshot,
}
#[derive(Debug, Clone, Copy)]
struct PendingHeartbeat {
heartbeat_id: u64,
snapshot: HeartbeatSnapshot,
cumulative: TunnelRequestMetricsSnapshot,
sent_at: Option<Instant>,
}
/// Spawn the heartbeat task. Returns a handle for forwarding ACKs.
pub fn spawn(
state: Arc<AppState>,
server: Arc<ServerContext>,
frame_tx: FrameSender,
mut shutdown: watch::Receiver<bool>,
) -> HeartbeatHandle {
let (ack_tx, mut ack_rx) = tokio::sync::mpsc::channel::<Bytes>(4);
tokio::spawn(async move {
// Read initial interval from dynamic config (may be updated by remote config).
let initial_interval = Duration::from_secs(server.dynamic.load().heartbeat_interval);
let mut current_interval = initial_interval;
// At most one in-flight heartbeat snapshot is tracked at a time.
// We keep the last ACKed cumulative snapshot so each payload can
// report both monotonic totals and the delta since the previous ACK.
let mut pending: Option<PendingHeartbeat> = None;
let mut last_acked_snapshot = TunnelRequestMetricsSnapshot::default();
let mut next_heartbeat_id: u64 = 1;
let heartbeat_session_id = format!(
"{}-{}",
std::process::id(),
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_nanos()
);
// Skip first immediate tick by sleeping first.
tokio::time::sleep(current_interval).await;
loop {
tokio::select! {
_ = tokio::time::sleep(current_interval) => {
let pending_entry = if let Some(entry) = pending {
entry
} else {
let cumulative = server.metrics.snapshot();
let id = next_heartbeat_id;
next_heartbeat_id = next_heartbeat_id.wrapping_add(1);
if next_heartbeat_id == 0 {
next_heartbeat_id = 1;
}
let window = cumulative.delta_since(last_acked_snapshot);
let entry = PendingHeartbeat {
heartbeat_id: id,
snapshot: HeartbeatSnapshot { cumulative, window },
cumulative,
sent_at: None,
};
pending = Some(entry);
entry
};
let payload = build_heartbeat_payload(
&state,
&server,
&heartbeat_session_id,
pending_entry.heartbeat_id,
pending_entry.snapshot
).await;
let frame = Frame::control(MsgType::HeartbeatData, payload);
if frame_tx.send(frame).await.is_err() {
break; // Writer closed
}
server.tunnel_metrics.record_heartbeat_sent();
if let Some(mut entry) = pending {
entry.sent_at = Some(Instant::now());
pending = Some(entry);
}
debug!("sent heartbeat data");
// Re-read interval from dynamic config (remote config may have
// updated it since the last heartbeat).
let new_interval = Duration::from_secs(
server.dynamic.load().heartbeat_interval
);
if new_interval != current_interval {
debug!(
old_secs = current_interval.as_secs(),
new_secs = new_interval.as_secs(),
"heartbeat interval updated from dynamic config"
);
current_interval = new_interval;
}
}
Some(ack_payload) = ack_rx.recv() => {
match handle_ack(&server, &ack_payload) {
AckDecision::Accept {
heartbeat_id: ack_id,
upgrade_to,
} => {
if let Some(entry) = pending {
if ack_id == entry.heartbeat_id {
if let Some(sent_at) = entry.sent_at {
server.tunnel_metrics.record_heartbeat_ack(sent_at.elapsed());
}
last_acked_snapshot = entry.cumulative;
pending = None;
}
}
maybe_trigger_upgrade(upgrade_to);
}
AckDecision::Ignore => {}
}
}
_ = shutdown.changed() => {
debug!("heartbeat task shutting down");
break;
}
}
}
});
HeartbeatHandle { ack_tx }
}
async fn build_heartbeat_payload(
state: &AppState,
server: &ServerContext,
heartbeat_session_id: &str,
heartbeat_id: u64,
snapshot: HeartbeatSnapshot,
) -> Bytes {
let node_id = server.node_id.read().unwrap().clone();
let tunnel_snapshot = server.tunnel_metrics.snapshot();
let recent_errors = server.tunnel_metrics.recent_errors(8);
let resource_usage = state.resource_monitor.snapshot();
let cumulative = snapshot.cumulative;
let window = snapshot.window;
let cumulative_metrics = serde_json::json!({
"total_requests": cumulative.total_requests,
"total_latency_ns": cumulative.total_latency_ns,
"avg_latency_ms": cumulative.average_latency_ms(),
"failed_requests": cumulative.failed_requests,
"dns_failures": cumulative.dns_failures,
"stream_errors": cumulative.stream_errors,
"slow_requests": cumulative.slow_requests,
});
let window_metrics = serde_json::json!({
"total_requests": window.total_requests,
"total_latency_ns": window.total_latency_ns,
"avg_latency_ms": window.average_latency_ms(),
"failed_requests": window.failed_requests,
"dns_failures": window.dns_failures,
"stream_errors": window.stream_errors,
"slow_requests": window.slow_requests,
});
let local_admission = state.stream_concurrency_snapshot().map(|snapshot| {
serde_json::json!({
"limit": snapshot.limit,
"in_flight": snapshot.in_flight,
"available_permits": snapshot.available_permits,
"high_watermark": snapshot.high_watermark,
"rejected_total": snapshot.rejected,
})
});
let distributed_admission = match state.distributed_stream_concurrency_snapshot().await {
Ok(Some(snapshot)) => Some(serde_json::json!({
"limit": snapshot.limit,
"in_flight": snapshot.in_flight,
"available_permits": snapshot.available_permits,
"high_watermark": snapshot.high_watermark,
"rejected_total": snapshot.rejected,
})),
Ok(None) => None,
Err(err) => Some(serde_json::json!({
"error": err.to_string(),
})),
};
let admission = match (local_admission, distributed_admission) {
(None, None) => None,
(local, distributed) => Some(serde_json::json!({
"local_streams": local,
"distributed_streams": distributed,
})),
};
let payload = serde_json::json!({
"node_id": node_id,
"heartbeat_session_id": heartbeat_session_id,
"heartbeat_id": heartbeat_id,
"heartbeat_interval": server.dynamic.load().heartbeat_interval,
"active_connections": server.active_connections.load(Ordering::Acquire),
"total_requests": cumulative.total_requests,
"avg_latency_ms": cumulative.average_latency_ms(),
"failed_requests": cumulative.failed_requests,
"dns_failures": cumulative.dns_failures,
"stream_errors": cumulative.stream_errors,
"slow_requests": cumulative.slow_requests,
"window_total_requests": window.total_requests,
"window_total_latency_ns": window.total_latency_ns,
"window_avg_latency_ms": window.average_latency_ms(),
"window_failed_requests": window.failed_requests,
"window_dns_failures": window.dns_failures,
"window_stream_errors": window.stream_errors,
"window_slow_requests": window.slow_requests,
"proxy_metrics": {
"cumulative": cumulative_metrics,
"window": window_metrics,
},
"proxy_metadata": {
"version": CURRENT_VERSION,
"admission": admission,
"resource_usage": resource_usage,
"tunnel_metrics": {
"connect_attempts": tunnel_snapshot.connect_attempts,
"connect_successes": tunnel_snapshot.connect_successes,
"connect_errors": tunnel_snapshot.connect_errors,
"disconnects": tunnel_snapshot.disconnects,
"last_connected_at_unix_secs": tunnel_snapshot.last_connected_at_unix_secs,
"last_disconnected_at_unix_secs": tunnel_snapshot.last_disconnected_at_unix_secs,
"last_connected_duration_ms": tunnel_snapshot.last_connected_duration_ms,
"connected_duration_total_ms": tunnel_snapshot.connected_duration_total_ms,
"heartbeat_sent": tunnel_snapshot.heartbeat_sent,
"heartbeat_ack": tunnel_snapshot.heartbeat_ack,
"heartbeat_rtt_last_ms": tunnel_snapshot.heartbeat_rtt_last_ms,
"heartbeat_rtt_avg_ms": tunnel_snapshot.heartbeat_rtt_avg_ms(),
"ws_in_frames": tunnel_snapshot.ws_in_frames,
"ws_in_bytes": tunnel_snapshot.ws_in_bytes,
"ws_out_frames": tunnel_snapshot.ws_out_frames,
"ws_out_bytes": tunnel_snapshot.ws_out_bytes,
"error_events_total": tunnel_snapshot.error_events_total,
},
"recent_tunnel_errors": recent_errors,
},
});
Bytes::from(serde_json::to_vec(&payload).unwrap_or_default())
}
fn handle_ack(server: &ServerContext, payload: &[u8]) -> AckDecision {
if payload.is_empty() {
warn!("received empty heartbeat ACK");
server
.tunnel_metrics
.record_error("heartbeat_ack_empty", "received empty heartbeat ACK");
return AckDecision::Ignore;
}
#[derive(serde::Deserialize)]
struct AckPayload {
#[serde(default)]
remote_config: Option<RemoteConfig>,
#[serde(default)]
config_version: u64,
heartbeat_id: u64,
#[serde(default)]
upgrade_to: Option<String>,
}
match serde_json::from_slice::<AckPayload>(payload) {
Ok(ack) => {
if let Some(ref rc) = ack.remote_config {
runtime::apply_remote_config(&server.dynamic, rc, ack.config_version);
}
AckDecision::Accept {
heartbeat_id: ack.heartbeat_id,
upgrade_to: ack.upgrade_to.and_then(normalize_upgrade_target),
}
}
Err(e) => {
warn!(error = %e, "failed to parse heartbeat ACK");
server
.tunnel_metrics
.record_error("heartbeat_ack_parse", &e.to_string());
AckDecision::Ignore
}
}
}
fn normalize_upgrade_target(raw: String) -> Option<String> {
let trimmed = raw.trim();
if trimmed.is_empty() {
return None;
}
let normalized = trimmed
.strip_prefix("tunnel-v")
.or_else(|| trimmed.strip_prefix("proxy-v"))
.unwrap_or(trimmed);
if normalized == CURRENT_VERSION {
return None;
}
Some(normalized.to_string())
}
fn maybe_trigger_upgrade(version: Option<String>) {
let Some(target_version) = version else {
return;
};
if !crate::setup::service::is_root() {
if NON_ROOT_UPGRADE_WARNED
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
.is_ok()
{
warn!(
target_version = %target_version,
"remote upgrade skipped: root privileges are required"
);
}
return;
}
if UPGRADE_IN_PROGRESS
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
.is_err()
{
debug!(target_version = %target_version, "upgrade already in progress, ignoring");
return;
}
tokio::spawn(async move {
info!(target_version = %target_version, "received remote upgrade instruction");
match crate::setup::upgrade::perform_upgrade(&target_version).await {
Ok(()) => {
info!(target_version = %target_version, "remote upgrade finished");
}
Err(e) => {
warn!(
target_version = %target_version,
error = %e,
"remote upgrade failed"
);
UPGRADE_IN_PROGRESS.store(false, Ordering::Release);
}
}
});
}
#[cfg(test)]
mod tests {
use std::sync::atomic::AtomicU64;
use std::sync::{Arc, RwLock};
use arc_swap::ArcSwap;
use clap::Parser;
use super::{build_heartbeat_payload, handle_ack, AckDecision, HeartbeatSnapshot};
use crate::registration::client::AetherClient;
use crate::runtime::DynamicConfig;
use crate::state::{AppState, ServerContext, TunnelMetrics, TunnelRequestMetrics};
fn sample_config() -> Arc<crate::config::Config> {
Arc::new(crate::config::Config::parse_from([
"aether-tunnel",
"--aether-url",
"https://example.com",
"--management-token",
"ae_test",
"--node-name",
"tunnel-test",
]))
}
fn sample_server() -> Arc<ServerContext> {
let config = sample_config();
Arc::new(ServerContext {
server_label: "heartbeat-test".to_string(),
aether_url: config.aether_url.clone(),
management_token: config.management_token.clone(),
node_name: config.node_name.clone(),
node_id: Arc::new(RwLock::new("node-123".to_string())),
aether_client: Arc::new(AetherClient::new(
&config,
&config.aether_url,
&config.management_token,
)),
dynamic: Arc::new(ArcSwap::from_pointee(DynamicConfig::from_config(&config))),
active_connections: Arc::new(AtomicU64::new(0)),
metrics: Arc::new(TunnelRequestMetrics::new()),
tunnel_metrics: Arc::new(TunnelMetrics::new()),
})
}
fn sample_state(config: Arc<crate::config::Config>) -> AppState {
let dns_cache = Arc::new(crate::target_filter::DnsCache::new(
std::time::Duration::from_secs(config.dns_cache_ttl_secs),
config.dns_cache_capacity,
));
AppState {
config: Arc::clone(&config),
dns_cache: Arc::clone(&dns_cache),
upstream_client_pool: crate::upstream_client::UpstreamClientPool::new(
config, dns_cache,
),
tunnel_tls_config: Arc::new(crate::tunnel::client::build_tls_config()),
resource_monitor: Arc::new(crate::hardware::RuntimeResourceMonitor::new()),
stream_gate: None,
distributed_stream_gate: None,
}
}
#[test]
fn heartbeat_ack_requires_heartbeat_id() {
let server = sample_server();
let decision = handle_ack(
&server,
br#"{"config_version":1,"remote_config":{"heartbeat_interval":9}}"#,
);
assert!(matches!(decision, AckDecision::Ignore));
assert_eq!(server.dynamic.load().heartbeat_interval, 5);
}
#[test]
fn heartbeat_ack_applies_remote_config_with_heartbeat_id() {
let server = sample_server();
let decision = handle_ack(
&server,
br#"{"heartbeat_id":7,"config_version":1,"remote_config":{"heartbeat_interval":9}}"#,
);
assert!(matches!(
decision,
AckDecision::Accept {
heartbeat_id: 7,
upgrade_to: None
}
));
assert_eq!(server.dynamic.load().heartbeat_interval, 9);
}
#[tokio::test]
async fn heartbeat_payload_reports_resource_usage_and_tunnel_error_diagnostics() {
let config = sample_config();
let server = sample_server();
server
.tunnel_metrics
.record_error("ws_write_error", "IO error: Connection reset by peer");
let state = sample_state(config);
let payload = build_heartbeat_payload(
&state,
&server,
"session-1",
42,
HeartbeatSnapshot::default(),
)
.await;
let payload: serde_json::Value =
serde_json::from_slice(&payload).expect("heartbeat payload should be JSON");
let resource_usage = payload
.pointer("/proxy_metadata/resource_usage")
.and_then(serde_json::Value::as_object)
.expect("resource usage should be reported");
assert!(resource_usage.contains_key("system_cpu_usage_percent"));
assert!(resource_usage.contains_key("process_memory_bytes"));
let recent_error = payload
.pointer("/proxy_metadata/recent_tunnel_errors/0")
.and_then(serde_json::Value::as_object)
.expect("recent tunnel error should be reported");
assert!(recent_error
.get("timestamp_unix_ms")
.and_then(serde_json::Value::as_u64)
.is_some());
assert_eq!(
recent_error
.get("component")
.and_then(serde_json::Value::as_str),
Some("tunnel_write")
);
assert_eq!(
recent_error
.get("severity")
.and_then(serde_json::Value::as_str),
Some("error")
);
}
}

View File

@@ -0,0 +1,562 @@
pub mod client;
pub mod dispatcher;
pub mod heartbeat;
pub mod protocol;
pub mod stream_handler;
pub mod writer;
use std::sync::Arc;
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
use tokio::sync::watch;
use tracing::{debug, error, info};
use crate::state::{AppState, ServerContext};
/// If a tunnel stays connected at least this long, treat the next disconnect
/// as a non-failure and reset reconnect backoff.
const STABLE_SESSION_RESET_AFTER: Duration = Duration::from_secs(30);
/// Startup staggering step per secondary connection, used to avoid
/// simultaneous bursts when a pool of tunnels starts together.
const STARTUP_STAGGER_STEP_MS: u64 = 150;
/// Upper bound for startup staggering.
const MAX_STARTUP_STAGGER_MS: u64 = 1_500;
/// Keep a tiny floor for repeated reconnects; first retry is still immediate.
const MIN_RECONNECT_DELAY_MS: u64 = 50;
/// Even under sustained failures, keep probing frequently so recovery is fast
/// once cross-border network quality improves.
const RECONNECT_PROBE_MAX_DELAY_MS: u64 = 3_000;
/// Run the tunnel mode main loop (connect, dispatch, reconnect).
///
/// `conn_idx` identifies which connection in the pool this is (0-based).
/// Only connection 0 sends heartbeats to avoid resetting shared metrics.
pub async fn run(
state: &Arc<AppState>,
server: &Arc<ServerContext>,
conn_idx: usize,
mut shutdown: watch::Receiver<bool>,
mut drain: watch::Receiver<bool>,
) {
info!(server = %server.server_label, conn = conn_idx, "starting tunnel");
let reconnect_salt = compute_connection_salt(server, conn_idx);
if *drain.borrow() {
info!(server = %server.server_label, conn = conn_idx, "tunnel drain requested before startup");
return;
}
let startup_delay = compute_startup_stagger(conn_idx, reconnect_salt);
if !startup_delay.is_zero() {
info!(
server = %server.server_label,
conn = conn_idx,
delay_ms = startup_delay.as_millis(),
"startup stagger before first connect"
);
tokio::select! {
_ = tokio::time::sleep(startup_delay) => {}
_ = shutdown.changed() => {
info!(server = %server.server_label, conn = conn_idx, "shutdown requested during startup stagger");
return;
}
_ = drain.changed() => {
if *drain.borrow() {
info!(server = %server.server_label, conn = conn_idx, "tunnel drain requested during startup stagger");
return;
}
}
}
}
let mut consecutive_failures: u32 = 0;
loop {
if *drain.borrow() {
info!(server = %server.server_label, conn = conn_idx, "tunnel drained, exiting slot");
return;
}
server.tunnel_metrics.record_connect_attempt();
let started_at = Instant::now();
match client::connect_and_run(state, server, conn_idx, &mut shutdown, drain.clone()).await {
Ok(client::TunnelOutcome::Shutdown) => {
info!(server = %server.server_label, conn = conn_idx, "tunnel shut down gracefully");
return;
}
Ok(client::TunnelOutcome::Disconnected) => {
debug!(server = %server.server_label, conn = conn_idx, "tunnel disconnected, reconnecting");
}
Err(e) => {
server.tunnel_metrics.record_connect_error();
server
.tunnel_metrics
.record_error("tunnel_connect_error", &e.to_string());
error!(server = %server.server_label, conn = conn_idx, error = %e, "tunnel connection error, reconnecting");
}
}
if *shutdown.borrow() {
info!(server = %server.server_label, conn = conn_idx, "shutdown requested, not reconnecting");
return;
}
if *drain.borrow() {
info!(server = %server.server_label, conn = conn_idx, "tunnel drained after disconnect");
return;
}
// Reset backoff after a stable session to keep recovery snappy when
// failures are only occasional.
let connected_for = started_at.elapsed();
if connected_for >= STABLE_SESSION_RESET_AFTER {
consecutive_failures = 0;
} else {
consecutive_failures = consecutive_failures.saturating_add(1);
}
let reconnect_delay = compute_reconnect_delay(
state.config.tunnel_reconnect_base_ms,
state.config.tunnel_reconnect_max_ms,
consecutive_failures,
reconnect_salt,
);
if reconnect_delay.is_zero() && consecutive_failures <= 1 {
debug!(
server = %server.server_label,
conn = conn_idx,
failures = consecutive_failures,
delay_ms = reconnect_delay.as_millis(),
"waiting before reconnect"
);
} else {
info!(
server = %server.server_label,
conn = conn_idx,
failures = consecutive_failures,
delay_ms = reconnect_delay.as_millis(),
"waiting before reconnect"
);
}
tokio::select! {
_ = tokio::time::sleep(reconnect_delay) => {}
_ = shutdown.changed() => {
info!(server = %server.server_label, conn = conn_idx, "shutdown requested during reconnect wait");
return;
}
_ = drain.changed() => {
if *drain.borrow() {
info!(server = %server.server_label, conn = conn_idx, "tunnel drain requested during reconnect wait");
return;
}
}
}
}
}
fn compute_connection_salt(server: &ServerContext, conn_idx: usize) -> u64 {
// FNV-1a style hash over server label + connection index.
let mut h: u64 = 0xcbf29ce484222325;
for &b in server.server_label.as_bytes() {
h ^= b as u64;
h = h.wrapping_mul(0x100000001b3);
}
h ^= conn_idx as u64;
mix_u64(h)
}
fn compute_startup_stagger(conn_idx: usize, salt: u64) -> Duration {
if conn_idx == 0 {
return Duration::ZERO;
}
let base = (conn_idx as u64).saturating_mul(STARTUP_STAGGER_STEP_MS);
let jitter = mix_u64(salt) % 301; // 0..=300ms
Duration::from_millis((base + jitter).min(MAX_STARTUP_STAGGER_MS))
}
fn compute_reconnect_delay(
base_ms: u64,
max_ms: u64,
consecutive_failures: u32,
salt: u64,
) -> Duration {
// First retry should be immediate to maximize recovery speed on transient
// blips (the user's primary expectation in poor networks).
if consecutive_failures <= 1 {
return Duration::ZERO;
}
// Keep a sane minimum for repeated failures.
let base_ms = base_ms.max(MIN_RECONNECT_DELAY_MS);
let max_ms = max_ms.max(base_ms);
let cap_ms = compute_reconnect_cap_ms(base_ms, max_ms, consecutive_failures)
.min(RECONNECT_PROBE_MAX_DELAY_MS.max(base_ms));
// Equal-jitter: randomize in [cap/2, cap], preventing synchronized reconnect
// storms while keeping reconnect latency bounded.
if cap_ms <= 1 {
return Duration::from_millis(cap_ms);
}
let half = cap_ms / 2;
let span = cap_ms - half;
let now_nanos = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.subsec_nanos() as u64)
.unwrap_or(0);
let mixed = mix_u64(now_nanos ^ salt);
let jitter = if span == 0 { 0 } else { mixed % (span + 1) };
Duration::from_millis(half + jitter)
}
fn compute_reconnect_cap_ms(base_ms: u64, max_ms: u64, consecutive_failures: u32) -> u64 {
if consecutive_failures <= 1 {
return base_ms.min(max_ms);
}
let shift = (consecutive_failures - 1).min(31);
let factor = 1u64 << shift;
base_ms.saturating_mul(factor).min(max_ms)
}
fn mix_u64(mut x: u64) -> u64 {
// SplitMix64 finalizer - cheap bit mixing for pseudo-random jitter.
x ^= x >> 30;
x = x.wrapping_mul(0xbf58476d1ce4e5b9);
x ^= x >> 27;
x = x.wrapping_mul(0x94d049bb133111eb);
x ^ (x >> 31)
}
#[cfg(test)]
mod tests {
use std::sync::atomic::AtomicU64;
use std::sync::{Arc, Once};
use std::time::Duration;
use aether_gateway::{build_router_with_state, AppState as GatewayAppState};
use arc_swap::ArcSwap;
use axum::Router;
use reqwest::StatusCode;
use tokio::sync::watch;
use crate::config::Config;
use crate::registration::client::AetherClient;
use crate::runtime::DynamicConfig;
use crate::state::{
AppState as TunnelAppState, ServerContext, TunnelMetrics, TunnelRequestMetrics,
};
use crate::target_filter::DnsCache;
use crate::tunnel::protocol;
use crate::upstream_client;
use super::{
compute_reconnect_cap_ms, compute_reconnect_delay, compute_startup_stagger, run,
MAX_STARTUP_STAGGER_MS, RECONNECT_PROBE_MAX_DELAY_MS, STARTUP_STAGGER_STEP_MS,
};
#[test]
fn reconnect_cap_grows_exponentially_and_caps() {
let base = 500;
let max = 30_000;
assert_eq!(compute_reconnect_cap_ms(base, max, 0), 500);
assert_eq!(compute_reconnect_cap_ms(base, max, 1), 500);
assert_eq!(compute_reconnect_cap_ms(base, max, 2), 1_000);
assert_eq!(compute_reconnect_cap_ms(base, max, 3), 2_000);
assert_eq!(compute_reconnect_cap_ms(base, max, 4), 4_000);
assert_eq!(compute_reconnect_cap_ms(base, max, 5), 8_000);
assert_eq!(compute_reconnect_cap_ms(base, max, 6), 16_000);
assert_eq!(compute_reconnect_cap_ms(base, max, 7), 30_000);
assert_eq!(compute_reconnect_cap_ms(base, max, 20), 30_000);
}
#[test]
fn startup_stagger_is_zero_for_primary_and_bounded_for_secondary() {
assert_eq!(compute_startup_stagger(0, 42), Duration::ZERO);
let d1 = compute_startup_stagger(1, 42);
let d2 = compute_startup_stagger(2, 42);
assert!(d1 >= Duration::from_millis(STARTUP_STAGGER_STEP_MS));
assert!(d1 <= Duration::from_millis(MAX_STARTUP_STAGGER_MS));
assert!(d2 >= Duration::from_millis(STARTUP_STAGGER_STEP_MS * 2));
assert!(d2 <= Duration::from_millis(MAX_STARTUP_STAGGER_MS));
}
#[test]
fn reconnect_delay_is_immediate_on_first_failure() {
assert_eq!(compute_reconnect_delay(700, 45_000, 1, 123), Duration::ZERO);
}
#[test]
fn reconnect_delay_stays_within_probe_ceiling_after_many_failures() {
let d = compute_reconnect_delay(500, 45_000, 100, 12345);
assert!(d <= Duration::from_millis(RECONNECT_PROBE_MAX_DELAY_MS));
}
#[tokio::test]
async fn tunnel_reconnects_after_gateway_restart() {
ensure_rustls_provider();
let gateway_port = reserve_local_port().expect("gateway port should reserve");
let gateway_base_url = format!("http://127.0.0.1:{gateway_port}");
let (gateway_state, mut gateway_handle) = start_gateway_on_port(gateway_port)
.await
.expect("gateway should start");
let state = sample_state(sample_config(&gateway_base_url));
let server = sample_server(&state, "node-recovery");
let (shutdown_tx, shutdown_rx) = watch::channel(false);
let tunnel_task = tokio::spawn({
let state = Arc::clone(&state);
let server = Arc::clone(&server);
let (_drain_tx, drain_rx) = watch::channel(false);
async move {
run(&state, &server, 0, shutdown_rx, drain_rx).await;
}
});
wait_until_relay_status(
&gateway_base_url,
"node-recovery",
StatusCode::GATEWAY_TIMEOUT,
)
.await;
assert_eq!(gateway_state.force_close_all_tunnel_proxies(), 1);
tokio::time::sleep(Duration::from_millis(200)).await;
gateway_handle.abort();
let (_restarted_gateway_state, restarted_gateway_handle) =
start_gateway_on_port_retry(gateway_port)
.await
.expect("gateway should restart on fixed port");
gateway_handle = restarted_gateway_handle;
wait_until_relay_status(
&gateway_base_url,
"node-recovery",
StatusCode::GATEWAY_TIMEOUT,
)
.await;
let _ = shutdown_tx.send(true);
tokio::time::timeout(Duration::from_secs(5), tunnel_task)
.await
.expect("tunnel task should stop")
.expect("tunnel task should join");
gateway_handle.abort();
}
async fn wait_until_relay_status(gateway_base_url: &str, node_id: &str, expected: StatusCode) {
let deadline = tokio::time::Instant::now() + Duration::from_secs(10);
let mut last_observed = None::<String>;
loop {
if let Some((status, body)) = probe_relay_status(gateway_base_url, node_id).await {
last_observed = Some(format!("{status} body={body}"));
if status == expected {
return;
}
}
assert!(
tokio::time::Instant::now() < deadline,
"relay status did not become {expected} within timeout; last={:?}",
last_observed
);
tokio::time::sleep(Duration::from_millis(25)).await;
}
}
async fn probe_relay_status(
gateway_base_url: &str,
node_id: &str,
) -> Option<(StatusCode, String)> {
let response = reqwest::Client::new()
.post(format!(
"{gateway_base_url}/api/internal/tunnel/relay/{node_id}"
))
.header("content-type", "application/octet-stream")
.body(relay_probe_envelope())
.send()
.await
.ok()?;
let status = response.status();
let body = response.text().await.unwrap_or_default();
Some((status, body))
}
fn relay_probe_envelope() -> Vec<u8> {
let meta = protocol::RequestMeta {
provider_id: None,
endpoint_id: None,
key_id: None,
method: "GET".to_string(),
url: "http://127.0.0.1:80/blocked".to_string(),
headers: std::collections::HashMap::new(),
timeout: 5,
follow_redirects: None,
http1_only: false,
transport_profile: None,
};
let meta_json =
serde_json::to_vec(&meta).expect("tunnel relay probe metadata should serialize");
let mut envelope = Vec::with_capacity(4 + meta_json.len());
envelope.extend_from_slice(&(meta_json.len() as u32).to_be_bytes());
envelope.extend_from_slice(&meta_json);
envelope
}
async fn start_gateway_on_port(
port: u16,
) -> Result<(GatewayAppState, tokio::task::JoinHandle<()>), std::io::Error> {
let state = GatewayAppState::new().expect("gateway test state should build");
let router = build_router_with_state(state.clone());
let handle = spawn_router_on_port(port, router).await?;
Ok((state, handle))
}
async fn start_gateway_on_port_retry(
port: u16,
) -> Result<(GatewayAppState, tokio::task::JoinHandle<()>), std::io::Error> {
let mut attempts = 0usize;
loop {
match start_gateway_on_port(port).await {
Ok(server) => return Ok(server),
Err(err) => {
attempts += 1;
if attempts >= 20 {
return Err(err);
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
}
}
}
async fn spawn_router_on_port(
port: u16,
app: Router,
) -> Result<tokio::task::JoinHandle<()>, std::io::Error> {
let listener = tokio::net::TcpListener::bind(("127.0.0.1", port)).await?;
Ok(tokio::spawn(async move {
axum::serve(
listener,
app.into_make_service_with_connect_info::<std::net::SocketAddr>(),
)
.await
.expect("gateway test server should run");
}))
}
fn reserve_local_port() -> Result<u16, std::io::Error> {
let listener = std::net::TcpListener::bind("127.0.0.1:0")?;
let port = listener.local_addr()?.port();
drop(listener);
Ok(port)
}
fn sample_state(config: Config) -> Arc<TunnelAppState> {
let config = Arc::new(config);
let dns_cache = Arc::new(DnsCache::new(Duration::from_secs(60), 128));
let upstream_client_pool =
upstream_client::UpstreamClientPool::new(Arc::clone(&config), Arc::clone(&dns_cache));
Arc::new(TunnelAppState {
config,
dns_cache,
upstream_client_pool,
tunnel_tls_config: Arc::new(crate::tunnel::client::build_tls_config()),
resource_monitor: Arc::new(crate::hardware::RuntimeResourceMonitor::new()),
stream_gate: None,
distributed_stream_gate: None,
})
}
fn sample_server(state: &Arc<TunnelAppState>, node_id: &str) -> Arc<ServerContext> {
let config = Arc::clone(&state.config);
Arc::new(ServerContext {
server_label: "gateway-owned-tunnel".to_string(),
aether_url: config.aether_url.clone(),
management_token: config.management_token.clone(),
node_name: config.node_name.clone(),
node_id: Arc::new(std::sync::RwLock::new(node_id.to_string())),
aether_client: Arc::new(AetherClient::new(
&config,
&config.aether_url,
&config.management_token,
)),
dynamic: Arc::new(ArcSwap::from_pointee(DynamicConfig::from_config(&config))),
active_connections: Arc::new(AtomicU64::new(0)),
metrics: Arc::new(TunnelRequestMetrics::new()),
tunnel_metrics: Arc::new(TunnelMetrics::new()),
})
}
fn sample_config(aether_url: &str) -> Config {
Config {
aether_url: aether_url.to_string(),
management_token: "token".to_string(),
public_ip: None,
node_name: "tunnel-test".to_string(),
node_region: None,
heartbeat_interval: 1,
allowed_ports: vec![80, 443],
allow_private_targets: false,
aether_request_timeout_secs: 10,
aether_connect_timeout_secs: 2,
aether_pool_max_idle_per_host: 8,
aether_pool_idle_timeout_secs: 90,
aether_tcp_keepalive_secs: 60,
aether_tcp_nodelay: true,
aether_http2: true,
aether_outbound_proxy_url: None,
aether_retry_max_attempts: 1,
aether_retry_base_delay_ms: 50,
aether_retry_max_delay_ms: 100,
diagnostics_bind: None,
max_concurrent_connections: None,
max_in_flight_streams: None,
distributed_stream_limit: None,
distributed_stream_redis_url: None,
distributed_stream_redis_key_prefix: None,
distributed_stream_lease_ttl_ms: 30_000,
distributed_stream_renew_interval_ms: 10_000,
distributed_stream_command_timeout_ms: 1_000,
dns_cache_ttl_secs: 60,
dns_cache_capacity: 128,
upstream_connect_timeout_secs: 30,
upstream_pool_max_idle_per_host: 4,
upstream_pool_idle_timeout_secs: 60,
upstream_tcp_keepalive_secs: 60,
upstream_tcp_nodelay: true,
upstream_proxy_url: None,
redirect_replay_budget_bytes: crate::config::DEFAULT_REDIRECT_REPLAY_BUDGET_BYTES,
emit_proxy_timing_header: true,
log_level: "info".to_string(),
log_destination: crate::config::TunnelLogDestinationArg::Stdout,
log_dir: None,
log_rotation: crate::config::TunnelLogRotationArg::Daily,
log_retention_days: 7,
log_max_files: 30,
tunnel_reconnect_base_ms: 50,
tunnel_reconnect_max_ms: 250,
tunnel_ping_interval_ms: 1_000,
tunnel_max_streams: Some(8),
tunnel_connect_timeout_ms: 2_000,
tunnel_tcp_keepalive_secs: 30,
tunnel_tcp_nodelay: true,
tunnel_stale_timeout_ms: 5_000,
tunnel_connections: Some(1),
tunnel_connections_max: Some(1),
tunnel_scale_check_interval_ms: 1_000,
tunnel_scale_up_threshold_percent: 70,
tunnel_scale_down_threshold_percent: 35,
tunnel_scale_down_grace_secs: 15,
}
}
fn ensure_rustls_provider() {
static INIT: Once = Once::new();
INIT.call_once(|| {
let _ = rustls::crypto::ring::default_provider().install_default();
});
}
}

View File

@@ -0,0 +1 @@
pub use aether_contracts::tunnel::*;

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,308 @@
//! Dedicated WebSocket writer task.
//!
//! All frame writes go through an mpsc channel to a single writer task,
//! avoiding contention on the WebSocket sink. The writer also sends
//! periodic WebSocket Ping frames to keep the connection alive through
//! intermediary proxies (Nginx, Cloudflare, etc.).
use std::sync::Arc;
use std::time::Duration;
use aether_contracts::tunnel::{MsgType, HEADER_SIZE};
#[cfg(test)]
use aether_runtime::QueueSnapshot;
use aether_runtime::{bounded_queue, BoundedQueueSender, QueueSendError};
use futures_util::SinkExt;
use tokio::task::JoinHandle;
use tokio_tungstenite::tungstenite::Message;
use tracing::{debug, error, trace};
use crate::state::TunnelMetrics;
use super::protocol::Frame;
const HIGH_PRIORITY_QUEUE_CAPACITY: usize = 64;
const NORMAL_PRIORITY_QUEUE_CAPACITY: usize = 256;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum FramePriority {
High,
Normal,
}
#[cfg(test)]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct FrameQueueSnapshots {
pub high: QueueSnapshot,
pub normal: QueueSnapshot,
}
/// Sender half — cloned by stream handlers and heartbeat.
#[derive(Debug, Clone)]
pub struct FrameSender {
high_tx: BoundedQueueSender<Frame>,
normal_tx: BoundedQueueSender<Frame>,
}
impl FrameSender {
pub async fn send(&self, frame: Frame) -> Result<(), QueueSendError<Frame>> {
match classify_frame_priority(&frame) {
FramePriority::High => self.high_tx.send(frame).await,
FramePriority::Normal => self.normal_tx.send(frame).await,
}
}
pub fn try_send(&self, frame: Frame) -> Result<(), QueueSendError<Frame>> {
match classify_frame_priority(&frame) {
FramePriority::High => self.high_tx.try_send(frame),
FramePriority::Normal => self.normal_tx.try_send(frame),
}
}
#[cfg(test)]
pub fn snapshots(&self) -> FrameQueueSnapshots {
FrameQueueSnapshots {
high: self.high_tx.snapshot(),
normal: self.normal_tx.snapshot(),
}
}
#[cfg(test)]
pub(crate) fn from_test_queues(
high_tx: BoundedQueueSender<Frame>,
normal_tx: BoundedQueueSender<Frame>,
) -> Self {
Self { high_tx, normal_tx }
}
}
/// Spawn the writer task. Returns the sender and a JoinHandle for cleanup.
///
/// `ping_interval` controls WebSocket-level Ping frequency (typically 15s).
/// This keeps the connection alive through intermediary proxies/load-balancers.
#[cfg(test)]
pub fn spawn_writer<S>(sink: S, ping_interval: Duration) -> (FrameSender, JoinHandle<()>)
where
S: SinkExt<Message, Error = tokio_tungstenite::tungstenite::Error> + Unpin + Send + 'static,
{
spawn_writer_with_metrics(sink, ping_interval, None)
}
/// Spawn the writer task with optional tunnel metrics instrumentation.
pub fn spawn_writer_with_metrics<S>(
mut sink: S,
ping_interval: Duration,
tunnel_metrics: Option<Arc<TunnelMetrics>>,
) -> (FrameSender, JoinHandle<()>)
where
S: SinkExt<Message, Error = tokio_tungstenite::tungstenite::Error> + Unpin + Send + 'static,
{
let (high_tx, mut high_rx) = bounded_queue::<Frame>(HIGH_PRIORITY_QUEUE_CAPACITY);
let (normal_tx, mut normal_rx) = bounded_queue::<Frame>(NORMAL_PRIORITY_QUEUE_CAPACITY);
let tx = FrameSender { high_tx, normal_tx };
let handle = tokio::spawn(async move {
let mut ping_ticker = tokio::time::interval(ping_interval);
let mut high_open = true;
let mut normal_open = true;
ping_ticker.tick().await; // skip first immediate tick
loop {
if let Ok(frame) = high_rx.try_recv() {
if !write_frame(&mut sink, frame, tunnel_metrics.as_deref()).await {
break;
}
continue;
}
if !high_open && !normal_open {
break;
}
tokio::select! {
biased;
frame = high_rx.recv(), if high_open => {
match frame {
Some(frame) => {
if !write_frame(&mut sink, frame, tunnel_metrics.as_deref()).await {
break;
}
}
None => high_open = false,
}
}
_ = ping_ticker.tick(), if high_open || normal_open => {
if let Err(e) = sink.send(Message::Ping(vec![])).await {
error!(error = %e, "failed to send WebSocket ping");
if let Some(metrics) = tunnel_metrics.as_deref() {
metrics.record_error("ws_ping_error", &e.to_string());
}
break;
}
trace!("sent WebSocket ping");
}
frame = normal_rx.recv(), if normal_open => {
match frame {
Some(frame) => {
if !write_frame(&mut sink, frame, tunnel_metrics.as_deref()).await {
break;
}
}
None => normal_open = false,
}
}
}
}
debug!("writer task exiting");
let _ = sink.close().await;
});
(tx, handle)
}
fn classify_frame_priority(frame: &Frame) -> FramePriority {
match frame.msg_type {
MsgType::ResponseHeaders
| MsgType::StreamError
| MsgType::Ping
| MsgType::Pong
| MsgType::GoAway
| MsgType::HeartbeatData
| MsgType::HeartbeatAck => FramePriority::High,
MsgType::RequestHeaders
| MsgType::RequestBody
| MsgType::ResponseBody
| MsgType::StreamEnd => FramePriority::Normal,
}
}
async fn write_frame<S>(sink: &mut S, frame: Frame, tunnel_metrics: Option<&TunnelMetrics>) -> bool
where
S: SinkExt<Message, Error = tokio_tungstenite::tungstenite::Error> + Unpin + Send + 'static,
{
let stream_id = frame.stream_id;
let msg_type = frame.msg_type;
let flags = frame.flags;
let data = frame.encode();
let wire_len = data.len().max(HEADER_SIZE);
if let Err(e) = sink.send(Message::Binary(data.into())).await {
error!(
stream_id = stream_id,
msg_type = ?msg_type,
flags = flags,
wire_len = wire_len,
error = %e,
"failed to write frame to WebSocket"
);
if let Some(metrics) = tunnel_metrics {
metrics.record_error("ws_write_error", &e.to_string());
}
return false;
}
if let Some(metrics) = tunnel_metrics {
metrics.record_ws_outgoing_frame(wire_len);
}
true
}
#[cfg(test)]
mod tests {
use std::pin::Pin;
use std::sync::{Arc, Mutex};
use std::task::{Context, Poll};
use std::time::Duration;
use futures_util::Sink;
use tokio_tungstenite::tungstenite::{Error, Message};
use super::spawn_writer;
use crate::tunnel::protocol::Frame;
use aether_contracts::tunnel::MsgType;
#[derive(Clone, Default)]
struct VecSink {
sent: Arc<Mutex<Vec<Message>>>,
}
impl Sink<Message> for VecSink {
type Error = Error;
fn poll_ready(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn start_send(self: Pin<&mut Self>, item: Message) -> Result<(), Self::Error> {
self.sent.lock().expect("sink lock").push(item);
Ok(())
}
fn poll_flush(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn poll_close(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
}
#[tokio::test]
async fn prioritizes_control_frames_ahead_of_buffered_body_frames() {
let sink = VecSink::default();
let sent = Arc::clone(&sink.sent);
let (sender, handle) = spawn_writer(sink, Duration::from_secs(60));
for idx in 0..8u8 {
sender
.try_send(Frame::new(
7,
MsgType::ResponseBody,
0,
bytes::Bytes::from(vec![idx; 32]),
))
.expect("frame send should succeed");
}
sender
.try_send(Frame::new(
7,
MsgType::StreamError,
0,
bytes::Bytes::from_static(b"boom"),
))
.expect("frame send should succeed");
let snapshots = sender.snapshots();
assert!(snapshots.high.enqueued_total >= 1);
assert!(snapshots.normal.enqueued_total >= 8);
tokio::time::sleep(Duration::from_millis(30)).await;
drop(sender);
handle.await.expect("writer should exit cleanly");
let sent = sent.lock().expect("sink lock");
assert!(
sent.len() >= 2,
"writer should flush both body and control frames"
);
let first = match &sent[0] {
Message::Binary(data) => {
Frame::decode(data.clone().into()).expect("frame should decode")
}
other => panic!("unexpected first message: {other:?}"),
};
let second = match &sent[1] {
Message::Binary(data) => {
Frame::decode(data.clone().into()).expect("frame should decode")
}
other => panic!("unexpected second message: {other:?}"),
};
assert_eq!(first.msg_type, MsgType::StreamError);
assert_eq!(second.msg_type, MsgType::ResponseBody);
}
}

View File

@@ -0,0 +1,969 @@
use std::collections::HashMap;
use std::convert::Infallible;
use std::future::Future;
use std::io;
use std::net::IpAddr;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::Mutex;
use std::task::{Context, Poll};
use std::time::Duration;
use aether_contracts::{
ResolvedTransportProfile, TRANSPORT_BACKEND_HYPER_RUSTLS, TRANSPORT_BACKEND_REQWEST_RUSTLS,
TRANSPORT_HTTP_MODE_HTTP1_ONLY,
};
use bytes::Bytes;
use futures_util::Stream;
use http_body_util::combinators::UnsyncBoxBody;
use http_body_util::{BodyExt, Full, StreamBody};
use hyper::body::Frame;
use hyper::rt;
use hyper::Response;
use hyper::Uri;
pub use hyper_util::client::legacy::connect::capture_connection;
use hyper_util::client::legacy::connect::dns::Name;
use hyper_util::client::legacy::connect::{Connected, Connection, HttpConnector};
use hyper_util::client::legacy::Client;
use hyper_util::rt::{TokioExecutor, TokioIo, TokioTimer};
use rustls::pki_types::ServerName;
use rustls::ClientConfig;
use tokio::net::TcpStream;
use tokio_rustls::TlsConnector;
use tower_service::Service;
use crate::config::Config;
use crate::egress_proxy::{
connect_proxy_tcp, http_connect, socks5_connect, ProxyConnectOptions, UpstreamProxyConfig,
UpstreamProxyScheme,
};
use crate::target_filter::{self, DnsCache};
type BoxError = Box<dyn std::error::Error + Send + Sync>;
type PlainStream = TokioIo<TcpStream>;
type TlsStream = TokioIo<tokio_rustls::client::TlsStream<TcpStream>>;
pub type UpstreamRequestBody = UnsyncBoxBody<Bytes, io::Error>;
pub type UpstreamClient = Client<InstrumentedConnector, UpstreamRequestBody>;
const DEFAULT_PROFILE_ID: &str = "default";
const DEFAULT_BACKEND: &str = TRANSPORT_BACKEND_HYPER_RUSTLS;
const DEFAULT_HTTP_MODE: &str = "auto";
#[derive(Clone, Debug, Eq, Hash, PartialEq)]
pub struct UpstreamClientPoolKey {
pub provider_id: String,
pub endpoint_id: String,
pub key_id: String,
pub profile_id: String,
pub backend: String,
pub http_mode: String,
}
#[derive(Clone)]
pub struct UpstreamClientPool {
config: Arc<Config>,
dns_cache: Arc<DnsCache>,
clients: Arc<Mutex<HashMap<UpstreamClientPoolKey, UpstreamClient>>>,
}
impl UpstreamClientPool {
pub fn new(config: Arc<Config>, dns_cache: Arc<DnsCache>) -> Self {
Self {
config,
dns_cache,
clients: Arc::new(Mutex::new(HashMap::new())),
}
}
pub fn get_or_build(&self, key: UpstreamClientPoolKey) -> Result<UpstreamClient, String> {
if let Some(client) = self
.clients
.lock()
.expect("client pool lock")
.get(&key)
.cloned()
{
return Ok(client);
}
validate_proxy_transport_backend(&key.backend)?;
let http1_only = key
.http_mode
.eq_ignore_ascii_case(TRANSPORT_HTTP_MODE_HTTP1_ONLY);
let client = build_upstream_client_with_protocol(
&self.config,
Arc::clone(&self.dns_cache),
http1_only,
)?;
self.clients
.lock()
.expect("client pool lock")
.insert(key, client.clone());
Ok(client)
}
}
pub fn upstream_client_pool_key(
provider_id: Option<&str>,
endpoint_id: Option<&str>,
key_id: Option<&str>,
profile: Option<&ResolvedTransportProfile>,
http1_only: bool,
) -> UpstreamClientPoolKey {
let profile_http_mode = profile
.map(|profile| profile.http_mode.trim())
.filter(|value| !value.is_empty())
.unwrap_or(DEFAULT_HTTP_MODE);
let http_mode = if http1_only {
TRANSPORT_HTTP_MODE_HTTP1_ONLY
} else {
profile_http_mode
};
UpstreamClientPoolKey {
provider_id: normalized_pool_key_part(provider_id),
endpoint_id: normalized_pool_key_part(endpoint_id),
key_id: normalized_pool_key_part(key_id),
profile_id: profile
.map(|profile| profile.profile_id.trim())
.filter(|value| !value.is_empty())
.unwrap_or(DEFAULT_PROFILE_ID)
.to_string(),
backend: profile
.map(|profile| profile.backend.trim())
.filter(|value| !value.is_empty())
.unwrap_or(DEFAULT_BACKEND)
.to_string(),
http_mode: http_mode.to_string(),
}
}
fn normalized_pool_key_part(value: Option<&str>) -> String {
value
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or("-")
.to_string()
}
fn validate_proxy_transport_backend(backend: &str) -> Result<(), String> {
if backend.eq_ignore_ascii_case(TRANSPORT_BACKEND_HYPER_RUSTLS)
|| backend.eq_ignore_ascii_case(TRANSPORT_BACKEND_REQWEST_RUSTLS)
{
return Ok(());
}
Err(format!("unsupported transport profile backend: {backend}"))
}
pub fn http_proxy_authorization_header(proxy_url: Option<&str>) -> Option<String> {
let proxy = proxy_url
.map(str::trim)
.filter(|value| !value.is_empty())
.and_then(|value| UpstreamProxyConfig::parse(value).ok())?;
if proxy.scheme() == UpstreamProxyScheme::Http {
proxy.basic_auth_header()
} else {
None
}
}
pub fn stream_request_body<S>(stream: S) -> UpstreamRequestBody
where
S: Stream<Item = Result<Frame<Bytes>, io::Error>> + Send + 'static,
{
StreamBody::new(stream).boxed_unsync()
}
pub fn full_request_body(body: Bytes) -> UpstreamRequestBody {
Full::new(body)
.map_err(|err: Infallible| match err {})
.boxed_unsync()
}
#[derive(Clone, Copy, Debug, Default)]
pub struct ConnectTiming {
pub connect_ms: u64,
pub tls_ms: u64,
}
#[derive(Clone, Copy, Debug, Default)]
pub struct RequestTiming {
pub connection_acquire_ms: u64,
pub connect_ms: u64,
pub tls_ms: u64,
pub response_wait_ms: u64,
pub connection_reused: bool,
}
#[derive(Clone)]
pub struct ValidatedResolver {
dns_cache: Arc<DnsCache>,
allow_private: bool,
}
impl ValidatedResolver {
pub fn new(dns_cache: Arc<DnsCache>, allow_private: bool) -> Self {
Self {
dns_cache,
allow_private,
}
}
}
pub struct ValidatedAddrs {
inner: std::vec::IntoIter<std::net::SocketAddr>,
}
impl Iterator for ValidatedAddrs {
type Item = std::net::SocketAddr;
fn next(&mut self) -> Option<Self::Item> {
self.inner.next()
}
}
impl Service<Name> for ValidatedResolver {
type Response = ValidatedAddrs;
type Error = io::Error;
type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn call(&mut self, name: Name) -> Self::Future {
let dns_cache = Arc::clone(&self.dns_cache);
let allow_private = self.allow_private;
let host = name.as_str().to_string();
Box::pin(async move {
if let Some(addrs) = dns_cache.get_by_host(&host).await {
return Ok(ValidatedAddrs {
inner: (*addrs).clone().into_iter(),
});
}
let resolved =
target_filter::resolve_public_addrs(&host, 0, allow_private, dns_cache.as_ref())
.await
.map_err(|err| io::Error::other(err.to_string()))?;
Ok(ValidatedAddrs {
inner: resolved.into_iter(),
})
})
}
}
#[derive(Clone)]
pub struct InstrumentedConnector {
http: HttpConnector<ValidatedResolver>,
tls_config: Arc<ClientConfig>,
proxy: Option<UpstreamProxyConfig>,
connect_timeout: Duration,
tcp_nodelay: bool,
tcp_keepalive: Option<Duration>,
}
impl Service<Uri> for InstrumentedConnector {
type Response = TimedConn;
type Error = BoxError;
type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.http.poll_ready(cx).map_err(Into::into)
}
fn call(&mut self, dst: Uri) -> Self::Future {
let scheme = dst.scheme_str().map(|value| value.to_ascii_lowercase());
let tls_config = Arc::clone(&self.tls_config);
if let Some(proxy) = self.proxy.clone() {
let options = ProxyConnectOptions {
connect_timeout: self.connect_timeout,
tcp_nodelay: self.tcp_nodelay,
tcp_keepalive: self.tcp_keepalive,
};
let connect_start = std::time::Instant::now();
return Box::pin(async move {
connect_via_proxy(dst, scheme, tls_config, proxy, options, connect_start).await
});
}
let connecting = self.http.call(dst.clone());
let connect_start = std::time::Instant::now();
Box::pin(async move {
match scheme.as_deref() {
Some("http") => {
let tcp = connecting.await.map_err(|err| Box::new(err) as BoxError)?;
let connect_ms = connect_start.elapsed().as_millis() as u64;
Ok(TimedConn::new(
MaybeHttpsStream::Http {
stream: tcp,
is_proxy: false,
},
ConnectTiming {
connect_ms,
tls_ms: 0,
},
))
}
Some("https") => {
let server_name = resolve_server_name(&dst)?;
let tcp = connecting.await.map_err(|err| Box::new(err) as BoxError)?;
let connect_ms = connect_start.elapsed().as_millis() as u64;
let tls_start = std::time::Instant::now();
let tls_stream = TlsConnector::from(tls_config)
.connect(server_name, tcp.into_inner())
.await
.map_err(io::Error::other)?;
let tls_ms = tls_start.elapsed().as_millis() as u64;
Ok(TimedConn::new(
MaybeHttpsStream::Https(TokioIo::new(tls_stream)),
ConnectTiming { connect_ms, tls_ms },
))
}
Some(other) => Err(io::Error::other(format!("unsupported scheme {other}")).into()),
None => Err(io::Error::other("missing scheme").into()),
}
})
}
}
async fn connect_via_proxy(
dst: Uri,
scheme: Option<String>,
tls_config: Arc<ClientConfig>,
proxy: UpstreamProxyConfig,
options: ProxyConnectOptions,
connect_start: std::time::Instant,
) -> Result<TimedConn, BoxError> {
let scheme = scheme.ok_or_else(|| io::Error::other("missing scheme"))?;
let target_host = uri_host(&dst)?;
let target_port = uri_port_or_default(&dst, &scheme)?;
let mut tcp = connect_proxy_tcp(
&proxy,
options.connect_timeout,
options.tcp_nodelay,
options.tcp_keepalive,
)
.await?;
match proxy.scheme() {
UpstreamProxyScheme::Http => {
if scheme == "https" {
http_connect(
&mut tcp,
&target_authority(&target_host, target_port),
&proxy,
)
.await?;
} else if scheme != "http" {
return Err(io::Error::other(format!("unsupported scheme {scheme}")).into());
}
}
UpstreamProxyScheme::Socks5 | UpstreamProxyScheme::Socks5h => {
socks5_connect(&mut tcp, &proxy, &target_host, target_port).await?;
}
}
let connect_ms = connect_start.elapsed().as_millis() as u64;
match scheme.as_str() {
"http" => Ok(TimedConn::new(
MaybeHttpsStream::Http {
stream: TokioIo::new(tcp),
is_proxy: proxy.scheme() == UpstreamProxyScheme::Http,
},
ConnectTiming {
connect_ms,
tls_ms: 0,
},
)),
"https" => {
let tls_start = std::time::Instant::now();
let tls_stream = TlsConnector::from(tls_config)
.connect(resolve_server_name(&dst)?, tcp)
.await
.map_err(io::Error::other)?;
let tls_ms = tls_start.elapsed().as_millis() as u64;
Ok(TimedConn::new(
MaybeHttpsStream::Https(TokioIo::new(tls_stream)),
ConnectTiming { connect_ms, tls_ms },
))
}
other => Err(io::Error::other(format!("unsupported scheme {other}")).into()),
}
}
fn uri_host(uri: &Uri) -> Result<String, io::Error> {
uri.host()
.map(|host| {
host.trim_start_matches('[')
.trim_end_matches(']')
.to_string()
})
.filter(|host| !host.is_empty())
.ok_or_else(|| io::Error::other("missing host"))
}
fn uri_port_or_default(uri: &Uri, scheme: &str) -> Result<u16, io::Error> {
uri.port_u16()
.or(match scheme {
"http" => Some(80),
"https" => Some(443),
_ => None,
})
.ok_or_else(|| io::Error::other(format!("missing port for scheme {scheme}")))
}
fn target_authority(host: &str, port: u16) -> String {
if host.contains(':') && !host.starts_with('[') {
format!("[{host}]:{port}")
} else {
format!("{host}:{port}")
}
}
fn build_upstream_client_with_protocol(
config: &Config,
dns_cache: Arc<DnsCache>,
http1_only: bool,
) -> Result<UpstreamClient, String> {
let mut http = HttpConnector::new_with_resolver(ValidatedResolver::new(
dns_cache,
config.allow_private_targets,
));
http.enforce_http(false);
http.set_connect_timeout(Some(Duration::from_secs(
config.upstream_connect_timeout_secs,
)));
http.set_nodelay(config.upstream_tcp_nodelay);
if config.upstream_tcp_keepalive_secs > 0 {
http.set_keepalive(Some(Duration::from_secs(
config.upstream_tcp_keepalive_secs,
)));
} else {
http.set_keepalive(None);
}
let connector = InstrumentedConnector {
http,
tls_config: build_tls_config(http1_only),
proxy: config
.upstream_proxy_url
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(UpstreamProxyConfig::parse)
.transpose()?,
connect_timeout: Duration::from_secs(config.upstream_connect_timeout_secs),
tcp_nodelay: config.upstream_tcp_nodelay,
tcp_keepalive: (config.upstream_tcp_keepalive_secs > 0)
.then(|| Duration::from_secs(config.upstream_tcp_keepalive_secs)),
};
let mut builder = Client::builder(TokioExecutor::new());
builder.pool_max_idle_per_host(config.upstream_pool_max_idle_per_host);
builder.pool_idle_timeout(Duration::from_secs(config.upstream_pool_idle_timeout_secs));
builder.pool_timer(TokioTimer::new());
Ok(builder.build(connector))
}
pub fn resolve_request_timing<B>(
response: &Response<B>,
connection_acquire_ms: Option<u64>,
ttfb_ms: u64,
) -> RequestTiming {
let raw = response
.extensions()
.get::<ConnectTiming>()
.copied()
.unwrap_or_default();
let raw_connection_ms = raw.connect_ms.saturating_add(raw.tls_ms);
let measured_acquire_ms = connection_acquire_ms.unwrap_or(raw_connection_ms.min(ttfb_ms));
let likely_reused = measured_acquire_ms <= 5 && raw_connection_ms > 0;
let connector_matches_request = raw_connection_ms <= measured_acquire_ms.saturating_add(25);
let (connect_ms, tls_ms) = if likely_reused || !connector_matches_request {
(0, 0)
} else {
(raw.connect_ms, raw.tls_ms)
};
RequestTiming {
connection_acquire_ms: measured_acquire_ms,
connect_ms,
tls_ms,
response_wait_ms: ttfb_ms.saturating_sub(measured_acquire_ms),
connection_reused: likely_reused,
}
}
fn build_tls_config(http1_only: bool) -> Arc<ClientConfig> {
let _ = rustls::crypto::ring::default_provider().install_default();
let root_store =
rustls::RootCertStore::from_iter(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
let mut config = ClientConfig::builder()
.with_root_certificates(root_store)
.with_no_client_auth();
config.alpn_protocols = if http1_only {
vec![b"http/1.1".to_vec()]
} else {
vec![b"h2".to_vec(), b"http/1.1".to_vec()]
};
Arc::new(config)
}
fn resolve_server_name(uri: &Uri) -> Result<ServerName<'static>, BoxError> {
let host = uri.host().ok_or_else(|| io::Error::other("missing host"))?;
let host = host.trim_start_matches('[').trim_end_matches(']');
if let Ok(ip) = host.parse::<IpAddr>() {
return Ok(ServerName::from(ip));
}
Ok(ServerName::try_from(host.to_string())?)
}
pub struct TimedConn {
inner: MaybeHttpsStream,
timing: ConnectTiming,
}
impl TimedConn {
fn new(inner: MaybeHttpsStream, timing: ConnectTiming) -> Self {
Self { inner, timing }
}
}
impl Connection for TimedConn {
fn connected(&self) -> Connected {
self.inner.connected().extra(self.timing)
}
}
impl rt::Read for TimedConn {
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: rt::ReadBufCursor<'_>,
) -> Poll<Result<(), io::Error>> {
Pin::new(&mut self.inner).poll_read(cx, buf)
}
}
impl rt::Write for TimedConn {
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<Result<usize, io::Error>> {
Pin::new(&mut self.inner).poll_write(cx, buf)
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
Pin::new(&mut self.inner).poll_flush(cx)
}
fn poll_shutdown(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Result<(), io::Error>> {
Pin::new(&mut self.inner).poll_shutdown(cx)
}
fn is_write_vectored(&self) -> bool {
self.inner.is_write_vectored()
}
fn poll_write_vectored(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
bufs: &[std::io::IoSlice<'_>],
) -> Poll<Result<usize, io::Error>> {
Pin::new(&mut self.inner).poll_write_vectored(cx, bufs)
}
}
pub enum MaybeHttpsStream {
Http { stream: PlainStream, is_proxy: bool },
Https(TlsStream),
}
impl Connection for MaybeHttpsStream {
fn connected(&self) -> Connected {
match self {
Self::Http { stream, is_proxy } => stream.connected().proxy(*is_proxy),
Self::Https(stream) => {
let (tcp, tls) = stream.inner().get_ref();
if tls.alpn_protocol() == Some(b"h2") {
tcp.connected().negotiated_h2()
} else {
tcp.connected()
}
}
}
}
}
impl rt::Read for MaybeHttpsStream {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: rt::ReadBufCursor<'_>,
) -> Poll<Result<(), io::Error>> {
match Pin::get_mut(self) {
Self::Http { stream, .. } => Pin::new(stream).poll_read(cx, buf),
Self::Https(stream) => Pin::new(stream).poll_read(cx, buf),
}
}
}
impl rt::Write for MaybeHttpsStream {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<Result<usize, io::Error>> {
match Pin::get_mut(self) {
Self::Http { stream, .. } => Pin::new(stream).poll_write(cx, buf),
Self::Https(stream) => Pin::new(stream).poll_write(cx, buf),
}
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
match Pin::get_mut(self) {
Self::Http { stream, .. } => Pin::new(stream).poll_flush(cx),
Self::Https(stream) => Pin::new(stream).poll_flush(cx),
}
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
match Pin::get_mut(self) {
Self::Http { stream, .. } => Pin::new(stream).poll_shutdown(cx),
Self::Https(stream) => Pin::new(stream).poll_shutdown(cx),
}
}
fn is_write_vectored(&self) -> bool {
match self {
Self::Http { stream, .. } => stream.is_write_vectored(),
Self::Https(stream) => stream.is_write_vectored(),
}
}
fn poll_write_vectored(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
bufs: &[std::io::IoSlice<'_>],
) -> Poll<Result<usize, io::Error>> {
match Pin::get_mut(self) {
Self::Http { stream, .. } => Pin::new(stream).poll_write_vectored(cx, bufs),
Self::Https(stream) => Pin::new(stream).poll_write_vectored(cx, bufs),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use aether_contracts::ResolvedTransportProfile;
use clap::Parser;
use http_body_util::BodyExt;
use hyper::Response;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
use crate::egress_proxy::socks5_target_address;
#[test]
fn fresh_connection_uses_connector_breakdown() {
let mut response = Response::new(());
response.extensions_mut().insert(ConnectTiming {
connect_ms: 80,
tls_ms: 40,
});
let timing = resolve_request_timing(&response, Some(125), 600);
assert_eq!(timing.connection_acquire_ms, 125);
assert_eq!(timing.connect_ms, 80);
assert_eq!(timing.tls_ms, 40);
assert_eq!(timing.response_wait_ms, 475);
assert!(!timing.connection_reused);
}
#[test]
fn reused_connection_zeroes_stale_connect_timings() {
let mut response = Response::new(());
response.extensions_mut().insert(ConnectTiming {
connect_ms: 70,
tls_ms: 30,
});
let timing = resolve_request_timing(&response, Some(0), 310);
assert_eq!(timing.connection_acquire_ms, 0);
assert_eq!(timing.connect_ms, 0);
assert_eq!(timing.tls_ms, 0);
assert_eq!(timing.response_wait_ms, 310);
assert!(timing.connection_reused);
}
#[test]
fn falls_back_to_connector_timings_when_capture_missing() {
let mut response = Response::new(());
response.extensions_mut().insert(ConnectTiming {
connect_ms: 55,
tls_ms: 25,
});
let timing = resolve_request_timing(&response, None, 400);
assert_eq!(timing.connection_acquire_ms, 80);
assert_eq!(timing.connect_ms, 55);
assert_eq!(timing.tls_ms, 25);
assert_eq!(timing.response_wait_ms, 320);
assert!(!timing.connection_reused);
}
#[test]
fn upstream_client_pool_key_includes_profile_identity() {
let profile = ResolvedTransportProfile {
profile_id: "profile-a".to_string(),
backend: TRANSPORT_BACKEND_REQWEST_RUSTLS.to_string(),
http_mode: "auto".to_string(),
pool_scope: "key".to_string(),
header_fingerprint: None,
extra: None,
};
let pool_key = upstream_client_pool_key(
Some("provider-1"),
Some("endpoint-1"),
Some("key-1"),
Some(&profile),
false,
);
assert_eq!(pool_key.provider_id, "provider-1");
assert_eq!(pool_key.endpoint_id, "endpoint-1");
assert_eq!(pool_key.key_id, "key-1");
assert_eq!(pool_key.profile_id, "profile-a");
assert_eq!(pool_key.backend, TRANSPORT_BACKEND_REQWEST_RUSTLS);
assert_eq!(pool_key.http_mode, "auto");
}
#[test]
fn upstream_client_pool_rejects_unsupported_backend() {
let error = validate_proxy_transport_backend("utls").unwrap_err();
assert!(error.contains("unsupported transport profile backend"));
}
#[test]
fn http_proxy_authorization_header_uses_basic_auth_for_http_proxy() {
assert_eq!(
http_proxy_authorization_header(Some("http://user:pass@proxy.example:8080")).as_deref(),
Some("Basic dXNlcjpwYXNz")
);
assert_eq!(
http_proxy_authorization_header(Some("socks5h://user:pass@127.0.0.1:1080")),
None
);
}
#[tokio::test]
async fn socks5h_target_address_uses_domain_name() {
let request = socks5_target_address("example.com", 443, true)
.await
.expect("SOCKS target should build");
assert_eq!(
request,
[
&[0x05, 0x01, 0x00, 0x03, 11][..],
b"example.com",
&[0x01, 0xbb][..],
]
.concat()
);
}
#[tokio::test]
#[ignore = "requires loopback listener support"]
async fn upstream_client_sends_http_requests_through_http_proxy() {
let (proxy_url, request_rx) = spawn_http_proxy().await;
let client = proxied_client(&proxy_url);
let request = hyper::Request::builder()
.method(hyper::Method::GET)
.uri("http://example.com/tunnel-test")
.body(full_request_body(Bytes::new()))
.expect("request should build");
let response = client.request(request).await.expect("request should pass");
let status = response.status();
let body = response
.into_body()
.collect()
.await
.expect("body should collect")
.to_bytes();
let raw_request = request_rx.await.expect("proxy should receive request");
assert_eq!(status, hyper::StatusCode::OK);
assert_eq!(&body[..], b"ok");
assert!(
raw_request.starts_with("GET http://example.com/tunnel-test HTTP/1.1\r\n"),
"unexpected proxy request: {raw_request:?}"
);
}
#[tokio::test]
#[ignore = "requires loopback listener support"]
async fn upstream_client_sends_http_requests_through_socks5h_proxy() {
let (proxy_url, target_rx, request_rx) = spawn_socks5h_proxy().await;
let client = proxied_client(&proxy_url);
let request = hyper::Request::builder()
.method(hyper::Method::GET)
.uri("http://example.com/socks-test")
.body(full_request_body(Bytes::new()))
.expect("request should build");
let response = client.request(request).await.expect("request should pass");
let body = response
.into_body()
.collect()
.await
.expect("body should collect")
.to_bytes();
let target = target_rx.await.expect("SOCKS proxy should receive target");
let raw_request = request_rx
.await
.expect("SOCKS proxy should receive HTTP request");
assert_eq!(&body[..], b"ok");
assert_eq!(target, ("example.com".to_string(), 80));
assert!(
raw_request.starts_with("GET /socks-test HTTP/1.1\r\n"),
"unexpected SOCKS tunneled request: {raw_request:?}"
);
}
fn proxied_client(proxy_url: &str) -> UpstreamClient {
let _ = rustls::crypto::ring::default_provider().install_default();
let config = Config::try_parse_from([
"aether-tunnel",
"--aether-url",
"https://aether.example.com",
"--management-token",
"ae_test",
"--node-name",
"tunnel-test",
"--upstream-proxy-url",
proxy_url,
"--upstream-connect-timeout-secs",
"2",
])
.expect("config should parse");
build_upstream_client_with_protocol(
&config,
Arc::new(DnsCache::new(Duration::from_secs(60), 16)),
true,
)
.expect("client should build")
}
async fn spawn_http_proxy() -> (String, tokio::sync::oneshot::Receiver<String>) {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("listener should bind");
let addr = listener.local_addr().expect("local addr should exist");
let (request_tx, request_rx) = tokio::sync::oneshot::channel();
tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.expect("proxy should accept");
let request = read_http_headers(&mut stream).await;
let _ = request_tx.send(request);
stream
.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\nok")
.await
.expect("proxy response should write");
});
(format!("http://{addr}"), request_rx)
}
async fn spawn_socks5h_proxy() -> (
String,
tokio::sync::oneshot::Receiver<(String, u16)>,
tokio::sync::oneshot::Receiver<String>,
) {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("listener should bind");
let addr = listener.local_addr().expect("local addr should exist");
let (target_tx, target_rx) = tokio::sync::oneshot::channel();
let (request_tx, request_rx) = tokio::sync::oneshot::channel();
tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.expect("SOCKS proxy should accept");
let mut greeting = [0u8; 3];
stream
.read_exact(&mut greeting)
.await
.expect("SOCKS greeting should read");
assert_eq!(greeting, [0x05, 0x01, 0x00]);
stream
.write_all(&[0x05, 0x00])
.await
.expect("SOCKS method should write");
let mut request_head = [0u8; 5];
stream
.read_exact(&mut request_head)
.await
.expect("SOCKS request head should read");
assert_eq!(&request_head[..4], &[0x05, 0x01, 0x00, 0x03]);
let len = request_head[4] as usize;
let mut host = vec![0u8; len];
stream
.read_exact(&mut host)
.await
.expect("SOCKS host should read");
let mut port = [0u8; 2];
stream
.read_exact(&mut port)
.await
.expect("SOCKS port should read");
let host = String::from_utf8(host).expect("SOCKS host should be UTF-8");
let port = u16::from_be_bytes(port);
let _ = target_tx.send((host, port));
stream
.write_all(&[0x05, 0x00, 0x00, 0x01, 0, 0, 0, 0, 0, 0])
.await
.expect("SOCKS connect response should write");
let request = read_http_headers(&mut stream).await;
let _ = request_tx.send(request);
stream
.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\nok")
.await
.expect("SOCKS tunneled response should write");
});
(format!("socks5h://{addr}"), target_rx, request_rx)
}
async fn read_http_headers(stream: &mut TcpStream) -> String {
let mut buf = Vec::new();
let mut chunk = [0u8; 1024];
loop {
let n = stream.read(&mut chunk).await.expect("request should read");
assert!(n > 0, "connection closed before headers finished");
buf.extend_from_slice(&chunk[..n]);
if buf.windows(4).any(|window| window == b"\r\n\r\n") {
break;
}
}
String::from_utf8(buf).expect("headers should be UTF-8")
}
}