Compare commits

..
11 Commits
Author SHA1 Message Date
fawney19 4e96b9d870 Merge pull request #418 from Avilianb/feat/codex-responses-websocket
feat: support Codex Responses WebSocket transport
2026-05-15 12:49:45 +08:00
Avilianb d3c0b1aa7f feat: support Codex Responses WebSocket transport 2026-05-10 01:59:48 +08:00
phenetr0andphenetr0 66bfd3e592 fix: avoid double conversion for forced cli sync streams (#367)
Co-authored-by: phenetr0 <[email protected]>
2026-05-03 02:11:16 +08:00
DaoandYour Name 392c557831 优化 OAuth Token 导入解析与号池管理直刷空列表问题 (#330)
* docs: add aether-proxy C++ rewrite design

* chore: ignore local worktrees directory

* 修复 OAuth Token 导入识别与账号信息解析

* 修复号池管理直刷账号列表为空

---------

Co-authored-by: Your Name <[email protected]>
2026-04-30 18:07:00 +08:00
49 c34565b02b fix(provider-test): 修复测试供应商配置时 DetachedInstanceError (#270)
两处修复:

1. endpoint_checker.py: 移除 asyncio.to_thread 包装同步 DB 查询,
   避免跨线程导致 SQLAlchemy Session 失效

2. provider_query.py: 并发测试预加载 ProviderEndpoint 和 ProviderAPIKey 时
   添加 joinedload(provider),确保 expunge 后不会触发懒加载
2026-04-17 10:37:19 +08:00
fawney19 f57fe6e13e fix(migration): 用 UPDATE...FROM 子查询修复 total_tokens 自引用更新问题,增加批次上限防止死循环 2026-03-24 02:07:24 +08:00
fawney19 7c678b715f fix(migration): 修复 usage token 语义迁移脚本
- 将内联 SQL 提取为模块级常量 _UPGRADE_BACKFILL_SQL / _DOWNGRADE_BACKFILL_SQL
- 抽取 run_backfill_in_batches() 统一批量更新逻辑,通过 autocommit_block 释放 ALTER TABLE 锁
- 修正 upgrade 中 total_tokens 计算逻辑:用 COALESCE(input_tokens,0)+COALESCE(output_tokens,0) 补全 input_output_total_tokens
- WHERE 条件改为按实际字段差异筛选待回填行,替换原先仅过滤 input_context_tokens=0 AND total_tokens=0 的不完整条件
- 批量大小由 5000 调整为 500,减少单批锁持有时间
2026-03-24 01:56:36 +08:00
fawney19andNyaDoo bd4e5f3a5d feat(analytics): 重构统计分析模块,统一 API 与前端视图
close #260

- 新增 `src/services/analytics/query_service.py`,集中实现排行榜、性能、时间序列等查询逻辑
- 新增 `src/api/analytics/routes.py`,替代原 `stats/` 和 `dashboard/` 的分散路由
- 删除旧 `src/api/admin/stats/`、`src/api/dashboard/` 模块
- 重构 `src/api/user_me/routes.py` 与 `src/api/admin/usage/routes.py`,精简用量查询接口
- 新增 Alembic 迁移,修正 token 语义字段
- 前端新增 Analytics.vue、LeaderboardTab、PerformanceTab、ReportsTab 及 Reports 用户视图
- 新增 composables(useAnalyticsFilters、useReportsData、useLeaderboardData、usePerformanceData)
- 新增工具函数:analyticsGranularity、analyticsTimeseries、chartTheme、csvExport、usageBreakdown
- 删除旧 CostAnalysis、PerformanceAnalysis、UserStats 页面及相关组件
- 前端 API 层重组:新增 analytics.ts、request-details.ts,删除 dashboard.ts 和 usage.ts

Co-authored-by: NyaDoo <[email protected]>
2026-03-24 01:51:53 +08:00
York Zang 165d9eab8f fix(proxy): 修正代理连通性测试地址,避免 1.1.1.1 证书校验失败 (#257)
代理连通性测试此前使用 https://1.1.1.1/cdn-cgi/trace 作为探测地址,在标准 TLS 校验下会因为证书与 IP 不匹配而失败。
改为使用基于域名的 Cloudflare trace 地址,避免触发 CERTIFICATE_VERIFY_FAILED,并恢复代理测试结果的准确性。
2026-03-24 00:04:01 +08:00
fawney19andAAEE86 dfb95f09e1 feat(health): 增强健康监控面板,支持按 Key 查看详情与摘要统计
- 后端 health monitor 新增 summary 接口,按 api_format 聚合健康摘要
- 新增 GroupedFormatKey 类型与 key 分组查询接口
- HealthMonitorCard 重构为卡片+详情对话框,展示成功率、响应时间、Key 状态
- 抽取 HealthMonitorDetailDialog 独立组件
- 新增 useRouteQuery composable 用于 URL query 参数双向绑定
- PoolManagement/ProviderManagement 集成健康监控入口
- 新增 health monitor summary 单元测试

Closes #256

Co-authored-by: AAEE86 <[email protected]>
2026-03-24 00:00:08 +08:00
fawney19andhemo94931 4d5c591654 fix: 覆写规则条件可使用映射前请求体
新增 rules_original_body 参数贯穿请求构建链路,确保 body_rules/header_rules
条件评估使用模型映射前的原始请求体;附带将 handlers __init__ 改为延迟导入。

Closes #255

Co-authored-by: hemo94931 <[email protected]>
2026-03-23 17:31:29 +08:00
3739 changed files with 292043 additions and 1075766 deletions
+15 -5
View File
@@ -1,14 +1,24 @@
# Build artifacts
build/
target/
# Python
__pycache__/
*.py[cod]
*$py.class
*.so
*.egg
.Python
env/
venv/
ENV/
.venv
.uv/
*.egg-info/
dist/
!aether-hub/dist/aether-hub
build/
*.egg
# Frontend
frontend/node_modules/
frontend/dist/
frontend/.vite/
# frontend/dist/ - 注释掉,因为我们需要预构建的dist文件
# Development
.git/
+84 -62
View File
@@ -1,48 +1,19 @@
# ==================== 必须配置(启动前) ====================
# 以下配置项必须在项目启动前设置
# 应用端口(默认 8084)
APP_PORT=8084
# 对外访问地址,用于一键安装、CC Switch 导入、支付回调等需要生成公网 URL 的场景。
# 生产环境建议显式配置为不带内部端口的公网域名,例如 https://aether.example.com
# AETHER_PUBLIC_BASE_URL=https://aether.example.com
# Docker Compose 镜像(默认正式版 latest;提前测试可改 rc/beta;也可固定具体版本)
# 示例:
# APP_IMAGE=ghcr.io/fawney19/aether:latest
# APP_IMAGE=ghcr.io/fawney19/aether:rc
# APP_IMAGE=ghcr.io/fawney19/aether:beta
# APP_IMAGE=ghcr.io/fawney19/aether:0.7.0-rc.1
# API Key 前缀(默认 sk)
API_KEY_PREFIX=sk
# Rust 日志过滤(默认 aether_gateway=info)
# 示例: aether_gateway=debug,sqlx=warn
RUST_LOG=aether_gateway=info
# CORS 配置(跨域带 Cookie 时不要写 *,必须显式列出前端源)
# 示例: http://localhost:5173,https://app.example.com
# CORS_ORIGINS=http://localhost:5173
# CORS_ALLOW_CREDENTIALS=true
# 如果前后端跨站并依赖登录刷新 Cookie,还要配合:
# AUTH_REFRESH_COOKIE_SAMESITE=None
# AUTH_REFRESH_COOKIE_SECURE=true
# 数据库配置
DB_HOST=localhost
DB_PORT=5432
DB_USER=postgres
DB_NAME=aether
DB_PASSWORD=aether
DB_PASSWORD=your_secure_password_here
# Redis 配置
REDIS_HOST=localhost
REDIS_PORT=6379
REDIS_PASSWORD=aether
REDIS_PASSWORD=your_redis_password_here
# JWT密钥(使用 ./generate_keys.sh 生成)
# JWT密钥(使用 python generate_keys.py 生成)
# 用于用户登录 token 签名,更换后所有用户需重新登录
JWT_SECRET_KEY=change-this-to-a-secure-random-string
@@ -50,40 +21,91 @@ JWT_SECRET_KEY=change-this-to-a-secure-random-string
# 注意:更换此密钥后需要在管理面板重新配置所有 Provider API Key
ENCRYPTION_KEY=change-this-to-another-secure-random-string
# 启动自举管理员(仅在当前库里还没有活动管理员时生效)
# 手动部署时取消注释并设置;install.sh 首次生成配置时会提示输入。
# 支付回调共享密钥(公开 /api/payment/callback/* 入口必须携带 x-payment-callback-token)
# 建议使用 32+ 位随机字符串
PAYMENT_CALLBACK_SECRET=change-this-to-a-secure-callback-secret
# 管理员账号(仅首次初始化时使用, 创建完成后可在系统内修改密码)
ADMIN_EMAIL=[email protected]
ADMIN_USERNAME=admin123456
# ADMIN_PASSWORD=
ADMIN_USERNAME=admin
ADMIN_PASSWORD=admin123456
# ==================== 可选配置(有默认值) ====================
# 以下配置项有合理的默认值,可按需调整
# 可信反向代理 IP/CIDR,只有这些来源发送的 X-Real-IP / X-Forwarded-For 会被采用。
# 默认仅信任本机回环代理:127.0.0.0/8,::1/128。
# Docker/Nginx 位于独立容器时,请按实际容器网络设置,例如:172.16.0.0/12。
# AETHER_TRUSTED_PROXY_CIDRS=127.0.0.0/8,::1/128,172.16.0.0/12
# 应用端口(默认 8084)
# APP_PORT=8084
# docker compose 下 app 启动前自动执行 pending migration/backfill(默认 true)
# AETHER_GATEWAY_AUTO_PREPARE_DATABASE=true
# 生产部署镜像(deploy.sh 会读取)
# APP_IMAGE=ghcr.io/fawney19/aether:latest
# PostgreSQL 连接池配置(默认按 CPU 自动计算;正式高并发环境可显式预算)
# AETHER_GATEWAY_DATA_POSTGRES_MIN_CONNECTIONS=12
# AETHER_GATEWAY_DATA_POSTGRES_MAX_CONNECTIONS=80
# AETHER_GATEWAY_MAX_IN_FLIGHT_REQUESTS=2048
# AETHER_GATEWAY_REQUEST_BODY_BUFFER_BUDGET_MB=256
# AETHER_GATEWAY_REQUEST_BODY_READ_TIMEOUT_MS=120000
# AETHER_MAX_REQUEST_BODY_MB=64
# AETHER_GATEWAY_SECURITY_CACHE_TTL_MS=1000
# AETHER_MAX_REDACTED_SYNC_RESPONSE_BODY_MB=64
# AETHER_MAX_INTERNAL_BUFFERED_BODY_MB=128
# AETHER_TUNNEL_NODE_STATUS_QUEUE_CAPACITY=1024
# Gunicorn Worker 数量(默认 2)
# Tunnel 请求统一经 Hub 转发,可安全使用多 worker。
# 非 Docker 运行时若使用 ProxyNode tunnel,请确保 aether-hub 可达(默认 ws://127.0.0.1:8085)。
# GUNICORN_WORKERS=2
# PostgreSQL 容器调优:docker-compose.yml 已内置通用默认值,通常不用配置。
# 只有在 Postgres 独占大内存、或压测显示 DB 缓存/排序/维护任务成为瓶颈时再覆盖。
# 内置默认:shared_buffers=1GB, effective_cache_size=3GB, shm_size=512mb,
# work_mem=16MB, maintenance_work_mem=256MB。
# POSTGRES_SHARED_BUFFERS=8GB
# POSTGRES_EFFECTIVE_CACHE_SIZE=24GB
# POSTGRES_SHM_SIZE=2gb
# POSTGRES_WORK_MEM=16MB
# POSTGRES_MAINTENANCE_WORK_MEM=1GB
# Gunicorn Max Requests(默认 4000)
# Worker 处理指定数量请求后自动重启,防止内存泄漏
# max-requests-jitter 会自动设置为 MAX_REQUESTS/20 (5%)
# MAX_REQUESTS=4000
# glibc malloc arena 上限(默认 2)
# 降低 malloc 内存碎片,减少 gunicorn worker RSS
# MALLOC_ARENA_MAX=2
# HTTP 连接池上限(默认总预算约 200,按 worker 平分)
# 如果容器内存偏高,可继续下调;例如 2 worker 时设为 80-100
# HTTP_MAX_CONNECTIONS=100
# HTTP 保活连接数(默认约为 max_connections 的 30%)
# HTTP_KEEPALIVE_CONNECTIONS=30
# HTTP 代理/Tunnel 客户端空闲清理(默认每 5 分钟扫描,空闲 600 秒即关闭)
# HTTP_CLIENT_IDLE_CLEANUP_INTERVAL_MINUTES=5
# HTTP_CLIENT_IDLE_CLEANUP_MAX_SECONDS=600
# curl_cffi session 池上限(默认 20,按 impersonate + proxy 组合缓存)
# CURL_CFFI_MAX_SESSIONS=20
# 流式响应块缓存上限(单位 MB,默认 2)
# 说明:
# - 这是单个流式请求可保留的“解析后响应块”内存上限,不是全局上限
# - 粗略峰值内存 ≈ 并发流数量 × RESPONSE_CHUNKS_MAX_SIZE_MB
# 例如:100 并发、2MB 上限,理论峰值约 200MB
# - 建议:
# - 内存敏感环境:1
# - 通用生产环境:2(默认)
# - 需要更多调试上下文:4
# RESPONSE_CHUNKS_MAX_SIZE_MB=2
# 流式空闲超时(单位秒,默认 30)
# 当流已经开始但连续一段时间没有任何新 chunk 时,提前中断并返回 504,
# 避免一直等到 worker 超时(如 300s)
# STREAM_IDLE_TIMEOUT_SECONDS=30
# API Key 前缀(默认 sk)
# API_KEY_PREFIX=sk
# 日志级别(默认 INFO,可选:DEBUG, INFO, WARNING, ERROR)
# LOG_LEVEL=INFO
# CORS 配置(允许跨域的源,多个源用逗号分隔)
# 示例: http://localhost:3000,https://example.com
# 默认: * (允许所有源)
# CORS_ORIGINS=*
# 启动预热配置(默认启用,降低首请求冷启动延迟)
# 是否启用启动期预热任务(默认 true)
# STARTUP_WARMUP_ENABLED=true
# /readyz 是否等待预热完成(默认 true)
# STARTUP_WARMUP_GATE_READINESS=true
# 预热时优先 bootstrap 的 provider_type 列表(逗号分隔;留空表示自动探测)
# STARTUP_WARMUP_PROVIDER_TYPES=codex,kiro
# ==================== 计费系统(可选) ====================
# Video/Image/Audio 缺失 billing_rule 时是否拒绝请求(默认 false:允许请求但 cost=0 并告警)
# BILLING_REQUIRE_RULE=false
#
# required 维度缺失时是否拒绝请求/标记任务失败(默认 false:cost=0 + 标记 incomplete)
# BILLING_STRICT_MODE=false
#
+91
View File
@@ -0,0 +1,91 @@
name: Build aether-hub
on:
push:
tags: ['hub-v*']
workflow_dispatch:
permissions:
contents: write
jobs:
build:
name: ${{ matrix.name }}
runs-on: ubuntu-latest
strategy:
fail-fast: false
matrix:
include:
- name: linux-amd64
target: x86_64-unknown-linux-gnu
use_cross: true
- name: linux-arm64
target: aarch64-unknown-linux-gnu
use_cross: true
steps:
- uses: actions/checkout@v5
- name: Install Rust toolchain
uses: dtolnay/rust-toolchain@stable
with:
targets: ${{ matrix.target }}
- name: Rust cache
uses: Swatinem/rust-cache@v2
with:
workspaces: aether-hub -> target
key: ${{ matrix.target }}
- name: Install cross
if: matrix.use_cross
uses: taiki-e/install-action@cross
- name: Build
working-directory: aether-hub
shell: bash
run: |
if [ "${{ matrix.use_cross }}" = "true" ]; then
cross build --release --target ${{ matrix.target }}
else
cargo build --release --target ${{ matrix.target }}
fi
- name: Package
shell: bash
run: |
cd aether-hub/target/${{ matrix.target }}/release
chmod +x aether-hub
tar czf ../../../../aether-hub-${{ matrix.name }}.tar.gz aether-hub
- name: Upload artifact
uses: actions/upload-artifact@v5
with:
name: aether-hub-${{ matrix.name }}
path: aether-hub-*.tar.gz
if-no-files-found: error
release:
needs: build
runs-on: ubuntu-latest
if: startsWith(github.ref, 'refs/tags/')
steps:
- name: Download all artifacts
uses: actions/download-artifact@v5
with:
merge-multiple: true
path: artifacts
- name: Generate checksums
working-directory: artifacts
run: sha256sum aether-hub-* > SHA256SUMS.txt
- name: Create GitHub Release
uses: softprops/action-gh-release@v2
with:
name: "${{ github.ref_name }}"
generate_release_notes: true
files: |
artifacts/aether-hub-*
artifacts/SHA256SUMS.txt
fail_on_unmatched_files: true
+227
View File
@@ -0,0 +1,227 @@
name: Build aether-proxy
on:
push:
tags: ['proxy-v*']
workflow_dispatch:
permissions:
contents: write
packages: write
env:
REGISTRY: ghcr.io
GHCR_IMAGE: fawney19/aether-proxy
DOCKERHUB_IMAGE: fawney19/aether-proxy
jobs:
build:
name: ${{ matrix.name }}
runs-on: ${{ matrix.os }}
strategy:
fail-fast: false
matrix:
include:
- name: linux-amd64
target: x86_64-unknown-linux-gnu
os: ubuntu-latest
use_cross: true
- name: linux-arm64
target: aarch64-unknown-linux-gnu
os: ubuntu-latest
use_cross: true
- name: macos-amd64
target: x86_64-apple-darwin
os: macos-latest
use_cross: false
- name: macos-arm64
target: aarch64-apple-darwin
os: macos-latest
use_cross: false
- name: windows-amd64
target: x86_64-pc-windows-msvc
os: windows-latest
use_cross: false
steps:
- uses: actions/checkout@v5
- name: Install Rust toolchain
uses: dtolnay/rust-toolchain@stable
with:
targets: ${{ matrix.target }}
- name: Rust cache
uses: Swatinem/rust-cache@v2
with:
workspaces: aether-proxy -> target
key: ${{ matrix.target }}
- name: Install cross
if: matrix.use_cross
uses: taiki-e/install-action@cross
- name: Build
working-directory: aether-proxy
shell: bash
run: |
if [ "${{ matrix.use_cross }}" = "true" ]; then
cross build --release --target ${{ matrix.target }}
else
cargo build --release --target ${{ matrix.target }}
fi
- name: Package (Unix)
if: runner.os != 'Windows'
shell: bash
run: |
cd aether-proxy/target/${{ matrix.target }}/release
chmod +x aether-proxy
tar czf ../../../../aether-proxy-${{ matrix.name }}.tar.gz aether-proxy
- name: Package (Windows)
if: runner.os == 'Windows'
shell: bash
run: |
cd aether-proxy/target/${{ matrix.target }}/release
7z a ../../../../aether-proxy-${{ matrix.name }}.zip aether-proxy.exe
- name: Upload artifact
uses: actions/upload-artifact@v5
with:
name: aether-proxy-${{ matrix.name }}
path: |
aether-proxy-*.tar.gz
aether-proxy-*.zip
if-no-files-found: error
release:
needs: build
runs-on: ubuntu-latest
if: startsWith(github.ref, 'refs/tags/')
steps:
- name: Download all artifacts
uses: actions/download-artifact@v5
with:
merge-multiple: true
path: artifacts
- name: Generate checksums
working-directory: artifacts
run: sha256sum aether-proxy-* > SHA256SUMS.txt
- name: Create GitHub Release
uses: softprops/action-gh-release@v2
with:
name: "${{ github.ref_name }}"
generate_release_notes: true
files: |
artifacts/aether-proxy-*
artifacts/SHA256SUMS.txt
fail_on_unmatched_files: true
docker:
needs: build
runs-on: ubuntu-latest
if: startsWith(github.ref, 'refs/tags/')
steps:
- uses: actions/checkout@v5
- name: Download Linux artifacts
uses: actions/download-artifact@v5
with:
pattern: aether-proxy-linux-*
merge-multiple: true
path: artifacts
- name: Prepare binaries
run: |
mkdir -p aether-proxy/build/linux-amd64 aether-proxy/build/linux-arm64
tar xzf artifacts/aether-proxy-linux-amd64.tar.gz -C aether-proxy/build/linux-amd64
tar xzf artifacts/aether-proxy-linux-arm64.tar.gz -C aether-proxy/build/linux-arm64
- name: Set up QEMU
uses: docker/setup-qemu-action@v3
- name: Set up Docker Buildx
uses: docker/setup-buildx-action@v3
- name: Log in to GHCR
uses: docker/login-action@v3
with:
registry: ${{ env.REGISTRY }}
username: ${{ github.actor }}
password: ${{ secrets.GITHUB_TOKEN }}
- name: Log in to Docker Hub
uses: docker/login-action@v3
with:
username: ${{ secrets.DOCKERHUB_USERNAME }}
password: ${{ secrets.DOCKERHUB_TOKEN }}
- name: Extract metadata
id: meta
uses: docker/metadata-action@v5
with:
images: |
${{ env.REGISTRY }}/${{ env.GHCR_IMAGE }}
docker.io/${{ env.DOCKERHUB_IMAGE }}
tags: |
type=match,pattern=proxy-v(.*),group=1
type=match,pattern=proxy-v(\d+\.\d+),group=1
type=sha,prefix=
flavor: |
latest=auto
- name: Build and push
uses: docker/build-push-action@v6
with:
context: ./aether-proxy
push: true
tags: ${{ steps.meta.outputs.tags }}
labels: ${{ steps.meta.outputs.labels }}
platforms: linux/amd64,linux/arm64
update-readme:
needs: release
runs-on: ubuntu-latest
if: startsWith(github.ref, 'refs/tags/')
steps:
- uses: actions/checkout@v5
with:
ref: master
- name: Update README download links
env:
TAG: ${{ github.ref_name }}
run: |
VERSION="${TAG#proxy-v}"
BASE="https://github.com/fawney19/Aether/releases/download/${TAG}"
cd aether-proxy
TABLE="| Platform | Download |\n|----------|----------|\n"
TABLE+="| Linux x86_64 | [aether-proxy-linux-amd64.tar.gz](${BASE}/aether-proxy-linux-amd64.tar.gz) |\n"
TABLE+="| Linux ARM64 | [aether-proxy-linux-arm64.tar.gz](${BASE}/aether-proxy-linux-arm64.tar.gz) |\n"
TABLE+="| macOS x86_64 | [aether-proxy-macos-amd64.tar.gz](${BASE}/aether-proxy-macos-amd64.tar.gz) |\n"
TABLE+="| macOS ARM64 | [aether-proxy-macos-arm64.tar.gz](${BASE}/aether-proxy-macos-arm64.tar.gz) |\n"
TABLE+="| Windows x86_64 | [aether-proxy-windows-amd64.zip](${BASE}/aether-proxy-windows-amd64.zip) |"
# Replace content between markers
if grep -q '<!-- DOWNLOAD_TABLE_START -->' README.md; then
awk -v table="$TABLE" '
/<!-- DOWNLOAD_TABLE_START -->/ { print; printf "%s\n", table; skip=1; next }
/<!-- DOWNLOAD_TABLE_END -->/ { skip=0 }
!skip { print }
' README.md > README.tmp && mv README.tmp README.md
fi
- name: Commit and push
run: |
cd aether-proxy
git config user.name "github-actions[bot]"
git config user.email "github-actions[bot]@users.noreply.github.com"
git add README.md
git diff --cached --quiet && exit 0
TAG="${GITHUB_REF#refs/tags/}"
git commit -m "chore(proxy): update download links for ${TAG}"
git push
-239
View File
@@ -1,239 +0,0 @@
name: Build aether-tunnel
on:
push:
tags: ['tunnel-v*']
workflow_dispatch:
permissions:
contents: write
concurrency:
group: build-tunnel-${{ github.ref }}
cancel-in-progress: false
jobs:
preflight:
runs-on: ubuntu-latest
if: startsWith(github.ref, 'refs/tags/')
steps:
- uses: actions/checkout@v5
- name: Ensure tunnel tag matches Cargo version
shell: bash
run: |
TAG="${GITHUB_REF_NAME}"
EXPECTED="${TAG#tunnel-v}"
ACTUAL="$(cargo metadata --manifest-path apps/aether-tunnel/Cargo.toml --locked --no-deps --format-version 1 | jq -r '.packages[] | select(.name == "aether-tunnel") | .version')"
echo "tag version: ${EXPECTED}"
echo "cargo version: ${ACTUAL}"
if [ -z "${ACTUAL}" ]; then
echo "Could not resolve aether-tunnel package version" >&2
exit 1
fi
if [ "${EXPECTED}" != "${ACTUAL}" ]; then
echo "tunnel tag ${TAG} does not match apps/aether-tunnel/Cargo.toml version ${ACTUAL}" >&2
exit 1
fi
build:
needs: preflight
if: always() && (needs.preflight.result == 'success' || needs.preflight.result == 'skipped')
name: ${{ matrix.name }}
runs-on: ${{ matrix.os }}
strategy:
fail-fast: false
matrix:
include:
- name: linux-amd64
target: x86_64-unknown-linux-gnu
os: ubuntu-latest
use_cross: true
- name: linux-arm64
target: aarch64-unknown-linux-gnu
os: ubuntu-latest
use_cross: true
- name: linux-musl-amd64
target: x86_64-unknown-linux-musl
os: ubuntu-latest
use_cross: true
- name: linux-musl-arm64
target: aarch64-unknown-linux-musl
os: ubuntu-latest
use_cross: true
- name: macos-amd64
target: x86_64-apple-darwin
os: macos-15-intel
use_cross: false
- name: macos-arm64
target: aarch64-apple-darwin
os: macos-15
use_cross: false
- name: windows-amd64
target: x86_64-pc-windows-msvc
os: windows-latest
use_cross: false
steps:
- uses: actions/checkout@v5
- name: Install Rust toolchain
uses: dtolnay/rust-toolchain@stable
with:
targets: ${{ matrix.target }}
- name: Ensure Rust target is installed
run: rustup target add ${{ matrix.target }}
- name: Rust cache
uses: Swatinem/rust-cache@v2
with:
workspaces: apps/aether-tunnel -> target
key: ${{ matrix.target }}
- name: Install cross
if: matrix.use_cross
uses: taiki-e/install-action@cross
- name: Build
working-directory: apps/aether-tunnel
shell: bash
run: |
if [ "${{ matrix.use_cross }}" = "true" ]; then
cross build --release --locked --target ${{ matrix.target }}
else
cargo build --release --locked --target ${{ matrix.target }}
fi
- name: Package (Unix)
if: runner.os != 'Windows'
shell: bash
run: |
cd target/${{ matrix.target }}/release
chmod +x aether-tunnel
tar czf ../../../aether-tunnel-${{ matrix.name }}.tar.gz aether-tunnel
- name: Package (Windows)
if: runner.os == 'Windows'
shell: bash
run: |
cd target/${{ matrix.target }}/release
7z a ../../../aether-tunnel-${{ matrix.name }}.zip aether-tunnel.exe
- name: Upload artifact
uses: actions/upload-artifact@v5
with:
name: aether-tunnel-${{ matrix.name }}
path: |
aether-tunnel-*.tar.gz
aether-tunnel-*.zip
if-no-files-found: error
retention-days: 1
release:
needs: build
runs-on: ubuntu-latest
if: startsWith(github.ref, 'refs/tags/')
steps:
- name: Download all artifacts
uses: actions/download-artifact@v5
with:
merge-multiple: true
path: artifacts
- name: Generate checksums
working-directory: artifacts
run: sha256sum aether-tunnel-* > SHA256SUMS.txt
- name: Delete stale draft releases for tag
env:
GH_TOKEN: ${{ github.token }}
RELEASE_TAG: ${{ github.ref_name }}
REPOSITORY: ${{ github.repository }}
shell: bash
run: |
set -euo pipefail
draft_ids="$(gh api "repos/${REPOSITORY}/releases" --paginate --jq '.[] | select(.tag_name == env.RELEASE_TAG and .draft == true) | .id')"
if [[ -z "${draft_ids}" ]]; then
echo "No stale draft releases for ${RELEASE_TAG}"
exit 0
fi
while IFS= read -r release_id; do
[[ -z "${release_id}" ]] && continue
echo "Deleting stale draft release ${release_id} for ${RELEASE_TAG}"
gh api -X DELETE "repos/${REPOSITORY}/releases/${release_id}"
done <<< "${draft_ids}"
- name: Create GitHub Release
uses: softprops/action-gh-release@v2
with:
name: "${{ github.ref_name }}"
generate_release_notes: true
files: |
artifacts/aether-tunnel-*
artifacts/SHA256SUMS.txt
fail_on_unmatched_files: true
update-readme:
needs: release
runs-on: ubuntu-latest
if: startsWith(github.ref, 'refs/tags/')
steps:
- uses: actions/checkout@v5
with:
ref: main
- name: Update README download links
env:
TAG: ${{ github.ref_name }}
run: |
VERSION="${TAG#tunnel-v}"
BASE="https://github.com/fawney19/Aether/releases/download/${TAG}"
if [ -d apps/aether-tunnel ]; then
TUNNEL_DIR="apps/aether-tunnel"
else
TUNNEL_DIR="aether-tunnel"
fi
cd "$TUNNEL_DIR"
TABLE="| Platform | Download |\n|----------|----------|\n"
TABLE+="| Linux x86_64 (GNU) | [aether-tunnel-linux-amd64.tar.gz](${BASE}/aether-tunnel-linux-amd64.tar.gz) |\n"
TABLE+="| Linux ARM64 (GNU) | [aether-tunnel-linux-arm64.tar.gz](${BASE}/aether-tunnel-linux-arm64.tar.gz) |\n"
TABLE+="| Linux x86_64 (musl) | [aether-tunnel-linux-musl-amd64.tar.gz](${BASE}/aether-tunnel-linux-musl-amd64.tar.gz) |\n"
TABLE+="| Linux ARM64 (musl) | [aether-tunnel-linux-musl-arm64.tar.gz](${BASE}/aether-tunnel-linux-musl-arm64.tar.gz) |\n"
TABLE+="| macOS x86_64 | [aether-tunnel-macos-amd64.tar.gz](${BASE}/aether-tunnel-macos-amd64.tar.gz) |\n"
TABLE+="| macOS ARM64 | [aether-tunnel-macos-arm64.tar.gz](${BASE}/aether-tunnel-macos-arm64.tar.gz) |\n"
TABLE+="| Windows x86_64 | [aether-tunnel-windows-amd64.zip](${BASE}/aether-tunnel-windows-amd64.zip) |"
# Replace content between markers
if grep -q '<!-- DOWNLOAD_TABLE_START -->' README.md; then
awk -v table="$TABLE" '
/<!-- DOWNLOAD_TABLE_START -->/ { print; printf "%s\n", table; skip=1; next }
/<!-- DOWNLOAD_TABLE_END -->/ { skip=0 }
!skip { print }
' README.md > README.tmp && mv README.tmp README.md
fi
- name: Commit and push
run: |
if [ -d apps/aether-tunnel ]; then
TUNNEL_DIR="apps/aether-tunnel"
else
TUNNEL_DIR="aether-tunnel"
fi
cd "$TUNNEL_DIR"
git config user.name "github-actions[bot]"
git config user.email "github-actions[bot]@users.noreply.github.com"
git add README.md
git diff --cached --quiet && exit 0
TAG="${GITHUB_REF#refs/tags/}"
git commit -m "chore(tunnel): update download links for ${TAG}"
git push
-28
View File
@@ -15,35 +15,7 @@ concurrency:
cancel-in-progress: false
jobs:
preflight:
runs-on: ubuntu-latest
outputs:
deploy_pages: ${{ steps.classify.outputs.deploy_pages }}
steps:
- name: Ensure stable Pages release tag
id: classify
shell: bash
run: |
set -euo pipefail
echo "deploy_pages=false" >> "${GITHUB_OUTPUT}"
if [[ "${GITHUB_REF_TYPE}" != "tag" ]]; then
echo "Manual Pages deployment."
echo "deploy_pages=true" >> "${GITHUB_OUTPUT}"
exit 0
fi
tag="${GITHUB_REF_NAME}"
if [[ ! "${tag}" =~ ^v[0-9]+\.[0-9]+\.[0-9]+$ ]]; then
echo "Skipping Pages deploy for non-stable release tag: ${tag}"
exit 0
fi
echo "deploy_pages=true" >> "${GITHUB_OUTPUT}"
build:
needs: preflight
if: needs.preflight.outputs.deploy_pages == 'true'
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v5
+327
View File
@@ -0,0 +1,327 @@
name: Build and Publish Docker Image
on:
push:
tags: ['v*']
workflow_dispatch:
inputs:
build_base:
description: 'Rebuild base image'
required: false
default: false
type: boolean
env:
REGISTRY: ghcr.io
BASE_IMAGE_NAME: fawney19/aether-base
APP_IMAGE_NAME: fawney19/aether
GITHUB_REPO: fawney19/Aether
# Base image hash inputs:
# - Dockerfile.base
# - pyproject.toml (dependency fingerprint only; ignores tool/optional deps)
# - frontend/package-lock.json
jobs:
check-base-changes:
runs-on: ubuntu-latest
permissions:
contents: read
packages: read
outputs:
base_changed: ${{ steps.check.outputs.base_changed }}
steps:
- uses: actions/checkout@v5
- name: Log in to Container Registry
uses: docker/login-action@v3
with:
registry: ${{ env.REGISTRY }}
username: ${{ github.actor }}
password: ${{ secrets.GITHUB_TOKEN }}
- name: Check if base image needs rebuild
id: check
run: |
if [ "${{ github.event.inputs.build_base }}" == "true" ]; then
echo "base_changed=true" >> $GITHUB_OUTPUT
exit 0
fi
# Calculate current hash of base-related inputs (dependency-only fingerprint)
PY_FINGERPRINT=$(python3 - <<'PY'
import json
import pathlib
import tomllib
data = tomllib.loads(pathlib.Path("pyproject.toml").read_text("utf-8"))
project = data.get("project") or {}
build = data.get("build-system") or {}
fingerprint = {
"requires-python": project.get("requires-python"),
"dependencies": sorted(project.get("dependencies") or []),
"build-backend": build.get("build-backend"),
"build-requires": sorted(build.get("requires") or []),
}
print(json.dumps(fingerprint, sort_keys=True, separators=(",", ":")))
PY
)
CURRENT_HASH=$(
(
cat Dockerfile.base
printf '%s\n' "$PY_FINGERPRINT"
cat frontend/package-lock.json
) | sha256sum | cut -d' ' -f1
)
echo "Current base hash: $CURRENT_HASH"
# Try to get hash label from remote image config
# Pull the image config and extract labels
REMOTE_HASH=""
if docker pull ${{ env.REGISTRY }}/${{ env.BASE_IMAGE_NAME }}:latest; then
REMOTE_HASH=$(docker inspect ${{ env.REGISTRY }}/${{ env.BASE_IMAGE_NAME }}:latest --format '{{ index .Config.Labels "org.opencontainers.image.base.hash" }}' 2>/dev/null) || true
else
echo "WARN: failed to pull remote base image; forcing base rebuild."
echo "base_changed=true" >> $GITHUB_OUTPUT
exit 0
fi
if [ -z "$REMOTE_HASH" ] || [ "$REMOTE_HASH" == "<no value>" ]; then
# No remote image or no hash label, need to rebuild
echo "No remote base image or hash label found, need rebuild"
echo "base_changed=true" >> $GITHUB_OUTPUT
elif [ "$CURRENT_HASH" != "$REMOTE_HASH" ]; then
echo "Hash mismatch: remote=$REMOTE_HASH, current=$CURRENT_HASH"
echo "base_changed=true" >> $GITHUB_OUTPUT
else
echo "Hash matches, no rebuild needed"
echo "base_changed=false" >> $GITHUB_OUTPUT
fi
build-base:
needs: check-base-changes
if: needs.check-base-changes.outputs.base_changed == 'true'
runs-on: ubuntu-latest
permissions:
contents: read
packages: write
steps:
- uses: actions/checkout@v5
- name: Set up Docker Buildx
uses: docker/setup-buildx-action@v3
- name: Log in to Container Registry
uses: docker/login-action@v3
with:
registry: ${{ env.REGISTRY }}
username: ${{ github.actor }}
password: ${{ secrets.GITHUB_TOKEN }}
- name: Calculate base files hash
id: hash
run: |
PY_FINGERPRINT=$(python3 - <<'PY'
import json
import pathlib
import tomllib
data = tomllib.loads(pathlib.Path("pyproject.toml").read_text("utf-8"))
project = data.get("project") or {}
build = data.get("build-system") or {}
fingerprint = {
"requires-python": project.get("requires-python"),
"dependencies": sorted(project.get("dependencies") or []),
"build-backend": build.get("build-backend"),
"build-requires": sorted(build.get("requires") or []),
}
print(json.dumps(fingerprint, sort_keys=True, separators=(",", ":")))
PY
)
HASH=$(
(
cat Dockerfile.base
printf '%s\n' "$PY_FINGERPRINT"
cat frontend/package-lock.json
) | sha256sum | cut -d' ' -f1
)
echo "hash=$HASH" >> $GITHUB_OUTPUT
- name: Extract metadata for base image
id: meta
uses: docker/metadata-action@v5
with:
images: ${{ env.REGISTRY }}/${{ env.BASE_IMAGE_NAME }}
tags: |
type=raw,value=latest
type=sha,prefix=
labels: |
org.opencontainers.image.base.hash=${{ steps.hash.outputs.hash }}
- name: Build and push base image
uses: docker/build-push-action@v6
with:
context: .
file: ./Dockerfile.base
push: true
tags: ${{ steps.meta.outputs.tags }}
labels: ${{ steps.meta.outputs.labels }}
cache-from: type=gha,scope=base
cache-to: type=gha,mode=max,scope=base
platforms: linux/amd64,linux/arm64
download-hub:
runs-on: ubuntu-latest
permissions:
contents: read
outputs:
hub_tag: ${{ steps.hub-tag.outputs.tag }}
steps:
- name: Get latest hub release tag
id: hub-tag
env:
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
run: |
TAG=$(gh release list --repo "${{ env.GITHUB_REPO }}" --limit 50 --json tagName,isDraft,isPrerelease \
--jq '[.[] | select(.tagName | startswith("hub-v")) | select(.isDraft == false and .isPrerelease == false)] | .[0].tagName')
if [ -z "$TAG" ] || [ "$TAG" = "null" ]; then
echo "No hub release found"
exit 1
fi
echo "tag=$TAG" >> $GITHUB_OUTPUT
echo "Hub release tag: $TAG"
build-app:
needs: [check-base-changes, build-base, download-hub]
if: always() && (needs.build-base.result == 'success' || needs.build-base.result == 'skipped') && needs.download-hub.result == 'success'
runs-on: ubuntu-latest
permissions:
contents: read
packages: write
steps:
- uses: actions/checkout@v5
- name: Set up Docker Buildx
uses: docker/setup-buildx-action@v3
- name: Log in to Container Registry
uses: docker/login-action@v3
with:
registry: ${{ env.REGISTRY }}
username: ${{ github.actor }}
password: ${{ secrets.GITHUB_TOKEN }}
- name: Log in to Docker Hub
uses: docker/login-action@v3
with:
username: ${{ secrets.DOCKERHUB_USERNAME }}
password: ${{ secrets.DOCKERHUB_TOKEN }}
- name: Extract metadata for app image
id: meta
uses: docker/metadata-action@v5
with:
images: |
${{ env.REGISTRY }}/${{ env.APP_IMAGE_NAME }}
docker.io/fawney19/aether
tags: |
type=semver,pattern={{version}}
type=semver,pattern={{major}}.{{minor}}
type=raw,value=pre,enable=${{ contains(github.ref, '-') }}
type=raw,value=fix,enable=${{ contains(github.ref, '-fix') }}
type=sha,prefix=
flavor: |
latest=auto
- name: Extract version from tag
id: version
run: |
# 从 tag 提取版本号,如 v0.2.5 -> 0.2.5
VERSION="${GITHUB_REF#refs/tags/v}"
if [ "$VERSION" = "$GITHUB_REF" ]; then
# 不是 tag 触发,使用 git describe
VERSION=$(git describe --tags --always | sed 's/^v//')
fi
echo "version=$VERSION" >> $GITHUB_OUTPUT
echo "Extracted version: $VERSION"
- name: Update Dockerfile.app to use registry base image
run: |
sed -i "s|FROM aether-base:latest AS builder|FROM ${{ env.REGISTRY }}/${{ env.BASE_IMAGE_NAME }}:latest AS builder|g" Dockerfile.app
- name: Generate version file
run: |
# 生成 _version.py 文件
cat > src/_version.py << EOF
# Auto-generated by CI
__version__ = '${{ steps.version.outputs.version }}'
__version_tuple__ = tuple(int(x) for x in '${{ steps.version.outputs.version }}'.split('.') if x.isdigit())
version = __version__
version_tuple = __version_tuple__
EOF
- name: Resolve hub release for build args
run: |
echo "Hub release tag: ${{ needs.download-hub.outputs.hub_tag }}"
- name: Build and push app image (amd64)
id: build-amd64
uses: docker/build-push-action@v6
with:
context: .
file: ./Dockerfile.app
labels: ${{ steps.meta.outputs.labels }}
no-cache-filters: builder
cache-from: type=gha,scope=app-amd64
cache-to: type=gha,mode=min,scope=app-amd64
build-args: |
HUB_RELEASE_REPO=${{ env.GITHUB_REPO }}
HUB_TAG=${{ needs.download-hub.outputs.hub_tag }}
platforms: linux/amd64
outputs: type=image,"name=${{ env.REGISTRY }}/${{ env.APP_IMAGE_NAME }},docker.io/fawney19/aether",push-by-digest=true,name-canonical=true,push=true
- name: Build and push app image (arm64)
id: build-arm64
uses: docker/build-push-action@v6
with:
context: .
file: ./Dockerfile.app
labels: ${{ steps.meta.outputs.labels }}
no-cache-filters: builder
cache-from: type=gha,scope=app-arm64
cache-to: type=gha,mode=min,scope=app-arm64
build-args: |
HUB_RELEASE_REPO=${{ env.GITHUB_REPO }}
HUB_TAG=${{ needs.download-hub.outputs.hub_tag }}
platforms: linux/arm64
outputs: type=image,"name=${{ env.REGISTRY }}/${{ env.APP_IMAGE_NAME }},docker.io/fawney19/aether",push-by-digest=true,name-canonical=true,push=true
- name: Create multi-arch manifest and push
run: |
# Extract digests
AMD64_DIGEST="${{ steps.build-amd64.outputs.digest }}"
ARM64_DIGEST="${{ steps.build-arm64.outputs.digest }}"
echo "amd64 digest: $AMD64_DIGEST"
echo "arm64 digest: $ARM64_DIGEST"
# For each tag, create multi-arch manifest on each registry
TAGS=$(echo "${{ steps.meta.outputs.tags }}" | tr '\n' ' ')
for FULL_TAG in $TAGS; do
# Determine which registry this tag belongs to
if [[ "$FULL_TAG" == ghcr.io/* ]]; then
REPO="${{ env.REGISTRY }}/${{ env.APP_IMAGE_NAME }}"
elif [[ "$FULL_TAG" == docker.io/* ]]; then
REPO="docker.io/fawney19/aether"
else
continue
fi
echo "Creating manifest for $FULL_TAG"
docker buildx imagetools create -t "$FULL_TAG" \
"$REPO@$AMD64_DIGEST" \
"$REPO@$ARM64_DIGEST"
done
-342
View File
@@ -1,342 +0,0 @@
name: Release Aether
on:
push:
tags: ['v*']
workflow_dispatch:
permissions:
contents: write
packages: write
concurrency:
group: release-aether-${{ github.ref }}
cancel-in-progress: false
env:
REGISTRY: ghcr.io
GHCR_IMAGE: fawney19/aether
DOCKERHUB_IMAGE: fawney19/aether
jobs:
preflight:
name: Release preflight
runs-on: ubuntu-latest
outputs:
publish: ${{ steps.classify.outputs.publish }}
version_tag: ${{ steps.classify.outputs.version_tag }}
prerelease: ${{ steps.classify.outputs.prerelease }}
make_latest: ${{ steps.classify.outputs.make_latest }}
steps:
- name: Classify release tag
id: classify
shell: bash
run: |
set -euo pipefail
echo "publish=false" >> "${GITHUB_OUTPUT}"
echo "version_tag=" >> "${GITHUB_OUTPUT}"
echo "prerelease=false" >> "${GITHUB_OUTPUT}"
echo "make_latest=false" >> "${GITHUB_OUTPUT}"
if [[ "${GITHUB_REF_TYPE}" != "tag" ]]; then
echo "Manual release build; publish jobs will be skipped."
exit 0
fi
tag="${GITHUB_REF_NAME}"
if [[ ! "${tag}" =~ ^v[0-9]+\.[0-9]+\.[0-9]+(-(beta|rc)\.[0-9]+)?$ ]]; then
echo "Unsupported release tag: ${tag}" >&2
echo "Expected vX.Y.Z, vX.Y.Z-beta.N, or vX.Y.Z-rc.N." >&2
exit 1
fi
echo "version_tag=${tag}" >> "${GITHUB_OUTPUT}"
if [[ "${tag}" == *-* ]]; then
echo "prerelease=true" >> "${GITHUB_OUTPUT}"
else
echo "make_latest=true" >> "${GITHUB_OUTPUT}"
fi
if [[ "${GITHUB_EVENT_NAME}" == "push" ]]; then
echo "publish=true" >> "${GITHUB_OUTPUT}"
else
echo "Manual release build for ${tag}; publish jobs will be skipped."
fi
frontend:
name: Build frontend
needs: preflight
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v5
- name: Setup Node.js
uses: actions/setup-node@v4
with:
node-version: 22
cache: npm
cache-dependency-path: frontend/package-lock.json
- name: Install & build
working-directory: frontend
run: |
npm ci
npm run build
- name: Upload frontend artifact
uses: actions/upload-artifact@v5
with:
name: frontend-dist
path: frontend/dist/
if-no-files-found: error
retention-days: 1
build:
name: Build ${{ matrix.name }}
needs: preflight
runs-on: ${{ matrix.os }}
strategy:
fail-fast: true
matrix:
include:
- name: linux-amd64
target: x86_64-unknown-linux-musl
platform: linux
arch: amd64
os: ubuntu-latest
use_cross: true
- name: linux-arm64
target: aarch64-unknown-linux-musl
platform: linux
arch: arm64
os: ubuntu-latest
use_cross: true
- name: macos-amd64
target: x86_64-apple-darwin
platform: macos
arch: amd64
os: macos-15-intel
use_cross: false
- name: macos-arm64
target: aarch64-apple-darwin
platform: macos
arch: arm64
os: macos-15
use_cross: false
steps:
- uses: actions/checkout@v5
- name: Install Rust toolchain
uses: dtolnay/rust-toolchain@stable
with:
targets: ${{ matrix.target }}
- name: Rust cache
uses: Swatinem/rust-cache@v2
with:
shared-key: release-${{ matrix.target }}
workspaces: . -> target
- name: Install cross
if: matrix.use_cross
uses: taiki-e/install-action@cross
- name: Build
env:
AETHER_VERSION: ${{ needs.preflight.outputs.version_tag }}
AETHER_BUILD_TYPE: release
CARGO_TERM_COLOR: always
shell: bash
run: |
if [[ "${{ matrix.use_cross }}" == "true" ]]; then
cross build --release --locked -p aether-gateway --target ${{ matrix.target }}
else
cargo build --release --locked -p aether-gateway --target ${{ matrix.target }}
fi
- name: Upload binary artifact
uses: actions/upload-artifact@v5
with:
name: aether-gateway-${{ matrix.platform }}-${{ matrix.arch }}
path: target/${{ matrix.target }}/release/aether-gateway
if-no-files-found: error
retention-days: 1
docker:
name: Docker multi-arch
needs: [preflight, frontend, build]
if: needs.preflight.outputs.publish == 'true'
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v5
- name: Download all artifacts
uses: actions/download-artifact@v5
with:
path: artifacts
- name: Prepare dist layout
run: |
mkdir -p dist
cp artifacts/aether-gateway-linux-amd64/aether-gateway dist/aether-gateway-amd64
cp artifacts/aether-gateway-linux-arm64/aether-gateway dist/aether-gateway-arm64
chmod +x dist/aether-gateway-amd64 dist/aether-gateway-arm64
cp -r artifacts/frontend-dist dist/frontend
- name: Set up QEMU
uses: docker/setup-qemu-action@v3
- name: Set up Docker Buildx
uses: docker/setup-buildx-action@v3
- name: Log in to GHCR
uses: docker/login-action@v3
with:
registry: ${{ env.REGISTRY }}
username: ${{ github.actor }}
password: ${{ secrets.GITHUB_TOKEN }}
- name: Log in to Docker Hub
uses: docker/login-action@v3
with:
username: ${{ secrets.DOCKERHUB_USERNAME }}
password: ${{ secrets.DOCKERHUB_TOKEN }}
- name: Extract metadata
id: meta
uses: docker/metadata-action@v5
with:
images: |
${{ env.REGISTRY }}/${{ env.GHCR_IMAGE }}
docker.io/${{ env.DOCKERHUB_IMAGE }}
tags: |
type=semver,pattern={{version}}
type=semver,pattern={{major}}.{{minor}},enable=${{ needs.preflight.outputs.make_latest == 'true' }}
type=raw,value=latest,enable=${{ needs.preflight.outputs.make_latest == 'true' }}
type=raw,value=beta,enable=${{ contains(github.ref_name, '-beta.') }}
type=raw,value=rc,enable=${{ contains(github.ref_name, '-rc.') }}
type=sha,prefix=
flavor: |
latest=false
- name: Build and push
uses: docker/build-push-action@v6
with:
context: .
file: ./Dockerfile.app
push: true
tags: ${{ steps.meta.outputs.tags }}
labels: ${{ steps.meta.outputs.labels }}
platforms: linux/amd64,linux/arm64
package:
name: Release tarballs
needs: [preflight, frontend, build]
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v5
- name: Download all artifacts
uses: actions/download-artifact@v5
with:
path: artifacts
- name: Build release packages
run: |
set -euo pipefail
if [[ "${GITHUB_REF_TYPE}" == "tag" ]]; then
VERSION="${GITHUB_REF_NAME}"
SOURCE_REF="${GITHUB_REF_NAME}"
else
VERSION="snapshot-${GITHUB_SHA::7}"
SOURCE_REF="${GITHUB_SHA}"
fi
mkdir -p package release-assets
for platform in linux macos; do
for arch in amd64 arm64; do
bundle="aether-${VERSION}-${platform}-${arch}"
root="package/${bundle}"
mkdir -p \
"${root}/bin" \
"${root}/frontend"
install -m 0755 "artifacts/aether-gateway-${platform}-${arch}/aether-gateway" "${root}/bin/aether-gateway"
cp -R artifacts/frontend-dist/. "${root}/frontend/"
sed \
-e "s/^SOURCE_REF=\"\${AETHER_SOURCE_REF:-main}\"/SOURCE_REF=\"\${AETHER_SOURCE_REF:-${SOURCE_REF}}\"/" \
-e "s/^VERSION=\"\${AETHER_VERSION:-}\"/VERSION=\"\${AETHER_VERSION:-${VERSION}}\"/" \
install.sh > "${root}/install.sh"
chmod 0755 "${root}/install.sh"
install -m 0755 update.sh "${root}/update.sh"
install -m 0644 docker-compose.yml "${root}/docker-compose.yml"
install -m 0644 docker-compose.single-node.yml "${root}/docker-compose.single-node.yml"
install -m 0644 .env.example "${root}/.env.example"
install -m 0755 generate_keys.sh "${root}/generate_keys.sh"
install -m 0644 README.md "${root}/README.md"
install -m 0644 LICENSE "${root}/LICENSE"
tar -C package -czf "release-assets/${bundle}.tar.gz" "${bundle}"
done
done
sed \
-e "s/^SOURCE_REF=\"\${AETHER_SOURCE_REF:-main}\"/SOURCE_REF=\"\${AETHER_SOURCE_REF:-${SOURCE_REF}}\"/" \
-e "s/^VERSION=\"\${AETHER_VERSION:-}\"/VERSION=\"\${AETHER_VERSION:-${VERSION}}\"/" \
install.sh > release-assets/install.sh
chmod +x release-assets/install.sh
(cd release-assets && sha256sum *.tar.gz > SHA256SUMS)
- name: Upload release package artifact
uses: actions/upload-artifact@v5
with:
name: release-assets
path: release-assets/*
if-no-files-found: error
retention-days: 7
github-release:
name: GitHub Release assets
needs: [preflight, docker, package]
if: needs.preflight.outputs.publish == 'true'
runs-on: ubuntu-latest
steps:
- name: Download release package artifact
uses: actions/download-artifact@v5
with:
name: release-assets
path: release-assets
- name: Delete stale draft releases for tag
env:
GH_TOKEN: ${{ github.token }}
RELEASE_TAG: ${{ github.ref_name }}
REPOSITORY: ${{ github.repository }}
shell: bash
run: |
set -euo pipefail
draft_ids="$(gh api "repos/${REPOSITORY}/releases" --paginate --jq '.[] | select(.tag_name == env.RELEASE_TAG and .draft == true) | .id')"
if [[ -z "${draft_ids}" ]]; then
echo "No stale draft releases for ${RELEASE_TAG}"
exit 0
fi
while IFS= read -r release_id; do
[[ -z "${release_id}" ]] && continue
echo "Deleting stale draft release ${release_id} for ${RELEASE_TAG}"
gh api -X DELETE "repos/${REPOSITORY}/releases/${release_id}"
done <<< "${draft_ids}"
- name: Publish GitHub Release assets
uses: softprops/action-gh-release@v2
with:
generate_release_notes: true
prerelease: ${{ needs.preflight.outputs.prerelease }}
make_latest: ${{ needs.preflight.outputs.make_latest }}
files: |
release-assets/*.tar.gz
release-assets/SHA256SUMS
release-assets/install.sh
-681
View File
@@ -1,681 +0,0 @@
name: Rust CI
on:
push:
branches:
- master
- main
paths:
- "Cargo.toml"
- "Cargo.lock"
- "crates/**"
- "apps/**"
- ".github/workflows/rust-ci.yml"
pull_request:
paths:
- "Cargo.toml"
- "Cargo.lock"
- "crates/**"
- "apps/**"
- ".github/workflows/rust-ci.yml"
concurrency:
group: rust-ci-${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
cancel-in-progress: true
env:
CARGO_INCREMENTAL: 0
CARGO_PROFILE_DEV_DEBUG: 0
CARGO_PROFILE_TEST_DEBUG: 0
CARGO_TERM_COLOR: always
jobs:
fmt:
name: Format
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v5
- name: Install Rust toolchain
uses: dtolnay/rust-toolchain@stable
with:
toolchain: 1.95.0
components: rustfmt
- name: Format
run: cargo fmt --all --check
clippy_gateway:
name: Clippy (Gateway)
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v5
- name: Install Rust toolchain
uses: dtolnay/rust-toolchain@stable
with:
toolchain: 1.95.0
components: clippy
- name: Rust cache
uses: Swatinem/rust-cache@v2
with:
shared-key: rust-ci-${{ runner.os }}
workspaces: . -> target
- name: Setup sccache
uses: mozilla-actions/[email protected]
- name: Clippy
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
run: cargo clippy -p aether-gateway --lib --bins --examples -- -D warnings
- name: Show sccache stats
if: always()
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
run: sccache --show-stats
clippy_data:
name: Clippy (Data)
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v5
- name: Install Rust toolchain
uses: dtolnay/rust-toolchain@stable
with:
toolchain: 1.95.0
components: clippy
- name: Rust cache
uses: Swatinem/rust-cache@v2
with:
shared-key: rust-ci-${{ runner.os }}
workspaces: . -> target
- name: Setup sccache
uses: mozilla-actions/[email protected]
- name: Clippy
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
run: cargo clippy -p aether-data --all-targets -- -D warnings
- name: Show sccache stats
if: always()
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
run: sccache --show-stats
clippy_rest:
name: Clippy (Workspace Rest)
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v5
- name: Install Rust toolchain
uses: dtolnay/rust-toolchain@stable
with:
toolchain: 1.95.0
components: clippy
- name: Rust cache
uses: Swatinem/rust-cache@v2
with:
shared-key: rust-ci-${{ runner.os }}
workspaces: . -> target
- name: Setup sccache
uses: mozilla-actions/[email protected]
- name: Clippy
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
run: cargo clippy --workspace --exclude aether-gateway --exclude aether-data --exclude aether-integration-tests --all-targets -- -D warnings
- name: Show sccache stats
if: always()
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
run: sccache --show-stats
clippy:
name: Clippy
runs-on: ubuntu-latest
needs:
- clippy_gateway
- clippy_data
- clippy_rest
if: ${{ always() }}
steps:
- name: Verify clippy jobs
run: |
if [ "${{ needs.clippy_gateway.result }}" != "success" ] || \
[ "${{ needs.clippy_data.result }}" != "success" ] || \
[ "${{ needs.clippy_rest.result }}" != "success" ]; then
echo "Clippy failed"
exit 1
fi
test_gateway:
name: Test (Gateway)
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v5
- name: Install Rust toolchain
uses: dtolnay/rust-toolchain@stable
- name: Show Rust toolchain
run: rustup show active-toolchain
- name: Rust cache
uses: Swatinem/rust-cache@v2
with:
shared-key: rust-ci-${{ runner.os }}
workspaces: . -> target
- name: Setup sccache
uses: mozilla-actions/[email protected]
- name: Setup mold
uses: rui314/setup-mold@v1
- name: Install nextest
uses: taiki-e/install-action@nextest
- name: Test lib
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
RUST_MIN_STACK: "16777216"
RUSTFLAGS: "-C link-arg=-fuse-ld=mold"
run: cargo nextest run -p aether-gateway --lib
- name: Test bin
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
RUST_MIN_STACK: "16777216"
RUSTFLAGS: "-C link-arg=-fuse-ld=mold"
run: cargo nextest run -p aether-gateway --bin aether-gateway
- name: Show sccache stats
if: always()
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
run: sccache --show-stats
test_data:
name: Test (Data)
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v5
- name: Install Rust toolchain
uses: dtolnay/rust-toolchain@stable
- name: Show Rust toolchain
run: rustup show active-toolchain
- name: Rust cache
uses: Swatinem/rust-cache@v2
with:
shared-key: rust-ci-${{ runner.os }}
workspaces: . -> target
- name: Setup sccache
uses: mozilla-actions/[email protected]
- name: Install nextest
uses: taiki-e/install-action@nextest
- name: Test
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
run: cargo nextest run -p aether-data
- name: Show sccache stats
if: always()
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
run: sccache --show-stats
check_data_features:
name: Check (Data Feature - ${{ matrix.feature }})
runs-on: ubuntu-latest
strategy:
fail-fast: false
matrix:
feature:
- postgres
- mysql
- sqlite
- all-drivers
steps:
- uses: actions/checkout@v5
- name: Install Rust toolchain
uses: dtolnay/rust-toolchain@stable
- name: Rust cache
uses: Swatinem/rust-cache@v2
with:
shared-key: rust-ci-${{ runner.os }}
workspaces: . -> target
- name: Setup sccache
uses: mozilla-actions/[email protected]
- name: Check selected data driver
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
run: cargo check -p aether-data --no-default-features --features ${{ matrix.feature }}
- name: Show sccache stats
if: always()
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
run: sccache --show-stats
test_rest:
name: Test (Workspace Rest)
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v5
- name: Install Rust toolchain
uses: dtolnay/rust-toolchain@stable
- name: Show Rust toolchain
run: rustup show active-toolchain
- name: Rust cache
uses: Swatinem/rust-cache@v2
with:
shared-key: rust-ci-${{ runner.os }}
workspaces: . -> target
- name: Setup sccache
uses: mozilla-actions/[email protected]
- name: Install nextest
uses: taiki-e/install-action@nextest
- name: Test
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
run: cargo nextest run --workspace --exclude aether-gateway --exclude aether-data --exclude aether-integration-tests
- name: Show sccache stats
if: always()
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
run: sccache --show-stats
test_data_adapters:
name: Test (Data Adapter - ${{ matrix.package }})
runs-on: ubuntu-latest
strategy:
fail-fast: false
matrix:
package:
- aether-data-postgres
- aether-data-mysql
- aether-data-sqlite
steps:
- uses: actions/checkout@v5
- name: Install Rust toolchain
uses: dtolnay/rust-toolchain@stable
- name: Rust cache
uses: Swatinem/rust-cache@v2
with:
shared-key: rust-ci-${{ runner.os }}
workspaces: . -> target
- name: Setup sccache
uses: mozilla-actions/[email protected]
- name: Install nextest
uses: taiki-e/install-action@nextest
- name: Test adapter
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
run: cargo nextest run -p ${{ matrix.package }}
- name: Show sccache stats
if: always()
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
run: sccache --show-stats
check_integration_scenarios:
name: Test (Integration Scenarios)
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v5
- name: Install Rust toolchain
uses: dtolnay/rust-toolchain@stable
- name: Rust cache
uses: Swatinem/rust-cache@v2
with:
shared-key: rust-ci-${{ runner.os }}
workspaces: . -> target
- name: Setup sccache
uses: mozilla-actions/[email protected]
- name: Test scenario binaries
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
run: cargo test -p aether-integration-tests --bins
- name: Show sccache stats
if: always()
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
run: sccache --show-stats
test:
name: Test
runs-on: ubuntu-latest
needs:
- test_gateway
- test_data
- check_data_features
- test_rest
- test_data_adapters
- check_integration_scenarios
if: ${{ always() }}
steps:
- name: Verify test jobs
run: |
if [ "${{ needs.test_gateway.result }}" != "success" ] || \
[ "${{ needs.test_data.result }}" != "success" ] || \
[ "${{ needs.check_data_features.result }}" != "success" ] || \
[ "${{ needs.test_rest.result }}" != "success" ] || \
[ "${{ needs.test_data_adapters.result }}" != "success" ] || \
[ "${{ needs.check_integration_scenarios.result }}" != "success" ]; then
echo "Tests failed"
exit 1
fi
data_db_smoke_sqlite:
name: Data DB Smoke (SQLite)
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v5
- name: Install Rust toolchain
uses: dtolnay/rust-toolchain@stable
- name: Show Rust toolchain
run: rustup show active-toolchain
- name: Rust cache
uses: Swatinem/rust-cache@v2
with:
shared-key: rust-ci-${{ runner.os }}
workspaces: . -> target
- name: Setup sccache
uses: mozilla-actions/[email protected]
- name: Run SQLite data smoke tests
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
run: cargo test -p aether-data --all-features sqlite --lib
- name: Show sccache stats
if: always()
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
run: sccache --show-stats
data_db_smoke_postgres:
name: Data DB Smoke (Postgres)
runs-on: ubuntu-latest
services:
postgres:
image: postgres:16
env:
POSTGRES_DB: aether_test
POSTGRES_USER: aether
POSTGRES_PASSWORD: aether
ports:
- 5432:5432
options: >-
--health-cmd="pg_isready -h 127.0.0.1 -U aether -d aether_test"
--health-interval=5s
--health-timeout=5s
--health-retries=20
steps:
- uses: actions/checkout@v5
- name: Install Rust toolchain
uses: dtolnay/rust-toolchain@stable
- name: Show Rust toolchain
run: rustup show active-toolchain
- name: Rust cache
uses: Swatinem/rust-cache@v2
with:
shared-key: rust-ci-${{ runner.os }}
workspaces: . -> target
- name: Setup sccache
uses: mozilla-actions/[email protected]
- name: Add PostgreSQL server binaries to PATH
run: echo "$(pg_config --bindir)" >> "$GITHUB_PATH"
- name: Run Postgres migration smoke test
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
AETHER_TEST_POSTGRES_URL: postgres://aether:[email protected]:5432/aether_test
run: cargo test -p aether-data --all-features postgres_migrations_create_core_config_tables_when_url_is_set --lib -- --nocapture
- name: Run Postgres provider metadata migration smoke test
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
AETHER_TEST_POSTGRES_URL: postgres://aether:[email protected]:5432/aether_test
run: cargo test -p aether-data --all-features postgres_provider_upstream_metadata_migration_preserves_json_when_url_is_set --lib -- --nocapture
- name: Run Postgres API key lifecycle tests
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
AETHER_REQUIRE_LOCAL_POSTGRES_TESTS: "true"
run: |
cargo test -p aether-data --all-features lifecycle::migrate::tests::postgres_request_candidates_preserve_deleted_api_key_identity --lib -- --exact --nocapture
cargo test -p aether-data --all-features lifecycle::migrate::tests::postgres_request_candidate_migration_decouples_legacy_api_key_foreign_key --lib -- --exact --nocapture
cargo test -p aether-data --all-features lifecycle::migrate::tests::postgres_stats_daily_api_key_migration_decouples_legacy_foreign_key --lib -- --exact --nocapture
cargo test -p aether-data --all-features lifecycle::migrate::tests::postgres_expired_api_key_cleanup_preserves_historical_identity --lib -- --exact --nocapture
cargo test -p aether-data --all-features lifecycle::migrate::tests::postgres_api_key_leaderboard_user_filter_preserves_aggregate_history --lib -- --exact --nocapture
- name: Run Postgres core export smoke test
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
AETHER_TEST_POSTGRES_URL: postgres://aether:[email protected]:5432/aether_test
run: cargo test -p aether-data --all-features postgres_core_export_reads_migrated_database_rows_when_url_is_set --lib -- --nocapture
- name: Run SQLite-to-Postgres import smoke test
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
AETHER_TEST_POSTGRES_URL: postgres://aether:[email protected]:5432/aether_test
run: cargo test -p aether-data --all-features sqlite_core_export_reads_migrated_database_rows --lib -- --nocapture
- name: Show sccache stats
if: always()
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
run: sccache --show-stats
data_db_smoke_mysql:
name: Data DB Smoke (MySQL)
runs-on: ubuntu-latest
services:
mysql:
image: mysql:8.0
env:
MYSQL_DATABASE: aether_test
MYSQL_USER: aether
MYSQL_PASSWORD: aether
MYSQL_ROOT_PASSWORD: aether_root
ports:
- 3306:3306
options: >-
--health-cmd="mysqladmin ping -h 127.0.0.1 -uaether -paether --silent"
--health-interval=5s
--health-timeout=5s
--health-retries=20
steps:
- uses: actions/checkout@v5
- name: Install Rust toolchain
uses: dtolnay/rust-toolchain@stable
- name: Show Rust toolchain
run: rustup show active-toolchain
- name: Rust cache
uses: Swatinem/rust-cache@v2
with:
shared-key: rust-ci-${{ runner.os }}
workspaces: . -> target
- name: Setup sccache
uses: mozilla-actions/[email protected]
- name: Run MySQL migration smoke test
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
AETHER_TEST_MYSQL_URL: mysql://aether:[email protected]:3306/aether_test
run: cargo test -p aether-data --all-features mysql_migrations_create_core_config_tables_when_url_is_set --lib -- --nocapture
- name: Run MySQL usage write smoke test
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
AETHER_TEST_MYSQL_URL: mysql://aether:[email protected]:3306/aether_test
run: cargo test -p aether-data-mysql mysql_usage_write_repository_upserts_when_url_is_set --lib -- --nocapture
- name: Run MySQL usage read smoke test
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
AETHER_TEST_MYSQL_URL: mysql://aether:[email protected]:3306/aether_test
run: cargo test -p aether-data-mysql mysql_usage_read_repository_reads_usage_contract_views_when_url_is_set --lib -- --nocapture
- name: Run MySQL provider catalog smoke test
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
AETHER_TEST_MYSQL_URL: mysql://aether:[email protected]:3306/aether_test
run: cargo test -p aether-data-mysql mysql_provider_catalog_repository_round_trips_when_url_is_set --lib -- --nocapture
- name: Run MySQL core export smoke test
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
AETHER_TEST_MYSQL_URL: mysql://aether:[email protected]:3306/aether_test
run: cargo test -p aether-data --all-features mysql_core_export_reads_migrated_database_rows_when_url_is_set --lib -- --nocapture
- name: Run MySQL wallet read smoke test
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
AETHER_TEST_MYSQL_URL: mysql://aether:[email protected]:3306/aether_test
run: cargo test -p aether-data-mysql mysql_wallet_read_repository_reads_wallet_contract_views --lib -- --nocapture
- name: Run MySQL wallet daily usage aggregation smoke test
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
AETHER_TEST_MYSQL_URL: mysql://aether:[email protected]:3306/aether_test
run: cargo test -p aether-data --all-features mysql_wallet_daily_usage_aggregation_uses_settlement_wallets_when_url_is_set --lib -- --nocapture
- name: Run MySQL stats aggregation smoke test
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
AETHER_TEST_MYSQL_URL: mysql://aether:[email protected]:3306/aether_test
run: cargo test -p aether-data --all-features mysql_stats_aggregation_runs_after_mysql_migrations_when_url_is_set --lib -- --nocapture
- name: Show sccache stats
if: always()
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
run: sccache --show-stats
data_db_smoke:
name: Data DB Smoke
runs-on: ubuntu-latest
needs:
- data_db_smoke_sqlite
- data_db_smoke_postgres
- data_db_smoke_mysql
if: ${{ always() }}
steps:
- name: Verify database smoke jobs
run: |
if [ "${{ needs.data_db_smoke_sqlite.result }}" != "success" ] || \
[ "${{ needs.data_db_smoke_postgres.result }}" != "success" ] || \
[ "${{ needs.data_db_smoke_mysql.result }}" != "success" ]; then
echo "Data DB smoke failed"
exit 1
fi
check:
name: check
runs-on: ubuntu-latest
needs:
- fmt
- clippy
- test
- data_db_smoke
if: ${{ always() }}
steps:
- name: Verify required jobs
run: |
if [ "${{ needs.fmt.result }}" != "success" ] || \
[ "${{ needs.clippy.result }}" != "success" ] || \
[ "${{ needs.test.result }}" != "success" ] || \
[ "${{ needs.data_db_smoke.result }}" != "success" ]; then
echo "Rust CI failed"
exit 1
fi
+2 -10
View File
@@ -1,17 +1,12 @@
# Created by https://www.toptal.com/developers/gitignore/api/python
# Edit at https://www.toptal.com/developers/gitignore?templates=python
*.rsa
*_rsa
# AI Assistant Configuration
.codex/
.claude/
.deepseek/
.serena/
.gemini*/
.plans
.playwright-mcp/
### Python ###
*.db
@@ -209,6 +204,7 @@ logs/
# Git backup
.git.backup/
.worktrees/
# Database backups
backups/
@@ -218,10 +214,6 @@ backups/
# Runtime lock files
.locks/
# Local Rust/Cargo configuration
.cargo/
# Demo and test files
frontend/public/*-demo.html
frontend/public/*-measure.html
@@ -247,4 +239,4 @@ src/_version.py
# Analysis folder (third-party code for reference)
analysis/
new-api/
apps/aether-tunnel/aether-tunnel.toml
/aether-proxy/target/
-2
View File
@@ -1,2 +0,0 @@
[tools]
rust = "latest"
+1
View File
@@ -0,0 +1 @@
3.13
Generated
-6256
View File
File diff suppressed because it is too large Load Diff
-156
View File
@@ -1,156 +0,0 @@
[workspace]
members = [
"apps/aether-tunnel",
"crates/aether-ai/formats",
"crates/aether-admin",
"crates/aether-admission-core",
"crates/aether-ai/serving",
"crates/aether-pool-core",
"crates/aether-provider/core",
"crates/aether-provider/pool",
"crates/aether-routing-core",
"crates/aether-data/contracts",
"crates/aether-data/adapters/postgres",
"crates/aether-data/adapters/mysql",
"crates/aether-data/adapters/sqlite",
"crates/aether-data/query",
"crates/aether-data/schema",
"crates/aether-dispatch-core",
"crates/aether-cache",
"crates/aether-billing",
"crates/aether-wallet",
"crates/aether-crypto",
"crates/aether-contracts",
"crates/aether-data/runtime",
"crates/aether-model-fetch",
"crates/aether-oauth",
"crates/aether-provider/transport",
"crates/aether-scheduler-core",
"crates/aether-runtime/state",
"crates/aether-task/runtime",
"crates/aether-task/core",
"crates/aether-gateway/frontdoor",
"crates/aether-gateway/control",
"crates/aether-gateway/execution",
"crates/aether-gateway/workers",
"crates/aether-gateway/tunnel",
"crates/aether-testing/loadtools",
"crates/aether-testing/integration",
"crates/aether-usage/core",
"crates/aether-testing/support",
"crates/aether-usage/runtime",
"crates/aether-video-tasks-core",
"apps/aether-gateway",
"crates/aether-http",
"crates/aether-runtime/base",
"crates/aether-testing/testkit",
]
default-members = [
"apps/aether-gateway",
]
resolver = "2"
[workspace.package]
edition = "2021"
license = "LicenseRef-Aether-NonCommercial"
repository = "https://github.com/fawney19/Aether.git"
[workspace.dependencies]
aether-admin = { path = "crates/aether-admin" }
aether-admission-core = { path = "crates/aether-admission-core" }
aether-ai-formats = { path = "crates/aether-ai/formats" }
aether-ai-serving = { path = "crates/aether-ai/serving" }
aether-pool-core = { path = "crates/aether-pool-core" }
aether-provider-core = { path = "crates/aether-provider/core" }
aether-provider-pool = { path = "crates/aether-provider/pool" }
aether-routing-core = { path = "crates/aether-routing-core" }
aether-data-contracts = { path = "crates/aether-data/contracts" }
aether-data-postgres = { path = "crates/aether-data/adapters/postgres" }
aether-data-mysql = { path = "crates/aether-data/adapters/mysql" }
aether-data-sqlite = { path = "crates/aether-data/adapters/sqlite" }
aether-data-query = { path = "crates/aether-data/query" }
aether-data-schema = { path = "crates/aether-data/schema" }
aether-dispatch-core = { path = "crates/aether-dispatch-core" }
aether-cache = { path = "crates/aether-cache" }
aether-billing = { path = "crates/aether-billing" }
aether-wallet = { path = "crates/aether-wallet" }
aether-crypto = { path = "crates/aether-crypto" }
aether-contracts = { path = "crates/aether-contracts" }
aether-data = { path = "crates/aether-data/runtime" }
aether-model-fetch = { path = "crates/aether-model-fetch" }
aether-oauth = { path = "crates/aether-oauth" }
aether-provider-transport = { path = "crates/aether-provider/transport" }
aether-scheduler-core = { path = "crates/aether-scheduler-core" }
aether-runtime-state = { path = "crates/aether-runtime/state" }
aether-task-runtime = { path = "crates/aether-task/runtime" }
aether-task-core = { path = "crates/aether-task/core" }
aether-gateway-frontdoor = { path = "crates/aether-gateway/frontdoor" }
aether-gateway-control = { path = "crates/aether-gateway/control" }
aether-gateway-execution = { path = "crates/aether-gateway/execution" }
aether-gateway-workers = { path = "crates/aether-gateway/workers" }
aether-gateway-tunnel = { path = "crates/aether-gateway/tunnel" }
aether-loadtools = { path = "crates/aether-testing/loadtools" }
aether-integration-tests = { path = "crates/aether-testing/integration" }
aether-test-support = { path = "crates/aether-testing/support" }
aether-usage-core = { path = "crates/aether-usage/core" }
aether-usage-runtime = { path = "crates/aether-usage/runtime" }
aether-video-tasks-core = { path = "crates/aether-video-tasks-core" }
aether-gateway = { path = "apps/aether-gateway" }
aether-http = { path = "crates/aether-http" }
aether-runtime = { path = "crates/aether-runtime/base" }
aether-testkit = { path = "crates/aether-testing/testkit" }
aes = "0.8"
aes-gcm = "0.10"
async-stream = "0.3"
async-trait = "0.1"
axum = "0.8"
base64 = "0.22"
bcrypt = "0.16"
bytes = "1"
cbc = "0.1"
chrono = { version = "0.4", features = ["serde"] }
chrono-tz = "0.10"
flate2 = "1"
futures-util = "0.3"
hmac = "0.12"
http = "1"
object_store = { version = "0.12", default-features = false, features = ["aws"] }
pbkdf2 = { version = "0.12", default-features = false, features = ["hmac"] }
reqwest = { version = "0.12", default-features = false, features = ["json", "stream", "rustls-tls", "http2", "socks"] }
redis = { version = "0.28", default-features = false, features = ["tokio-comp", "script", "streams", "connection-manager"] }
regex = "1"
rustls = { version = "0.23", features = ["ring"] }
semver = "1"
serde = { version = "1", features = ["derive"] }
serde_json = { version = "1", features = ["preserve_order"] }
serde_path_to_error = "0.1"
sha2 = "0.10"
socket2 = "0.6"
tar = "0.4"
sqlx = { version = "0.8", default-features = false, features = ["runtime-tokio-rustls", "chrono"] }
thiserror = "2"
tokio = { version = "1", features = ["macros", "net", "rt-multi-thread", "signal", "sync", "time"] }
tokio-util = { version = "0.7", features = ["codec", "io-util"] }
tracing = "0.1"
tracing-subscriber = { version = "0.3", features = ["env-filter", "json"] }
uuid = { version = "1", features = ["serde", "v4", "v5"] }
webpki-roots = "0.26"
wreq = { version = "6.0.0-rc.28", default-features = false, features = ["json", "stream", "socks", "webpki-roots", "ws"] }
wreq-util = "3.0.0-rc.10"
url = "2"
zstd = "0.13"
[profile.dev]
# Keep file/line information for backtraces while avoiding full debug info
# generation on very large crates during local development builds.
debug = "line-tables-only"
[profile.test]
# The gateway test target pulls in a very large in-crate test tree, so use the
# lighter debug format here as well to reduce rustc peak memory.
debug = "line-tables-only"
[profile.release]
lto = "thin"
strip = true
codegen-units = 8
+285 -37
View File
@@ -1,43 +1,291 @@
# syntax=docker/dockerfile:1
# Aether Gateway runtime image (cross-compilation)
# Binary and frontend assets are pre-built by CI; this Dockerfile only packages them.
# Usage: docker buildx build --platform linux/amd64,linux/arm64 -f Dockerfile.app .
#
# Build context must contain:
# dist/aether-gateway-amd64 (x86_64-unknown-linux-musl cross-compiled binary)
# dist/aether-gateway-arm64 (aarch64-unknown-linux-musl cross-compiled binary)
# dist/frontend/ (npm run build output)
# 运行镜像:从 base 提取产物到精简运行时
# 构建命令: docker build -f Dockerfile.app -t aether-app:latest .
# 用于 GitHub Actions CI(官方源)
# --- layout stage: create /opt/aether directory structure with symlink ---
# distroless has no shell, so we use busybox to set up the symlink.
FROM busybox:1.37-musl AS layout
FROM aether-base:latest AS builder
WORKDIR /app
# 复制前端源码并构建(CI 通过 no-cache-filters=builder 确保每次重建)
COPY frontend/ ./frontend/
RUN cd frontend && npm run build
# ==================== 运行时镜像 ====================
FROM python:3.13-slim
WORKDIR /app
ARG HUB_RELEASE_REPO=fawney19/Aether
ARG HUB_TAG
ARG TARGETARCH
ARG GITHUB_TOKEN
RUN mkdir -p /opt/aether/releases/image/bin /opt/aether/releases/image/frontend /opt/aether/logs
COPY dist/aether-gateway-${TARGETARCH} /opt/aether/releases/image/bin/aether-gateway
RUN chmod 0755 /opt/aether/releases/image/bin/aether-gateway
COPY dist/frontend/ /opt/aether/releases/image/frontend/
RUN ln -s /opt/aether/releases/image /opt/aether/current
# --- final stage: distroless runtime ---
FROM gcr.io/distroless/static-debian12
COPY --from=layout /opt/aether /opt/aether
WORKDIR /opt/aether
ENV RUST_LOG=aether_gateway=info \
APP_PORT=8084 \
AETHER_UPDATE_STRATEGY=docker \
AETHER_GATEWAY_STATIC_DIR=/opt/aether/current/frontend
EXPOSE 8084
# 运行时依赖(无 gcc/nodejs/npm,使用 BuildKit 缓存加速)
RUN --mount=type=cache,target=/var/cache/apt,sharing=locked \
--mount=type=cache,target=/var/lib/apt,sharing=locked \
apt-get update && apt-get install -y --no-install-recommends \
nginx \
supervisor \
libpq5 \
curl \
libjemalloc2
RUN set -eux; \
jemalloc_path="$(find /usr/lib -type f -name 'libjemalloc.so.2' | head -n1)"; \
[ -n "$jemalloc_path" ]; \
ln -sf "$jemalloc_path" /usr/local/lib/libjemalloc.so.2
# 从 base 镜像复制 Python 包
COPY --from=builder /usr/local/lib/python3.13/site-packages /usr/local/lib/python3.13/site-packages
# 只复制需要的 Python 可执行文件
COPY --from=builder /usr/local/bin/gunicorn /usr/local/bin/
COPY --from=builder /usr/local/bin/uvicorn /usr/local/bin/
COPY --from=builder /usr/local/bin/alembic /usr/local/bin/
# Hub 预编译二进制(构建时从 GitHub Release 下载)
# GITHUB_TOKEN 可选:未认证 API 限流 60 次/小时,认证后 5000 次/小时
RUN set -eux; \
auth_header=""; \
if [ -n "${GITHUB_TOKEN:-}" ]; then \
auth_header="Authorization: token ${GITHUB_TOKEN}"; \
fi; \
tag="${HUB_TAG:-}"; \
if [ -z "$tag" ]; then \
tag="$(curl -sL ${auth_header:+-H "$auth_header"} "https://api.github.com/repos/${HUB_RELEASE_REPO}/releases" | python3 -c "import json,sys;print(next((r['tag_name'] for r in json.load(sys.stdin) if r.get('tag_name','').startswith('hub-v') and not r.get('draft') and not r.get('prerelease')),''))")"; \
fi; \
if [ -z "$tag" ]; then \
echo "Failed to resolve hub release tag"; \
exit 1; \
fi; \
arch="${TARGETARCH:-}"; \
if [ -z "$arch" ]; then \
arch="$(dpkg --print-architecture)"; \
fi; \
case "$arch" in \
amd64|arm64) ;; \
x86_64) arch="amd64" ;; \
aarch64) arch="arm64" ;; \
*) echo "Unsupported architecture: $arch"; exit 1 ;; \
esac; \
echo "Using Hub release tag: $tag"; \
url="https://github.com/${HUB_RELEASE_REPO}/releases/download/${tag}/aether-hub-linux-${arch}.tar.gz"; \
curl -L --fail -o /tmp/aether-hub.tar.gz "$url"; \
tar xzf /tmp/aether-hub.tar.gz -C /usr/local/bin; \
chmod +x /usr/local/bin/aether-hub; \
rm -f /tmp/aether-hub.tar.gz
# 从 builder 阶段复制前端构建产物
COPY --from=builder /app/frontend/dist /usr/share/nginx/html
RUN chmod -R 755 /usr/share/nginx/html
# 复制后端代码
COPY src/ ./src/
COPY alembic.ini ./
COPY alembic/ ./alembic/
COPY gunicorn_conf.py ./
# Nginx 配置模板
# 策略:白名单后端路由 → 后端代理,其余全部 → 前端 SPA(index.html)
# 智能处理 IP:有外层代理头就透传,没有就用直连 IP
RUN printf '%s\n' \
'map $http_x_real_ip $real_ip {' \
' default $http_x_real_ip;' \
' "" $remote_addr;' \
'}' \
'' \
'map $http_x_forwarded_for $forwarded_for {' \
' default $http_x_forwarded_for;' \
' "" $remote_addr;' \
'}' \
'' \
'map $http_upgrade $connection_upgrade {' \
' default upgrade;' \
' "" "";' \
'}' \
'' \
'server {' \
' listen 80;' \
' server_name _;' \
' root /usr/share/nginx/html;' \
' index index.html;' \
' client_max_body_size 100M;' \
'' \
' # gzip 压缩配置(对 base64 图片等非流式响应有效)' \
' gzip on;' \
' gzip_min_length 256;' \
' gzip_comp_level 5;' \
' gzip_vary on;' \
' gzip_proxied any;' \
' gzip_types application/json text/plain text/css text/javascript application/javascript application/octet-stream;' \
' gzip_disable "msie6";' \
'' \
' # 静态资源:长期缓存' \
' location ~* \.(js|css|png|jpg|jpeg|gif|ico|svg|woff|woff2|ttf|eot)$ {' \
' expires 1y;' \
' add_header Cache-Control "public, no-transform";' \
' try_files $uri =404;' \
' }' \
'' \
' # 安全:阻止访问源码目录' \
' location ~ ^/(src|node_modules)/ {' \
' deny all;' \
' return 404;' \
' }' \
'' \
' # WebSocket 隧道端点(aether-proxy tunnel 模式)' \
' location = /api/internal/proxy-tunnel {' \
' proxy_pass http://127.0.0.1:8085/proxy;' \
' proxy_http_version 1.1;' \
' proxy_set_header Host $host;' \
' proxy_set_header X-Real-IP $real_ip;' \
' proxy_set_header X-Forwarded-For $forwarded_for;' \
' proxy_set_header X-Forwarded-Proto $scheme;' \
' proxy_set_header Upgrade $http_upgrade;' \
' proxy_set_header Connection "upgrade";' \
' # 剥离 CF 头,防止泄露给上游或返回给客户端' \
' proxy_hide_header CF-Connecting-IP;' \
' proxy_hide_header CF-IPCountry;' \
' proxy_hide_header CF-Ray;' \
' proxy_hide_header CF-Visitor;' \
' proxy_hide_header CDN-Loop;' \
' proxy_hide_header True-Client-IP;' \
' proxy_hide_header CF-Worker;' \
' proxy_hide_header CF-EW-Via;' \
' proxy_hide_header CF-Warp-Tag-ID;' \
' proxy_set_header CF-Connecting-IP "";' \
' proxy_set_header CF-IPCountry "";' \
' proxy_set_header CF-Ray "";' \
' proxy_set_header CF-Visitor "";' \
' proxy_set_header CDN-Loop "";' \
' proxy_set_header True-Client-IP "";' \
' proxy_set_header CF-Worker "";' \
' proxy_set_header CF-EW-Via "";' \
' proxy_set_header CF-Warp-Tag-ID "";' \
' proxy_read_timeout 86400s;' \
' proxy_send_timeout 86400s;' \
' }' \
'' \
' # 后端 API 路由(白名单)→ 代理到后端' \
' location ~ ^/(api|v1|v1beta|upload|health)(/|$) {' \
' proxy_pass http://127.0.0.1:PORT_PLACEHOLDER;' \
' proxy_http_version 1.1;' \
' proxy_set_header Host $host;' \
' proxy_set_header X-Real-IP $real_ip;' \
' proxy_set_header X-Forwarded-For $forwarded_for;' \
' proxy_set_header X-Forwarded-Proto $scheme;' \
' proxy_set_header Upgrade $http_upgrade;' \
' proxy_set_header Connection $connection_upgrade;' \
' proxy_set_header Accept $http_accept;' \
' proxy_set_header Content-Type $content_type;' \
' proxy_set_header Authorization $http_authorization;' \
' proxy_set_header X-Api-Key $http_x_api_key;' \
' # 剥离 CF 头,防止泄露给上游或返回给客户端' \
' proxy_hide_header CF-Connecting-IP;' \
' proxy_hide_header CF-IPCountry;' \
' proxy_hide_header CF-Ray;' \
' proxy_hide_header CF-Visitor;' \
' proxy_hide_header CDN-Loop;' \
' proxy_hide_header True-Client-IP;' \
' proxy_hide_header CF-Worker;' \
' proxy_hide_header CF-EW-Via;' \
' proxy_hide_header CF-Warp-Tag-ID;' \
' proxy_set_header CF-Connecting-IP "";' \
' proxy_set_header CF-IPCountry "";' \
' proxy_set_header CF-Ray "";' \
' proxy_set_header CF-Visitor "";' \
' proxy_set_header CDN-Loop "";' \
' proxy_set_header True-Client-IP "";' \
' proxy_set_header CF-Worker "";' \
' proxy_set_header CF-EW-Via "";' \
' proxy_set_header CF-Warp-Tag-ID "";' \
' proxy_buffering off;' \
' proxy_cache off;' \
' proxy_request_buffering off;' \
' chunked_transfer_encoding on;' \
' gzip off;' \
' add_header X-Accel-Buffering no;' \
' proxy_connect_timeout 60s;' \
' proxy_send_timeout 3600s;' \
' proxy_read_timeout 3600s;' \
' }' \
'' \
' # API 文档路由 → 代理到后端' \
' location ~ ^/(docs|redoc|openapi\\.json)$ {' \
' proxy_pass http://127.0.0.1:PORT_PLACEHOLDER;' \
' proxy_http_version 1.1;' \
' proxy_set_header Host $host;' \
' proxy_set_header X-Real-IP $real_ip;' \
' proxy_set_header X-Forwarded-For $forwarded_for;' \
' proxy_set_header X-Forwarded-Proto $scheme;' \
' # 剥离 CF 头,防止泄露给上游或返回给客户端' \
' proxy_hide_header CF-Connecting-IP;' \
' proxy_hide_header CF-IPCountry;' \
' proxy_hide_header CF-Ray;' \
' proxy_hide_header CF-Visitor;' \
' proxy_hide_header CDN-Loop;' \
' proxy_hide_header True-Client-IP;' \
' proxy_hide_header CF-Worker;' \
' proxy_hide_header CF-EW-Via;' \
' proxy_hide_header CF-Warp-Tag-ID;' \
' proxy_set_header CF-Connecting-IP "";' \
' proxy_set_header CF-IPCountry "";' \
' proxy_set_header CF-Ray "";' \
' proxy_set_header CF-Visitor "";' \
' proxy_set_header CDN-Loop "";' \
' proxy_set_header True-Client-IP "";' \
' proxy_set_header CF-Worker "";' \
' proxy_set_header CF-EW-Via "";' \
' proxy_set_header CF-Warp-Tag-ID "";' \
' }' \
'' \
' # 所有其他路由 → 前端 SPA(先尝试静态文件,再回退到 index.html)' \
' location / {' \
' try_files $uri $uri/ /index.html;' \
' }' \
'}' > /etc/nginx/sites-available/default.template
# Supervisor 配置
RUN printf '%s\n' \
'[supervisord]' \
'nodaemon=true' \
'logfile=/var/log/supervisor/supervisord.log' \
'pidfile=/var/run/supervisord.pid' \
'' \
'[program:nginx]' \
'command=/bin/bash -c "sed \"s/PORT_PLACEHOLDER/8084/g\" /etc/nginx/sites-available/default.template > /etc/nginx/sites-available/default && /usr/sbin/nginx -g \"daemon off;\""' \
'autostart=true' \
'autorestart=true' \
'stdout_logfile=/var/log/nginx/access.log' \
'stderr_logfile=/var/log/nginx/error.log' \
'' \
'[program:app]' \
'command=/bin/bash -c "MAX_REQUESTS_JITTER=$((${MAX_REQUESTS:-50000}/20)); exec gunicorn src.main:app -c gunicorn_conf.py --preload -w %(ENV_GUNICORN_WORKERS)s -k uvicorn.workers.UvicornWorker --bind 127.0.0.1:8084 --max-requests ${MAX_REQUESTS:-50000} --max-requests-jitter $MAX_REQUESTS_JITTER --access-logfile - --error-logfile - --log-level info"' \
'directory=/app' \
'autostart=true' \
'autorestart=true' \
'stdout_logfile=/dev/stdout' \
'stdout_logfile_maxbytes=0' \
'stderr_logfile=/dev/stderr' \
'stderr_logfile_maxbytes=0' \
'environment=PYTHONUNBUFFERED=1,PYTHONIOENCODING=utf-8,LANG=C.UTF-8,LC_ALL=C.UTF-8,DOCKER_CONTAINER=true,LD_PRELOAD=/usr/local/lib/libjemalloc.so.2,MALLOC_CONF="background_thread:true,dirty_decay_ms:5000,muzzy_decay_ms:5000"' \
'' \
'[program:tunnel-hub]' \
'command=/usr/local/bin/aether-hub --bind 0.0.0.0:8085' \
'autostart=true' \
'autorestart=true' \
'stdout_logfile=/dev/stdout' \
'stdout_logfile_maxbytes=0' \
'stderr_logfile=/dev/stderr' \
'stderr_logfile_maxbytes=0' > /etc/supervisor/conf.d/supervisord.conf
# 创建目录
RUN mkdir -p /var/log/supervisor /app/logs /app/data
# 入口脚本(启动前执行迁移)
COPY entrypoint.sh /entrypoint.sh
RUN chmod +x /entrypoint.sh
# 环境变量
ENV PYTHONUNBUFFERED=1 \
PYTHONDONTWRITEBYTECODE=1 \
PYTHONIOENCODING=utf-8 \
LANG=C.UTF-8 \
LC_ALL=C.UTF-8 \
LD_PRELOAD=/usr/local/lib/libjemalloc.so.2 \
MALLOC_CONF=background_thread:true,dirty_decay_ms:5000,muzzy_decay_ms:5000 \
PORT=8084 \
GUNICORN_WORKERS=2 \
MAX_REQUESTS=4000
EXPOSE 80
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
CMD ["/opt/aether/current/bin/aether-gateway", "--healthcheck"]
USER root
ENTRYPOINT ["/opt/aether/current/bin/aether-gateway"]
CMD curl -f http://localhost/health || exit 1
ENTRYPOINT ["/entrypoint.sh"]
CMD ["/usr/bin/supervisord", "-c", "/etc/supervisor/conf.d/supervisord.conf"]
+298 -129
View File
@@ -1,151 +1,320 @@
# syntax=docker.m.daocloud.io/docker/dockerfile:1
# Aether 运行镜像:Rust gateway 直接服务 API + 前端静态文件(国内镜像源版本)
# 构建命令: docker build --build-arg AETHER_BUILD_VERSION=v0.7.2 -f Dockerfile.app.local -t aether-app:latest .
# syntax=docker/dockerfile:1
# 运行镜像:从 base 提取产物到精简运行时(国内镜像源版本)
# 构建命令: docker build -f Dockerfile.app.local -t aether-app:latest .
# 用于本地/国内服务器部署
ARG RUST_VERSION=1.95.0
ARG NODE_BASE_IMAGE=docker.m.daocloud.io/library/node:22-slim
ARG RUST_BASE_IMAGE=docker.m.daocloud.io/library/rust:${RUST_VERSION}-slim
FROM aether-base:latest AS builder
# ==================== 前端构建 ====================
FROM ${NODE_BASE_IMAGE} AS frontend-builder
ARG AETHER_BUILD_VERSION
ENV AETHER_BUILD_VERSION=${AETHER_BUILD_VERSION} \
AETHER_VERSION=${AETHER_BUILD_VERSION}
WORKDIR /app/frontend
COPY frontend/package*.json ./
RUN --mount=type=cache,id=aether-npm-cache,target=/root/.npm,sharing=locked \
npm config set registry https://registry.npmmirror.com && \
npm ci --no-audit --no-fund
COPY frontend/ ./
RUN npm run build
WORKDIR /app
# ==================== Rust gateway 构建 ====================
FROM ${RUST_BASE_IMAGE} AS gateway-base
WORKDIR /build
# 复制前端源码并构建
COPY frontend/ ./frontend/
RUN cd frontend && npm run build
# 生产级 release 构建:保留 thin LTO,同时用 lld 缩短最终链接阶段。
ENV CARGO_REGISTRIES_CRATES_IO_PROTOCOL=sparse \
CARGO_PROFILE_RELEASE_LTO=thin \
CARGO_PROFILE_RELEASE_CODEGEN_UNITS=16 \
RUSTFLAGS="-C linker=clang -C link-arg=-fuse-ld=lld"
# ==================== 运行时镜像 ====================
FROM python:3.13-slim
WORKDIR /app
ARG HUB_RELEASE_REPO=fawney19/Aether
ARG HUB_TAG
ARG TARGETARCH
ARG GITHUB_TOKEN
# GitHub 下载镜像前缀,国内构建时传入可用的镜像加速地址
# 用法: --build-arg GITHUB_MIRROR=https://ghfast.top
# 或: --build-arg GITHUB_MIRROR=https://gh-proxy.com
# 或: --build-arg GITHUB_MIRROR=https://mirror.ghproxy.com
ARG GITHUB_MIRROR
# 运行时依赖(使用清华镜像源 + BuildKit 缓存加速)
RUN --mount=type=cache,target=/var/cache/apt,sharing=locked \
--mount=type=cache,target=/var/lib/apt,sharing=locked \
sed -i 's/deb.debian.org/mirrors.tuna.tsinghua.edu.cn/g' /etc/apt/sources.list.d/debian.sources && \
apt-get update && apt-get install -y --no-install-recommends \
build-essential \
ca-certificates \
clang \
cmake \
git \
libclang-dev \
libssl-dev \
lld \
pkg-config \
perl
RUN --mount=type=cache,id=aether-cargo-registry,target=/usr/local/cargo/registry,sharing=locked \
--mount=type=cache,id=aether-cargo-git,target=/usr/local/cargo/git,sharing=locked \
cargo install cargo-chef --locked
FROM gateway-base AS gateway-planner
COPY Cargo.toml Cargo.lock ./
COPY apps/ ./apps/
COPY crates/ ./crates/
RUN cargo chef prepare --recipe-path recipe.json
FROM gateway-base AS gateway-builder
ARG AETHER_BUILD_VERSION
ENV AETHER_BUILD_VERSION=${AETHER_BUILD_VERSION} \
AETHER_VERSION=${AETHER_BUILD_VERSION}
COPY --from=gateway-planner /build/recipe.json ./recipe.json
RUN --mount=type=cache,id=aether-cargo-registry,target=/usr/local/cargo/registry,sharing=locked \
--mount=type=cache,id=aether-cargo-git,target=/usr/local/cargo/git,sharing=locked \
--mount=type=cache,id=aether-cargo-target-local,target=/build/target,sharing=locked \
cargo chef cook --release --locked --package aether-gateway --bin aether-gateway --features jemalloc --recipe-path recipe.json
COPY Cargo.toml Cargo.lock ./
COPY apps/ ./apps/
COPY crates/ ./crates/
RUN --mount=type=cache,id=aether-cargo-registry,target=/usr/local/cargo/registry,sharing=locked \
--mount=type=cache,id=aether-cargo-git,target=/usr/local/cargo/git,sharing=locked \
--mount=type=cache,id=aether-cargo-target-local,target=/build/target,sharing=locked \
set -eux; \
cargo build --release --locked -p aether-gateway --bin aether-gateway --features jemalloc; \
cp target/release/aether-gateway /tmp/aether-gateway
# ==================== 最小运行时打包 ====================
FROM gateway-builder AS runtime-prep
nginx \
supervisor \
libpq5 \
curl \
libjemalloc2
RUN set -eux; \
mkdir -p \
/runtime-root/app/data \
/runtime-root/app/logs \
/runtime-root/etc \
/runtime-root/etc/ssl \
/runtime-root/lib \
/runtime-root/lib64 \
/runtime-root/usr/local/bin; \
cp /tmp/aether-gateway /runtime-root/usr/local/bin/aether-gateway; \
: > /tmp/runtime-libs.txt; \
: > /tmp/runtime-scan-queue.txt; \
printf '%s\n' /tmp/aether-gateway >> /tmp/runtime-scan-queue.txt; \
while [ -s /tmp/runtime-scan-queue.txt ]; do \
current="$(head -n1 /tmp/runtime-scan-queue.txt)"; \
sed -i '1d' /tmp/runtime-scan-queue.txt; \
ldd "$current" | awk '/=>/ { print $3 } $1 ~ /^\// { print $1 }' | while read -r lib; do \
[ -n "$lib" ]; \
if ! grep -Fxq "$lib" /tmp/runtime-libs.txt; then \
printf '%s\n' "$lib" >> /tmp/runtime-libs.txt; \
printf '%s\n' "$lib" >> /tmp/runtime-scan-queue.txt; \
jemalloc_path="$(find /usr/lib -type f -name 'libjemalloc.so.2' | head -n1)"; \
[ -n "$jemalloc_path" ]; \
ln -sf "$jemalloc_path" /usr/local/lib/libjemalloc.so.2
# 从 base 镜像复制 Python 包
COPY --from=builder /usr/local/lib/python3.13/site-packages /usr/local/lib/python3.13/site-packages
# 只复制需要的 Python 可执行文件
COPY --from=builder /usr/local/bin/gunicorn /usr/local/bin/
COPY --from=builder /usr/local/bin/uvicorn /usr/local/bin/
COPY --from=builder /usr/local/bin/alembic /usr/local/bin/
# Hub 预编译二进制
# 国内构建: --build-arg GITHUB_MIRROR=https://ghfast.top 即可走镜像下载
# GITHUB_TOKEN 可选:未认证 API 限流 60 次/小时,认证后 5000 次/小时
RUN set -eux; \
arch="${TARGETARCH:-}"; \
if [ -z "$arch" ]; then \
arch="$(dpkg --print-architecture)"; \
fi; \
done; \
done; \
sort -u /tmp/runtime-libs.txt -o /tmp/runtime-libs.txt; \
while read -r lib; do \
[ -n "$lib" ]; \
dest="/runtime-root$(dirname "$lib")"; \
mkdir -p "$dest"; \
cp -L "$lib" "$dest/"; \
done < /tmp/runtime-libs.txt; \
for lib in \
/lib/x86_64-linux-gnu/libnss_dns.so.2 \
/lib/x86_64-linux-gnu/libnss_files.so.2 \
/lib/x86_64-linux-gnu/libresolv.so.2; do \
if [ -f "$lib" ]; then \
dest="/runtime-root$(dirname "$lib")"; \
mkdir -p "$dest"; \
cp -L "$lib" "$dest/"; \
case "$arch" in \
amd64|arm64) ;; \
x86_64) arch="amd64" ;; \
aarch64) arch="arm64" ;; \
*) echo "Unsupported architecture: $arch"; exit 1 ;; \
esac; \
auth_header=""; \
if [ -n "${GITHUB_TOKEN:-}" ]; then \
auth_header="Authorization: token ${GITHUB_TOKEN}"; \
fi; \
done; \
cp -a /usr/lib/ssl /runtime-root/usr/lib/; \
cp -a /etc/ssl/certs /runtime-root/etc/ssl/; \
if [ -f /etc/ssl/openssl.cnf ]; then \
cp /etc/ssl/openssl.cnf /runtime-root/etc/ssl/openssl.cnf; \
tag="${HUB_TAG:-}"; \
if [ -z "$tag" ]; then \
tag="$(curl -sL ${auth_header:+-H "$auth_header"} "https://api.github.com/repos/${HUB_RELEASE_REPO}/releases" | python3 -c "import json,sys;print(next((r['tag_name'] for r in json.load(sys.stdin) if r.get('tag_name','').startswith('hub-v') and not r.get('draft') and not r.get('prerelease')),''))")"; \
fi; \
if [ -f /etc/nsswitch.conf ]; then \
cp /etc/nsswitch.conf /runtime-root/etc/nsswitch.conf; \
fi
if [ -z "$tag" ]; then \
echo "Failed to resolve hub release tag"; \
exit 1; \
fi; \
echo "Using Hub release tag: $tag"; \
origin_url="https://github.com/${HUB_RELEASE_REPO}/releases/download/${tag}/aether-hub-linux-${arch}.tar.gz"; \
if [ -n "${GITHUB_MIRROR:-}" ]; then \
url="${GITHUB_MIRROR}/https://github.com/${HUB_RELEASE_REPO}/releases/download/${tag}/aether-hub-linux-${arch}.tar.gz"; \
echo "Using mirror: ${GITHUB_MIRROR}"; \
else \
url="$origin_url"; \
fi; \
curl -L --fail -o /tmp/aether-hub.tar.gz "$url"; \
tar xzf /tmp/aether-hub.tar.gz -C /usr/local/bin; \
chmod +x /usr/local/bin/aether-hub; \
rm -f /tmp/aether-hub.tar.gz
# ==================== 运行时镜像 ====================
FROM scratch
# 从 builder 阶段复制前端构建产物
COPY --from=builder /app/frontend/dist /usr/share/nginx/html
RUN chmod -R 755 /usr/share/nginx/html
# 复制 gateway 二进制
COPY --from=runtime-prep /runtime-root/ /
# 复制后端代码
COPY src/ ./src/
COPY alembic.ini ./
COPY alembic/ ./alembic/
COPY gunicorn_conf.py ./
# 复制前端构建产物
COPY --from=frontend-builder /app/frontend/dist /srv/frontend
WORKDIR /app
# Nginx 配置模板
# 策略:白名单后端路由 → 后端代理,其余全部 → 前端 SPA(index.html)
# 智能处理 IP:有外层代理头就透传,没有就用直连 IP
RUN printf '%s\n' \
'map $http_x_real_ip $real_ip {' \
' default $http_x_real_ip;' \
' "" $remote_addr;' \
'}' \
'' \
'map $http_x_forwarded_for $forwarded_for {' \
' default $http_x_forwarded_for;' \
' "" $remote_addr;' \
'}' \
'' \
'map $http_upgrade $connection_upgrade {' \
' default upgrade;' \
' "" "";' \
'}' \
'' \
'server {' \
' listen 80;' \
' server_name _;' \
' root /usr/share/nginx/html;' \
' index index.html;' \
' client_max_body_size 100M;' \
'' \
' # gzip 压缩配置(对 base64 图片等非流式响应有效)' \
' gzip on;' \
' gzip_min_length 256;' \
' gzip_comp_level 5;' \
' gzip_vary on;' \
' gzip_proxied any;' \
' gzip_types application/json text/plain text/css text/javascript application/javascript application/octet-stream;' \
' gzip_disable "msie6";' \
'' \
' # 静态资源:长期缓存' \
' location ~* \.(js|css|png|jpg|jpeg|gif|ico|svg|woff|woff2|ttf|eot)$ {' \
' expires 1y;' \
' add_header Cache-Control "public, no-transform";' \
' try_files $uri =404;' \
' }' \
'' \
' # 安全:阻止访问源码目录' \
' location ~ ^/(src|node_modules)/ {' \
' deny all;' \
' return 404;' \
' }' \
'' \
' # WebSocket 隧道端点(aether-proxy tunnel 模式)' \
' location = /api/internal/proxy-tunnel {' \
' proxy_pass http://127.0.0.1:8085/proxy;' \
' proxy_http_version 1.1;' \
' proxy_set_header Host $host;' \
' proxy_set_header X-Real-IP $real_ip;' \
' proxy_set_header X-Forwarded-For $forwarded_for;' \
' proxy_set_header X-Forwarded-Proto $scheme;' \
' proxy_set_header Upgrade $http_upgrade;' \
' proxy_set_header Connection "upgrade";' \
' # 剥离 CF 头,防止泄露给上游或返回给客户端' \
' proxy_hide_header CF-Connecting-IP;' \
' proxy_hide_header CF-IPCountry;' \
' proxy_hide_header CF-Ray;' \
' proxy_hide_header CF-Visitor;' \
' proxy_hide_header CDN-Loop;' \
' proxy_hide_header True-Client-IP;' \
' proxy_hide_header CF-Worker;' \
' proxy_hide_header CF-EW-Via;' \
' proxy_hide_header CF-Warp-Tag-ID;' \
' proxy_set_header CF-Connecting-IP "";' \
' proxy_set_header CF-IPCountry "";' \
' proxy_set_header CF-Ray "";' \
' proxy_set_header CF-Visitor "";' \
' proxy_set_header CDN-Loop "";' \
' proxy_set_header True-Client-IP "";' \
' proxy_set_header CF-Worker "";' \
' proxy_set_header CF-EW-Via "";' \
' proxy_set_header CF-Warp-Tag-ID "";' \
' proxy_read_timeout 86400s;' \
' proxy_send_timeout 86400s;' \
' }' \
'' \
' # 后端 API 路由(白名单)→ 代理到后端' \
' location ~ ^/(api|v1|v1beta|upload|health)(/|$) {' \
' proxy_pass http://127.0.0.1:PORT_PLACEHOLDER;' \
' proxy_http_version 1.1;' \
' proxy_set_header Host $host;' \
' proxy_set_header X-Real-IP $real_ip;' \
' proxy_set_header X-Forwarded-For $forwarded_for;' \
' proxy_set_header X-Forwarded-Proto $scheme;' \
' proxy_set_header Upgrade $http_upgrade;' \
' proxy_set_header Connection $connection_upgrade;' \
' proxy_set_header Accept $http_accept;' \
' proxy_set_header Content-Type $content_type;' \
' proxy_set_header Authorization $http_authorization;' \
' proxy_set_header X-Api-Key $http_x_api_key;' \
' # 剥离 CF 头,防止泄露给上游或返回给客户端' \
' proxy_hide_header CF-Connecting-IP;' \
' proxy_hide_header CF-IPCountry;' \
' proxy_hide_header CF-Ray;' \
' proxy_hide_header CF-Visitor;' \
' proxy_hide_header CDN-Loop;' \
' proxy_hide_header True-Client-IP;' \
' proxy_hide_header CF-Worker;' \
' proxy_hide_header CF-EW-Via;' \
' proxy_hide_header CF-Warp-Tag-ID;' \
' proxy_set_header CF-Connecting-IP "";' \
' proxy_set_header CF-IPCountry "";' \
' proxy_set_header CF-Ray "";' \
' proxy_set_header CF-Visitor "";' \
' proxy_set_header CDN-Loop "";' \
' proxy_set_header True-Client-IP "";' \
' proxy_set_header CF-Worker "";' \
' proxy_set_header CF-EW-Via "";' \
' proxy_set_header CF-Warp-Tag-ID "";' \
' proxy_buffering off;' \
' proxy_cache off;' \
' proxy_request_buffering off;' \
' chunked_transfer_encoding on;' \
' gzip off;' \
' add_header X-Accel-Buffering no;' \
' proxy_connect_timeout 60s;' \
' proxy_send_timeout 3600s;' \
' proxy_read_timeout 3600s;' \
' }' \
'' \
' # API 文档路由 → 代理到后端' \
' location ~ ^/(docs|redoc|openapi\\.json)$ {' \
' proxy_pass http://127.0.0.1:PORT_PLACEHOLDER;' \
' proxy_http_version 1.1;' \
' proxy_set_header Host $host;' \
' proxy_set_header X-Real-IP $real_ip;' \
' proxy_set_header X-Forwarded-For $forwarded_for;' \
' proxy_set_header X-Forwarded-Proto $scheme;' \
' # 剥离 CF 头,防止泄露给上游或返回给客户端' \
' proxy_hide_header CF-Connecting-IP;' \
' proxy_hide_header CF-IPCountry;' \
' proxy_hide_header CF-Ray;' \
' proxy_hide_header CF-Visitor;' \
' proxy_hide_header CDN-Loop;' \
' proxy_hide_header True-Client-IP;' \
' proxy_hide_header CF-Worker;' \
' proxy_hide_header CF-EW-Via;' \
' proxy_hide_header CF-Warp-Tag-ID;' \
' proxy_set_header CF-Connecting-IP "";' \
' proxy_set_header CF-IPCountry "";' \
' proxy_set_header CF-Ray "";' \
' proxy_set_header CF-Visitor "";' \
' proxy_set_header CDN-Loop "";' \
' proxy_set_header True-Client-IP "";' \
' proxy_set_header CF-Worker "";' \
' proxy_set_header CF-EW-Via "";' \
' proxy_set_header CF-Warp-Tag-ID "";' \
' }' \
'' \
' # 所有其他路由 → 前端 SPA(先尝试静态文件,再回退到 index.html)' \
' location / {' \
' try_files $uri $uri/ /index.html;' \
' }' \
'}' > /etc/nginx/sites-available/default.template
ENV LANG=C.UTF-8 \
# Supervisor 配置
RUN printf '%s\n' \
'[supervisord]' \
'nodaemon=true' \
'logfile=/var/log/supervisor/supervisord.log' \
'pidfile=/var/run/supervisord.pid' \
'' \
'[program:nginx]' \
'command=/bin/bash -c "sed \"s/PORT_PLACEHOLDER/${PORT:-8084}/g\" /etc/nginx/sites-available/default.template > /etc/nginx/sites-available/default && /usr/sbin/nginx -g \"daemon off;\""' \
'autostart=true' \
'autorestart=true' \
'stdout_logfile=/var/log/nginx/access.log' \
'stderr_logfile=/var/log/nginx/error.log' \
'' \
'[program:app]' \
'command=/bin/bash -c "MAX_REQUESTS_JITTER=$((${MAX_REQUESTS:-50000}/20)); exec gunicorn src.main:app -c gunicorn_conf.py --preload -w %(ENV_GUNICORN_WORKERS)s -k uvicorn.workers.UvicornWorker --bind 0.0.0.0:%(ENV_PORT)s --max-requests ${MAX_REQUESTS:-50000} --max-requests-jitter $MAX_REQUESTS_JITTER --access-logfile - --error-logfile - --log-level info"' \
'directory=/app' \
'autostart=true' \
'autorestart=true' \
'stdout_logfile=/dev/stdout' \
'stdout_logfile_maxbytes=0' \
'stderr_logfile=/dev/stderr' \
'stderr_logfile_maxbytes=0' \
'environment=PYTHONUNBUFFERED=1,PYTHONIOENCODING=utf-8,LANG=C.UTF-8,LC_ALL=C.UTF-8,DOCKER_CONTAINER=true,LD_PRELOAD=/usr/local/lib/libjemalloc.so.2,MALLOC_CONF="background_thread:true,dirty_decay_ms:5000,muzzy_decay_ms:5000"' \
'' \
'[program:tunnel-hub]' \
'command=/usr/local/bin/aether-hub --bind 0.0.0.0:8085' \
'autostart=true' \
'autorestart=true' \
'stdout_logfile=/dev/stdout' \
'stdout_logfile_maxbytes=0' \
'stderr_logfile=/dev/stderr' \
'stderr_logfile_maxbytes=0' > /etc/supervisor/conf.d/supervisord.conf
# 创建目录
RUN mkdir -p /var/log/supervisor /app/logs /app/data
# 入口脚本(启动前执行迁移)
COPY entrypoint.sh /entrypoint.sh
RUN sed -i 's/\r$//' /entrypoint.sh && chmod +x /entrypoint.sh
# 环境变量
ENV PYTHONUNBUFFERED=1 \
PYTHONDONTWRITEBYTECODE=1 \
PYTHONIOENCODING=utf-8 \
LANG=C.UTF-8 \
LC_ALL=C.UTF-8 \
RUST_LOG=aether_gateway=info \
APP_PORT=8084 \
AETHER_UPDATE_STRATEGY=manual \
AETHER_GATEWAY_STATIC_DIR=/srv/frontend
LD_PRELOAD=/usr/local/lib/libjemalloc.so.2 \
MALLOC_CONF=background_thread:true,dirty_decay_ms:5000,muzzy_decay_ms:5000 \
PORT=8084 \
GUNICORN_WORKERS=2 \
MAX_REQUESTS=4000
EXPOSE 8084
EXPOSE 80
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
CMD ["/usr/local/bin/aether-gateway", "--healthcheck"]
CMD curl -f http://localhost/health || exit 1
ENTRYPOINT ["/usr/local/bin/aether-gateway"]
ENTRYPOINT ["/entrypoint.sh"]
CMD ["/usr/bin/supervisord", "-c", "/etc/supervisor/conf.d/supervisord.conf"]
-150
View File
@@ -1,150 +0,0 @@
# syntax=docker.m.daocloud.io/docker/dockerfile:1
# Aether 本地发布版联调镜像
# 作用:用当前源码构建一个 release-layout 容器,专门测试管理后台在线更新流程。
ARG RUST_VERSION=1.95.0
ARG NODE_BASE_IMAGE=docker.m.daocloud.io/library/node:22-slim
ARG RUST_BASE_IMAGE=docker.m.daocloud.io/library/rust:${RUST_VERSION}-slim
# ==================== 前端构建 ====================
FROM ${NODE_BASE_IMAGE} AS frontend-builder
ARG AETHER_BUILD_VERSION
ENV AETHER_BUILD_VERSION=${AETHER_BUILD_VERSION} \
AETHER_VERSION=${AETHER_BUILD_VERSION}
WORKDIR /app/frontend
COPY frontend/package*.json ./
RUN --mount=type=cache,id=aether-npm-cache,target=/root/.npm,sharing=locked \
npm config set registry https://registry.npmmirror.com && \
npm ci --no-audit --no-fund
COPY frontend/ ./
RUN npm run build
# ==================== Rust gateway 构建 ====================
FROM ${RUST_BASE_IMAGE} AS gateway-base
WORKDIR /build
ENV CARGO_REGISTRIES_CRATES_IO_PROTOCOL=sparse \
CARGO_PROFILE_RELEASE_LTO=thin \
CARGO_PROFILE_RELEASE_CODEGEN_UNITS=16
RUN --mount=type=cache,target=/var/cache/apt,sharing=locked \
--mount=type=cache,target=/var/lib/apt,sharing=locked \
sed -i 's/deb.debian.org/mirrors.tuna.tsinghua.edu.cn/g' /etc/apt/sources.list.d/debian.sources && \
apt-get update && apt-get install -y --no-install-recommends \
build-essential \
ca-certificates \
cmake \
git \
libclang-dev \
libssl-dev \
pkg-config \
perl
RUN --mount=type=cache,id=aether-cargo-registry,target=/usr/local/cargo/registry,sharing=locked \
--mount=type=cache,id=aether-cargo-git,target=/usr/local/cargo/git,sharing=locked \
cargo install cargo-chef --locked
FROM gateway-base AS gateway-planner
COPY Cargo.toml Cargo.lock ./
COPY apps/ ./apps/
COPY crates/ ./crates/
RUN cargo chef prepare --recipe-path recipe.json
FROM gateway-base AS gateway-builder
ARG AETHER_BUILD_VERSION
ARG AETHER_BUILD_TYPE=release
ENV AETHER_BUILD_VERSION=${AETHER_BUILD_VERSION} \
AETHER_VERSION=${AETHER_BUILD_VERSION} \
AETHER_BUILD_TYPE=${AETHER_BUILD_TYPE}
COPY --from=gateway-planner /build/recipe.json ./recipe.json
RUN --mount=type=cache,id=aether-cargo-registry,target=/usr/local/cargo/registry,sharing=locked \
--mount=type=cache,id=aether-cargo-git,target=/usr/local/cargo/git,sharing=locked \
--mount=type=cache,id=aether-cargo-target-release-local,target=/build/target,sharing=locked \
cargo chef cook --release --locked --package aether-gateway --bin aether-gateway --features jemalloc --recipe-path recipe.json
COPY Cargo.toml Cargo.lock ./
COPY apps/ ./apps/
COPY crates/ ./crates/
RUN --mount=type=cache,id=aether-cargo-registry,target=/usr/local/cargo/registry,sharing=locked \
--mount=type=cache,id=aether-cargo-git,target=/usr/local/cargo/git,sharing=locked \
--mount=type=cache,id=aether-cargo-target-release-local,target=/build/target,sharing=locked \
cargo build --release --locked -p aether-gateway --features jemalloc && \
cp target/release/aether-gateway /tmp/aether-gateway
# ==================== 最小运行时打包 ====================
FROM gateway-builder AS runtime-prep
RUN set -eux; \
mkdir -p \
/runtime-root/app/data \
/runtime-root/etc \
/runtime-root/etc/ssl \
/runtime-root/lib \
/runtime-root/lib64 \
/runtime-root/usr/lib \
/runtime-root/opt/aether/logs \
/runtime-root/opt/aether/releases/image/bin \
/runtime-root/opt/aether/releases/image/frontend; \
cp /tmp/aether-gateway /runtime-root/opt/aether/releases/image/bin/aether-gateway; \
ln -s /opt/aether/releases/image /runtime-root/opt/aether/current; \
: > /tmp/runtime-libs.txt; \
: > /tmp/runtime-scan-queue.txt; \
printf '%s\n' /tmp/aether-gateway >> /tmp/runtime-scan-queue.txt; \
while [ -s /tmp/runtime-scan-queue.txt ]; do \
current="$(head -n1 /tmp/runtime-scan-queue.txt)"; \
sed -i '1d' /tmp/runtime-scan-queue.txt; \
ldd "$current" | awk '/=>/ { print $3 } $1 ~ /^\// { print $1 }' | while read -r lib; do \
[ -n "$lib" ]; \
if ! grep -Fxq "$lib" /tmp/runtime-libs.txt; then \
printf '%s\n' "$lib" >> /tmp/runtime-libs.txt; \
printf '%s\n' "$lib" >> /tmp/runtime-scan-queue.txt; \
fi; \
done; \
done; \
sort -u /tmp/runtime-libs.txt -o /tmp/runtime-libs.txt; \
while read -r lib; do \
[ -n "$lib" ]; \
dest="/runtime-root$(dirname "$lib")"; \
mkdir -p "$dest"; \
cp -L "$lib" "$dest/"; \
done < /tmp/runtime-libs.txt; \
for lib in \
/lib/x86_64-linux-gnu/libnss_dns.so.2 \
/lib/x86_64-linux-gnu/libnss_files.so.2 \
/lib/x86_64-linux-gnu/libresolv.so.2; do \
if [ -f "$lib" ]; then \
dest="/runtime-root$(dirname "$lib")"; \
mkdir -p "$dest"; \
cp -L "$lib" "$dest/"; \
fi; \
done; \
cp -a /usr/lib/ssl /runtime-root/usr/lib/; \
cp -a /etc/ssl/certs /runtime-root/etc/ssl/; \
if [ -f /etc/ssl/openssl.cnf ]; then \
cp /etc/ssl/openssl.cnf /runtime-root/etc/ssl/openssl.cnf; \
fi; \
if [ -f /etc/nsswitch.conf ]; then \
cp /etc/nsswitch.conf /runtime-root/etc/nsswitch.conf; \
fi
COPY --from=frontend-builder /app/frontend/dist /runtime-root/opt/aether/releases/image/frontend
# ==================== 运行时镜像 ====================
FROM scratch
COPY --from=runtime-prep /runtime-root/ /
WORKDIR /app
ENV LANG=C.UTF-8 \
LC_ALL=C.UTF-8 \
RUST_LOG=aether_gateway=info \
APP_PORT=8084 \
AETHER_BASE_DIR=/opt/aether \
AETHER_UPDATE_STRATEGY=self \
AETHER_GATEWAY_STATIC_DIR=/opt/aether/current/frontend
EXPOSE 8084
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
CMD ["/opt/aether/current/bin/aether-gateway", "--healthcheck"]
ENTRYPOINT ["/opt/aether/current/bin/aether-gateway"]
+28
View File
@@ -0,0 +1,28 @@
# syntax=docker/dockerfile:1
# 构建镜像:编译环境 + 预编译的依赖
# 用于 GitHub Actions CI 构建(不使用国内镜像源)
# 构建命令: docker build -f Dockerfile.base -t aether-base:latest .
# 只在 pyproject.toml 或 frontend/package*.json 变化时需要重建
FROM python:3.13-slim
WORKDIR /app
# 构建工具(使用 BuildKit 缓存加速)
RUN --mount=type=cache,target=/var/cache/apt,sharing=locked \
--mount=type=cache,target=/var/lib/apt,sharing=locked \
apt-get update && apt-get install -y --no-install-recommends \
libpq-dev \
gcc \
nodejs \
npm
# Python 依赖(使用 BuildKit 缓存加速)
COPY pyproject.toml README.md ./
RUN --mount=type=cache,target=/root/.cache/pip \
mkdir -p src && touch src/__init__.py && \
SETUPTOOLS_SCM_PRETEND_VERSION=0.1.0 pip install .
# 前端依赖(只安装,不构建,使用 BuildKit 缓存加速)
COPY frontend/package*.json ./frontend/
RUN --mount=type=cache,target=/root/.npm \
cd frontend && npm ci
+31
View File
@@ -0,0 +1,31 @@
# syntax=docker/dockerfile:1
# 构建镜像:编译环境 + 预编译的依赖(国内镜像源版本)
# 构建命令: docker build -f Dockerfile.base.local -t aether-base:latest .
# 只在 pyproject.toml 或 frontend/package*.json 变化时需要重建
FROM python:3.13-slim
WORKDIR /app
# 构建工具(使用清华镜像源 + BuildKit 缓存加速)
RUN --mount=type=cache,target=/var/cache/apt,sharing=locked \
--mount=type=cache,target=/var/lib/apt,sharing=locked \
sed -i 's/deb.debian.org/mirrors.tuna.tsinghua.edu.cn/g' /etc/apt/sources.list.d/debian.sources && \
apt-get update && apt-get install -y --no-install-recommends \
libpq-dev \
gcc \
nodejs \
npm
# pip 镜像源
RUN pip config set global.index-url https://pypi.tuna.tsinghua.edu.cn/simple
# Python 依赖(使用 BuildKit 缓存加速)
COPY pyproject.toml README.md ./
RUN --mount=type=cache,target=/root/.cache/pip \
mkdir -p src && touch src/__init__.py && \
SETUPTOOLS_SCM_PRETEND_VERSION=0.1.0 pip install .
# 前端依赖(只安装,不构建,使用淘宝镜像源 + BuildKit 缓存加速)
COPY frontend/package*.json ./frontend/
RUN --mount=type=cache,target=/root/.npm \
cd frontend && npm config set registry https://registry.npmmirror.com && npm ci
-567
View File
@@ -1,567 +0,0 @@
SHELL := /bin/bash
DEV_RUST_LOG := info,executor::candidate_loop=debug,stream::execution=debug
ifeq ($(origin RUST_LOG), command line)
DEV_RUST_LOG := $(RUST_LOG)
endif
export DEV_RUST_LOG
.PHONY: dev dev-backend dev-frontend migration backfill
define DEV_BACKEND_SCRIPT
set -euo pipefail
if [ ! -f .env ]; then
echo "=> 未找到 .env,请先执行: cp .env.example .env"
exit 1
fi
set -a
source .env
set +a
dotenv_has_key() {
local key="$$1"
grep -Eq "^[[:space:]]*$${key}=" .env
}
lowercase() {
printf '%s' "$$1" | tr '[:upper:]' '[:lower:]'
}
dev_uses_sqlite_database() {
local driver
local url
driver="$$(lowercase "$${AETHER_DATABASE_DRIVER:-}")"
url="$${AETHER_DATABASE_URL:-$${DATABASE_URL:-}}"
[[ "$${driver}" == "sqlite" || "$${url}" == sqlite:* ]]
}
dev_uses_postgres_database() {
local driver
local url
driver="$$(lowercase "$${AETHER_DATABASE_DRIVER:-}")"
url="$${AETHER_DATABASE_URL:-$${DATABASE_URL:-}}"
if [[ -z "$${driver}" && -z "$${url}" ]]; then
return 0
fi
[[ "$${driver}" == "postgres" || "$${driver}" == "postgresql" || "$${url}" == postgres:* || "$${url}" == postgresql:* ]]
}
dev_uses_redis_runtime() {
local backend
backend="$$(lowercase "$${AETHER_RUNTIME_BACKEND:-}")"
if [[ "$${backend}" == "memory" ]]; then
return 1
fi
if [[ "$${backend}" == "redis" ]]; then
return 0
fi
if dev_uses_sqlite_database; then
return 1
fi
return 0
}
print_dev_infra_hint() {
echo "=> 本地开发依赖未就绪。"
echo "=> 可手动启动 Postgres / Redis:"
echo "=> docker compose up -d postgres redis"
}
check_postgres_ready() {
local host="$$1"
local port="$$2"
if command -v pg_isready >/dev/null 2>&1; then
pg_isready -h "$${host}" -p "$${port}" >/dev/null 2>&1
return $$?
fi
if command -v nc >/dev/null 2>&1; then
nc -z "$${host}" "$${port}" >/dev/null 2>&1
return $$?
fi
return 0
}
check_redis_ready() {
local host="$$1"
local port="$$2"
local password="$$3"
if command -v redis-cli >/dev/null 2>&1; then
REDISCLI_AUTH="$${password}" redis-cli -h "$${host}" -p "$${port}" ping >/dev/null 2>&1
return $$?
fi
if command -v nc >/dev/null 2>&1; then
nc -z "$${host}" "$${port}" >/dev/null 2>&1
return $$?
fi
return 0
}
is_local_host() {
case "$$1" in
localhost|127.0.0.1|::1)
return 0
;;
esac
return 1
}
ensure_dev_infra() {
local postgres_host="$${DB_HOST:-localhost}"
local postgres_port="$${DB_PORT:-5432}"
local redis_host="$${REDIS_HOST:-localhost}"
local redis_port="$${REDIS_PORT:-6379}"
local redis_password="$${REDIS_PASSWORD:-}"
local need_postgres=false
local need_redis=false
local services=()
if dev_uses_postgres_database; then
if ! check_postgres_ready "$${postgres_host}" "$${postgres_port}"; then
if is_local_host "$${postgres_host}"; then
need_postgres=true
services+=(postgres)
else
echo "=> PostgreSQL 不可用: $${postgres_host}:$${postgres_port}"
print_dev_infra_hint
return 1
fi
fi
fi
if dev_uses_redis_runtime; then
if ! check_redis_ready "$${redis_host}" "$${redis_port}" "$${redis_password}"; then
if is_local_host "$${redis_host}"; then
need_redis=true
services+=(redis)
else
echo "=> Redis 不可用: $${redis_host}:$${redis_port}"
print_dev_infra_hint
return 1
fi
fi
fi
if [ "$${#services[@]}" -eq 0 ]; then
return 0
fi
if ! command -v docker >/dev/null 2>&1; then
echo "=> 未找到 docker,无法自动启动本地开发依赖。"
print_dev_infra_hint
return 1
fi
echo "=> 本地开发依赖未就绪,正在启动: docker compose up -d $${services[*]}"
if ! docker compose up -d "$${services[@]}"; then
echo "=> docker compose 启动本地开发依赖失败。"
print_dev_infra_hint
return 1
fi
for _ in {1..100}; do
local ready=true
if [ "$${need_postgres}" = "true" ] && ! check_postgres_ready "$${postgres_host}" "$${postgres_port}"; then
ready=false
fi
if [ "$${need_redis}" = "true" ] && ! check_redis_ready "$${redis_host}" "$${redis_port}" "$${redis_password}"; then
ready=false
fi
if [ "$${ready}" = "true" ]; then
return 0
fi
sleep 0.2
done
if [ "$${need_postgres}" = "true" ] && ! check_postgres_ready "$${postgres_host}" "$${postgres_port}"; then
echo "=> PostgreSQL 不可用: $${postgres_host}:$${postgres_port}"
fi
if [ "$${need_redis}" = "true" ] && ! check_redis_ready "$${redis_host}" "$${redis_port}" "$${redis_password}"; then
echo "=> Redis 不可用: $${redis_host}:$${redis_port}"
fi
print_dev_infra_hint
return 1
}
print_startup_failure_hint() {
local log_file="$$1"
if [ -n "$${log_file}" ] && [ -f "$${log_file}" ]; then
if grep -Eq "database schema is behind" "$${log_file}"; then
echo "=> 检测到数据库 schema 落后,请执行: make migration"
return
fi
if grep -Eq "database backfills are behind" "$${log_file}"; then
echo "=> 检测到待执行 backfills,请执行: make backfill"
return
fi
fi
echo "=> 未识别到明确的修复动作,请根据上面的日志继续排查。"
}
wait_for_startup() {
local pid="$$1"
local timeout_seconds="$$2"
local service_name="$$3"
shift 3
STARTUP_WAIT_EARLY_EXIT=false
local attempts=$$((timeout_seconds * 10))
if [ "$${attempts}" -lt 1 ]; then
attempts=1
fi
for ((i = 0; i < attempts; i++)); do
if "$$@" >/dev/null 2>&1; then
return 0
fi
if ! kill -0 "$${pid}" >/dev/null 2>&1; then
STARTUP_WAIT_EARLY_EXIT=true
echo "=> $${service_name} 启动进程已提前退出,请检查上面的日志。"
print_startup_failure_hint "$${GATEWAY_LOG_FILE}"
return 1
fi
sleep 0.1
done
if "$$@" >/dev/null 2>&1; then
return 0
fi
if ! kill -0 "$${pid}" >/dev/null 2>&1; then
STARTUP_WAIT_EARLY_EXIT=true
echo "=> $${service_name} 启动进程已提前退出,请检查上面的日志。"
print_startup_failure_hint "$${GATEWAY_LOG_FILE}"
return 1
fi
echo "=> $${service_name} 在 $${timeout_seconds}s 内未通过启动检查。"
echo "=> 如果这是冷编译或存在并发 cargo 构建,可调大启动超时后重试。"
return 1
}
create_gateway_log_file() {
local tmp_root="$${TMPDIR:-/tmp}"
tmp_root="$${tmp_root%/}"
GATEWAY_LOG_DIR="$$(mktemp -d "$${tmp_root}/aether-dev-startup.XXXXXX")"
GATEWAY_LOG_FILE="$${GATEWAY_LOG_DIR}/gateway.log"
: > "$${GATEWAY_LOG_FILE}"
}
cleanup() {
local status="$${1:-0}"
trap - INT TERM EXIT
if [ -n "$${GATEWAY_PID:-}" ]; then
echo ""
echo "=> 停止 aether-gateway..."
kill "$${GATEWAY_PID}" >/dev/null 2>&1 || true
wait "$${GATEWAY_PID}" >/dev/null 2>&1 || true
fi
if [ -n "$${GATEWAY_LOG_FILE:-}" ] && [ -f "$${GATEWAY_LOG_FILE}" ]; then
rm -f "$${GATEWAY_LOG_FILE}"
fi
if [ -n "$${GATEWAY_LOG_DIR:-}" ] && [ -d "$${GATEWAY_LOG_DIR}" ]; then
rmdir "$${GATEWAY_LOG_DIR}" >/dev/null 2>&1 || true
fi
exit "$${status}"
}
trap 'cleanup 130' INT
trap 'cleanup 143' TERM
trap 'cleanup $$?' EXIT
export APP_PORT="$${APP_PORT:-8084}"
export RUST_LOG="$${DEV_RUST_LOG}"
RUST_SERVICE_STARTUP_TIMEOUT_SECONDS="$${RUST_SERVICE_STARTUP_TIMEOUT_SECONDS:-180}"
GATEWAY_STARTUP_TIMEOUT_SECONDS="$${GATEWAY_STARTUP_TIMEOUT_SECONDS:-$${RUST_SERVICE_STARTUP_TIMEOUT_SECONDS}}"
export AETHER_GATEWAY_VIDEO_TASK_TRUTH_SOURCE_MODE="$${AETHER_GATEWAY_VIDEO_TASK_TRUTH_SOURCE_MODE:-rust-authoritative}"
if dev_uses_postgres_database; then
export DATABASE_URL="postgresql://$${DB_USER:-postgres}:$${DB_PASSWORD:-}@$${DB_HOST:-localhost}:$${DB_PORT:-5432}/$${DB_NAME:-aether}"
if ! dotenv_has_key "AETHER_GATEWAY_DATA_POSTGRES_URL"; then
export AETHER_GATEWAY_DATA_POSTGRES_URL="$${DATABASE_URL}"
fi
fi
if dev_uses_redis_runtime; then
export REDIS_URL="redis://:$${REDIS_PASSWORD:-}@$${REDIS_HOST:-localhost}:$${REDIS_PORT:-6379}/0"
if ! dotenv_has_key "AETHER_GATEWAY_DATA_REDIS_URL"; then
export AETHER_GATEWAY_DATA_REDIS_URL="$${REDIS_URL}"
fi
else
unset REDIS_URL
unset AETHER_GATEWAY_DATA_REDIS_URL
fi
if ! dotenv_has_key "AETHER_GATEWAY_DATA_ENCRYPTION_KEY"; then
export AETHER_GATEWAY_DATA_ENCRYPTION_KEY="$${ENCRYPTION_KEY:-}"
fi
export DB_POOL_SIZE="$${DB_POOL_SIZE:-5}"
export DB_MAX_OVERFLOW="$${DB_MAX_OVERFLOW:-5}"
export HTTP_MAX_CONNECTIONS="$${HTTP_MAX_CONNECTIONS:-20}"
export HTTP_KEEPALIVE_CONNECTIONS="$${HTTP_KEEPALIVE_CONNECTIONS:-5}"
if ! command -v cargo >/dev/null 2>&1; then
echo "=> 未找到 cargo,无法启动 aether-gateway。请先安装 Rust toolchain。"
exit 1
fi
if ! command -v curl >/dev/null 2>&1; then
echo "=> 未找到 curl,无法检查 aether-gateway 健康状态。请先安装 curl。"
exit 1
fi
if [ -z "$${RUSTC_WRAPPER:-}" ] && command -v sccache >/dev/null 2>&1; then
export RUSTC_WRAPPER="$$(command -v sccache)"
echo "=> 启用 Rust 编译缓存: $${RUSTC_WRAPPER}"
fi
if ! ensure_dev_infra; then
exit 1
fi
GATEWAY_PID=""
GATEWAY_LOG_DIR=""
GATEWAY_LOG_FILE=""
STARTUP_WAIT_EARLY_EXIT=false
create_gateway_log_file
echo "=> 启动 aether-gateway (Rust frontdoor: 0.0.0.0:$${APP_PORT})..."
echo "=> 日志过滤: $${RUST_LOG}"
echo "=> 执行命令: cargo run -p aether-gateway -- --app-port $${APP_PORT}"
cargo run -p aether-gateway -- --app-port "$${APP_PORT}" > >(
tee -a "$${GATEWAY_LOG_FILE}"
) 2>&1 &
GATEWAY_PID=$$!
if ! wait_for_startup "$${GATEWAY_PID}" "$${GATEWAY_STARTUP_TIMEOUT_SECONDS}" "aether-gateway" curl -sf "http://127.0.0.1:$${APP_PORT}/_gateway/health"; then
if [ "$${STARTUP_WAIT_EARLY_EXIT}" = "true" ]; then
GATEWAY_PID=""
fi
exit 1
fi
if wait "$${GATEWAY_PID}"; then
gateway_exit_code=0
else
gateway_exit_code=$$?
fi
GATEWAY_PID=""
if [ "$${gateway_exit_code}" -ne 130 ] && [ "$${gateway_exit_code}" -ne 143 ]; then
echo "=> aether-gateway 运行失败并已退出,请检查上面的日志。"
print_startup_failure_hint "$${GATEWAY_LOG_FILE}"
fi
exit "$${gateway_exit_code}"
endef
export DEV_BACKEND_SCRIPT
define DEV_SCRIPT
set -euo pipefail
backend_pid=""
frontend_pid=""
cleanup() {
local status="$${1:-0}"
trap - INT TERM EXIT
if [ -n "$${backend_pid}" ] || [ -n "$${frontend_pid}" ]; then
echo ""
echo "=> 停止本地开发服务..."
if [ -n "$${backend_pid}" ]; then
kill "$${backend_pid}" >/dev/null 2>&1 || true
wait "$${backend_pid}" >/dev/null 2>&1 || true
fi
if [ -n "$${frontend_pid}" ]; then
kill "$${frontend_pid}" >/dev/null 2>&1 || true
wait "$${frontend_pid}" >/dev/null 2>&1 || true
fi
fi
exit "$${status}"
}
wait_for_backend_ready() {
while :; do
if curl -sf "http://127.0.0.1:$${APP_PORT}/_gateway/health" >/dev/null 2>&1; then
return 0
fi
if ! kill -0 "$${backend_pid}" >/dev/null 2>&1; then
if wait "$${backend_pid}"; then
status=0
else
status=$$?
fi
if [ "$${status}" -ne 0 ]; then
echo "=> 后端进程已退出 (status $${status})"
else
echo "=> 后端进程已退出"
fi
backend_pid=""
cleanup "$${status}"
fi
sleep 0.2
done
}
trap 'cleanup 130' INT
trap 'cleanup 143' TERM
trap 'cleanup $$?' EXIT
if [ -f .env ]; then
set -a
source .env
set +a
fi
export APP_PORT="$${APP_PORT:-8084}"
echo "=> 启动后端: RUST_LOG=$${DEV_RUST_LOG} cargo run -p aether-gateway -- --app-port $${APP_PORT:-8084}"
/bin/bash -euo pipefail -c "$$DEV_BACKEND_SCRIPT" &
backend_pid=$$!
echo "=> 等待后端健康检查: http://127.0.0.1:$${APP_PORT}/_gateway/health"
wait_for_backend_ready
echo "=> 启动前端: cd frontend && npm run dev"
( cd frontend && exec npm run dev ) &
frontend_pid=$$!
while :; do
if ! kill -0 "$${backend_pid}" >/dev/null 2>&1; then
if wait "$${backend_pid}"; then
status=0
else
status=$$?
fi
if [ "$${status}" -ne 0 ]; then
echo "=> 后端进程已退出 (status $${status})"
else
echo "=> 后端进程已退出"
fi
backend_pid=""
cleanup "$${status}"
fi
if ! kill -0 "$${frontend_pid}" >/dev/null 2>&1; then
if wait "$${frontend_pid}"; then
status=0
else
status=$$?
fi
if [ "$${status}" -ne 0 ]; then
echo "=> 前端进程已退出 (status $${status})"
else
echo "=> 前端进程已退出"
fi
frontend_pid=""
cleanup "$${status}"
fi
sleep 1
done
endef
export DEV_SCRIPT
define DB_TASK_SCRIPT
set -euo pipefail
if [ -z "$${DB_TASK_FLAG:-}" ] || [ -z "$${DB_TASK_LABEL:-}" ]; then
echo "=> 内部错误: DB_TASK_FLAG / DB_TASK_LABEL 未设置"
exit 1
fi
if [ ! -f .env ]; then
echo "=> 未找到 .env,请先执行: cp .env.example .env"
exit 1
fi
set -a
source .env
set +a
dotenv_has_key() {
local key="$$1"
grep -Eq "^[[:space:]]*$${key}=" .env
}
lowercase() {
printf '%s' "$$1" | tr '[:upper:]' '[:lower:]'
}
uses_postgres_database() {
local driver
local url
driver="$$(lowercase "$${AETHER_DATABASE_DRIVER:-}")"
url="$${AETHER_DATABASE_URL:-$${DATABASE_URL:-}}"
if [[ -z "$${driver}" && -z "$${url}" ]]; then
return 0
fi
[[ "$${driver}" == "postgres" || "$${driver}" == "postgresql" || "$${url}" == postgres:* || "$${url}" == postgresql:* ]]
}
if uses_postgres_database; then
export DATABASE_URL="postgresql://$${DB_USER:-postgres}:$${DB_PASSWORD:-}@$${DB_HOST:-localhost}:$${DB_PORT:-5432}/$${DB_NAME:-aether}"
if ! dotenv_has_key "AETHER_GATEWAY_DATA_POSTGRES_URL"; then
export AETHER_GATEWAY_DATA_POSTGRES_URL="$${DATABASE_URL}"
fi
fi
if ! dotenv_has_key "AETHER_GATEWAY_DATA_ENCRYPTION_KEY"; then
export AETHER_GATEWAY_DATA_ENCRYPTION_KEY="$${ENCRYPTION_KEY:-}"
fi
if ! command -v cargo >/dev/null 2>&1; then
echo "=> 未找到 cargo,无法执行 $${DB_TASK_LABEL}。请先安装 Rust toolchain。"
exit 1
fi
echo "=> 执行 $${DB_TASK_LABEL}: cargo run -p aether-gateway -- $${DB_TASK_FLAG}"
exec cargo run -p aether-gateway -- "$${DB_TASK_FLAG}"
endef
export DB_TASK_SCRIPT
dev:
@$(SHELL) -euo pipefail -c "$$DEV_SCRIPT"
dev-backend:
@$(SHELL) -euo pipefail -c "$$DEV_BACKEND_SCRIPT"
dev-frontend:
@cd frontend && npm run dev
migration:
@DB_TASK_FLAG=--migrate DB_TASK_LABEL="数据库迁移" $(SHELL) -euo pipefail -c "$$DB_TASK_SCRIPT"
backfill:
@DB_TASK_FLAG=--apply-backfills DB_TASK_LABEL="数据库 backfill" $(SHELL) -euo pipefail -c "$$DB_TASK_SCRIPT"
+100 -97
View File
@@ -11,7 +11,6 @@
<p align="center">
<a href="#简介">简介</a> •
<a href="#部署">部署</a> •
<a href="#api-文档">API 文档</a> •
<a href="#环境变量">环境变量</a> •
<a href="#qa">Q&A</a>
</p>
@@ -44,122 +43,126 @@ cd Aether
# 2. 配置环境变量
cp .env.example .env
# 生成 JWT_SECRET_KEY / ENCRYPTION_KEY, 并填入 .env
./generate_keys.sh
# 编辑 .env 设置 ADMIN_PASSWORD
python generate_keys.py # 生成密钥, 并将生成的密钥填入 .env
# 3. 首次部署 / 更新 (从以下部署形态任选其一)
# Postgres + Redis (适用于企业或多人使用)
# 3. 部署 / 更新(自动执行数据库迁移)
docker compose pull && docker compose up -d
# Single Node (适用于个人用户或朋友分享)
docker compose -f docker-compose.single-node.yml pull && docker compose -f docker-compose.single-node.yml up -d
# 4. 升级前备份 (可选)
docker compose exec postgres pg_dump -U postgres aether | gzip > backup_$(date +%Y%m%d_%H%M%S).sql.gz
```
### 一键更新
Docker Compose 部署后,可在部署目录直接执行:
### Docker Compose(本地构建镜像)
```bash
./update.sh
```
# 1. 克隆代码
git clone https://github.com/fawney19/Aether.git
cd Aether
`update.sh` 会拉取最新 `app` 镜像并重建 `app` 容器,Docker named volumes、`./data` 和 `./logs` 不会被删除。Single Node 部署也可显式指定:
# 2. 配置环境变量
cp .env.example .env
python generate_keys.py # 生成密钥, 并将生成的密钥填入 .env
```bash
./update.sh --mode single-node
```
仓库自带的 Docker Compose 默认把应用日志输出到容器 `stdout/stderr`,直接用 `docker compose logs -f app` 查看,并由 Docker 轮转日志,避免正式发布镜像切换到非 root 用户后再被宿主机挂载日志目录的权限问题拖垮启动。如果你确实需要文件日志,需要在 compose 里把 `AETHER_LOG_DESTINATION` 改成 `file|both`,并额外挂载一个容器用户可写的目录到 `/opt/aether/logs`。
管理后台右上角“版本信息”会检测新版本。Docker Compose 部署只提示版本,实际更新继续执行 `./update.sh`;systemd / launchd / 二进制部署才使用后台自更新,流程是下载对应平台的 GitHub Release 包、强制校验 `SHA256SUMS`、解压到 `/opt/aether/releases/<version>`,再切换 `/opt/aether/current` 并退出进程,交给 systemd / launchd 拉起新版本。
源码或本地构建版本不会启用后台在线更新,请继续使用源码更新流程。Docker Compose 用户如果希望“容器重建后也保持镜像层面的新版本”,仍建议定期运行 `./update.sh` 拉取并重建 app 镜像。服务器访问 GitHub 需要代理时,可设置 `AETHER_UPDATE_PROXY_URL`,也兼容 `UPDATE_PROXY_URL`、`HTTPS_PROXY`、`ALL_PROXY`、`HTTP_PROXY` 以及 `NO_PROXY`。共享出口触发 GitHub API 限流时,可设置只读 `AETHER_UPDATE_GITHUB_TOKEN`,也兼容 `GITHUB_TOKEN` / `GH_TOKEN`。下载总超时默认 600 秒,连续无响应/无数据默认 30 秒,可通过 `AETHER_UPDATE_DOWNLOAD_TIMEOUT_SECS` 和 `AETHER_UPDATE_DOWNLOAD_IDLE_TIMEOUT_SECS` 调整。
标准 Docker Compose 使用 Docker named volumes 存放 Postgres/Redis/MySQL 数据;Single Node 使用部署目录下的 `./data` 存放 SQLite 数据。
如果是本地源码构建镜像的部署,继续使用:
```bash
# 3. 部署 / 更新(自动构建、启动、迁移)
git pull
./deploy.sh
```
如果要在本机联调“管理后台在线更新”本身,可启动仓库内置的 release-layout 测试环境:
### 本地开发
```bash
docker compose -f docker-compose.release-local.yml up -d --build
# 启动依赖
docker compose -f docker-compose.build.yml up -d postgres redis
# 后端
uv sync
./dev.sh
# 前端
cd frontend && npm install && npm run dev
```
这套环境会用当前源码构建一个本地测试镜像,但编译为 `release` 类型,并默认伪装成 `v0.7.0`,这样后台会按正式发布版逻辑开放“立即更新”。默认监听 `http://127.0.0.1:18085`,数据目录使用 `./data-release-local`;日志默认走 `docker logs`,不会影响你正在跑的源码构建容器。
## Aether Proxy (可选)
如果这套容器在 `prepare-update` 时访问 GitHub 失败,而你本机是通过代理出网,请在 `.env` 里把 `AETHER_UPDATE_PROXY_URL` 写成宿主机地址,例如 `http://host.docker.internal:7890`;容器内的 `127.0.0.1` 指向容器自身,不是宿主机。
如果想重置这套联调环境(包括 `/opt/aether/current` 和已下载的历史版本),执行:
```bash
docker compose -f docker-compose.release-local.yml down -v
```
可选变量:
- `AETHER_RELEASE_LOCAL_VERSION`:本地联调镜像对外声明的当前版本,默认 `v0.7.0`
- `AETHER_RELEASE_LOCAL_PORT`:本地联调端口,默认 `18085`
- `LOCAL_RELEASE_APP_IMAGE`:本地联调镜像名,默认 `aether-app:release-local`
### 一键安装(默认 Single Node:Linux systemd / macOS launchd + SQLite)
```bash
git clone https://github.com/fawney19/Aether.git
cd Aether
curl -fsSL https://raw.githubusercontent.com/fawney19/Aether/main/install.sh | sudo bash
```
## 本地开发
依赖 Docker、Rust toolchain、Node.js 和 make。
```bash
make dev
```
`make dev` 会同时启动后端 `aether-gateway` 和前端 `frontend` 的 Vite dev server。需要单独启动时可使用 `make dev-backend` 或 `make dev-frontend`。
Postgres / Redis 本地依赖未就绪时,`make dev` 会自动执行 `docker compose up -d postgres redis`。
## Aether Tunnel (可选)
Aether Tunnel 是配套的正向代理节点,部署在海外 VPS 上,为墙内的 Aether 实例中转 API 流量。
Aether Proxy 是配套的正向代理节点,部署在海外 VPS 上,为墙内的 Aether 实例中转 API 流量。或者部署在其他服务器为指定的提供商、账号、Key使用不同的节点访问。支持 TUI 向导一键配置、systemd 服务管理、TLS 加密、DNS 缓存及连接池调优。
- Docker Compose 部署或下载预编译二进制直接运行
- 提供 macOS/Linux 与 Windows 一键脚本,自动下载最新 `tunnel-v*` 制品并向现有 `aether-tunnel.toml` 追加 `[[servers]]`
- 通过 `aether-tunnel setup` 完成交互式配置,自动注册为系统服务
- 详细文档见 [apps/aether-tunnel/README.md](apps/aether-tunnel/README.md)
## API 文档
- Embeddings: [OpenAI compatible `POST /v1/embeddings`](docs/api/embeddings.md)
- Rerank: [OpenAI/Jina compatible `POST /v1/rerank`](docs/api/rerank.md)
- 通过 `aether-proxy setup` 完成交互式配置,自动注册为系统服务
- 详细文档见 [aether-proxy/README.md](aether-proxy/README.md)
## 环境变量
- `APP_PORT`:`aether-gateway` 唯一监听端口,固定绑定 `0.0.0.0:${APP_PORT}`
- `DATABASE_URL`:数据库连接串;SQLite 例如 `sqlite:///opt/aether/data/aether.db`,Postgres 例如 `postgresql://postgres:aether@postgres:5432/aether`
- `AETHER_GATEWAY_DATA_POSTGRES_MIN_CONNECTIONS` / `AETHER_GATEWAY_DATA_POSTGRES_MAX_CONNECTIONS`:数据库连接池手动覆盖值;未配置时会自动推导,SQLite 固定 `1/1`,Postgres/MySQL 按 CPU 核心数计算并默认封顶 `100`
- `AETHER_GATEWAY_MAX_IN_FLIGHT_REQUESTS`:单实例请求并发上限;未配置时按 CPU 自动推导(基础范围 `512-65536`),低文件描述符预算时会进一步下调
- `AETHER_GATEWAY_REQUEST_BODY_BUFFER_BUDGET_MB`:单实例同时读取和解压请求体的加权内存预算,默认 `256MB`
- `AETHER_GATEWAY_REQUEST_BODY_READ_TIMEOUT_MS`:请求体完整读取超时,默认 `120000ms`
- `AETHER_MAX_REQUEST_BODY_MB`:单请求解压后的最大请求体,默认 `64MB`
- `AETHER_MAX_INTERNAL_BUFFERED_BODY_MB`:heartbeat、管理探测等内部必须整包读取的响应体上限,默认 `128MB`
- `AETHER_TUNNEL_NODE_STATUS_QUEUE_CAPACITY`:隧道节点状态上报队列容量,默认 `1024`;满载时拒绝新事件,避免控制面故障导致无界内存增长
- `AETHER_GATEWAY_SECURITY_CACHE_TTL_MS`:IP 黑白名单本地缓存时间,默认 `1000ms`,写操作会主动失效相关缓存
- `AETHER_MAX_REDACTED_SYNC_RESPONSE_BODY_MB`:启用 PII 恢复时同步响应允许缓冲的最大大小,默认 `64MB`
- `REDIS_URL`:Redis 连接串;仅 Postgres + Redis 的 Docker Compose 部署需要配置
- `AETHER_RUNTIME_BACKEND=memory|redis`:运行时缓存/协调后端。SQLite 默认用 `memory`,不会连接 Redis
- `AETHER_GATEWAY_AUTO_PREPARE_DATABASE`:常规启动前自动执行挂起的 schema migration 和 backfill;仓库自带的 `docker-compose.yml` 默认开启
- `JWT_SECRET_KEY` / `ENCRYPTION_KEY`:认证和敏感数据加密所需密钥
- `API_KEY_PREFIX`:用户和管理员新建 API Key 时使用的前缀,默认 `sk`
- `ADMIN_USERNAME` / `ADMIN_PASSWORD` / `ADMIN_EMAIL`:首次启动时自举首个本地管理员;`install.sh` 会提示输入管理员密码
- `CORS_ORIGINS` / `CORS_ALLOW_CREDENTIALS`:前端跨域来源控制;如果要跨域带登录 Cookie,`CORS_ORIGINS` 不能写 `*`
- `RUST_LOG`:Rust 日志过滤,例如 `aether_gateway=info`、`aether_gateway=debug,sqlx=warn`
- Docker Compose 的 `DB_PASSWORD` / `REDIS_PASSWORD` 默认使用 `aether`
### 必需配置
| 变量 | 说明 |
|------|------|
| `DB_PASSWORD` | PostgreSQL 数据库密码 |
| `REDIS_PASSWORD` | Redis 密码 |
| `JWT_SECRET_KEY` | JWT 签名密钥(使用 `generate_keys.py` 生成) |
| `ENCRYPTION_KEY` | API Key 加密密钥(更换后需重新配置 Provider Key) |
| `ADMIN_EMAIL` | 初始管理员邮箱 |
| `ADMIN_USERNAME` | 初始管理员用户名 |
| `ADMIN_PASSWORD` | 初始管理员密码 |
### 可选配置
| 变量 | 默认值 | 说明 |
|------|--------|------|
| `APP_PORT` | 8084 | 应用端口 |
| `API_KEY_PREFIX` | sk | API Key 前缀 |
| `LOG_LEVEL` | INFO | 日志级别 (DEBUG/INFO/WARNING/ERROR) |
| `GUNICORN_WORKERS` | 2 | Gunicorn 工作进程数 |
| `DB_PORT` | 5432 | PostgreSQL 端口 |
| `REDIS_PORT` | 6379 | Redis 端口 |
## Q&A
### Q: 如何开启/关闭请求体记录?
管理员在 **系统设置** 中配置日志记录的详细程度:
| 级别 | 记录内容 |
|------|----------|
| Base | 基本请求信息 |
| Headers | Base + 请求头 |
| Full | Headers + 请求体 |
### Q: 更新出问题如何回滚?
**有备份的情况(推荐):**
```bash
# 1. 停止应用
docker compose stop app
# 2. 恢复数据库(先清空再导入)
docker compose exec -T postgres psql -U postgres -c "DROP DATABASE aether; CREATE DATABASE aether;"
gunzip < backup_xxx.sql.gz | docker compose exec -T postgres psql -U postgres -d aether
# 3. 拉取旧版本镜像并重启
# 方式一:使用具体版本 tag(如果有发布版本号)
# 将 docker-compose.yml 中 image 从 ghcr.io/fawney19/aether:latest 改为指定版本
# 方式二:使用之前记录的镜像 digest
# 将 image 改为 ghcr.io/fawney19/aether@sha256:xxxxx
docker compose up -d app
```
> 可以在升级前通过 `docker inspect ghcr.io/fawney19/aether:latest --format '{{index .RepoDigests 0}}'` 记录当前镜像 digest,方便回滚时使用。
**没有备份的情况:**
```bash
# 1. 用当前容器回退数据库迁移(回退 1 步,按需调整数字)
docker compose exec app alembic downgrade -1
# 2. 查看回退后的版本确认正确
docker compose exec app alembic current
# 3. 切回旧镜像并重启(同上方式修改 docker-compose.yml 中的 image)
docker compose up -d app
```
> 注意:没有备份的回滚依赖 alembic downgrade,如果迁移涉及不可逆的数据变更(如删除列),可能无法完全恢复数据。因此强烈建议升级前备份。
---
@@ -177,4 +180,4 @@ Aether Tunnel 是配套的正向代理节点,部署在海外 VPS 上,为墙
## Star History
[![Star History Chart](https://api.star-history.com/svg?repos=fawney19/Aether&type=date&legend=top-left)](https://www.star-history.com/?repos=fawney19%2FAether&type=date&legend=top-left)
[![Star History Chart](https://api.star-history.com/svg?repos=fawney19/Aether&type=Date)](https://star-history.com/#fawney19/Aether&Date)
+3
View File
@@ -0,0 +1,3 @@
target/
.git/
.DS_Store
+2012
View File
File diff suppressed because it is too large Load Diff
+27
View File
@@ -0,0 +1,27 @@
[package]
name = "aether-hub"
version = "0.2.0"
edition = "2021"
description = "Tunnel Hub for Aether - frame router between workers and proxies"
[dependencies]
tokio = { version = "1", features = ["full"] }
axum = { version = "0.8", features = ["ws"] }
serde = { version = "1", features = ["derive"] }
serde_json = "1"
tracing = "0.1"
tracing-subscriber = { version = "0.3", features = ["env-filter"] }
clap = { version = "4", features = ["derive", "env"] }
dashmap = "6"
parking_lot = "0.12"
flate2 = "1"
futures-util = "0.3"
bytes = "1"
async-stream = "0.3"
http-body-util = "0.1"
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls"] }
[profile.release]
lto = true
strip = true
codegen-units = 1
+34
View File
@@ -0,0 +1,34 @@
# syntax=docker/dockerfile:1
FROM rust:1.85-slim AS builder
WORKDIR /build/aether-hub
# 可选:配置国内 Cargo 镜像源(本地构建时传 --build-arg CARGO_MIRROR=1)
ARG CARGO_MIRROR
RUN if [ -n "$CARGO_MIRROR" ]; then \
printf '[source.crates-io]\nreplace-with = "tuna"\n\n[source.tuna]\nregistry = "sparse+https://mirrors.tuna.tsinghua.edu.cn/crates.io-index/"\n' \
> /usr/local/cargo/config.toml; \
fi
# 先构建依赖层,最大化后续代码变更时的缓存命中
COPY Cargo.toml Cargo.lock ./
RUN mkdir src && printf 'fn main() {}\n' > src/main.rs
RUN --mount=type=cache,target=/usr/local/cargo/registry,sharing=locked \
--mount=type=cache,target=/build/aether-hub/target,sharing=locked \
cargo build --release --locked
RUN rm -rf src
COPY src ./src
RUN --mount=type=cache,target=/usr/local/cargo/registry,sharing=locked \
--mount=type=cache,target=/build/aether-hub/target,sharing=locked \
cargo build --release --locked && \
cp target/release/aether-hub /tmp/aether-hub
FROM debian:bookworm-slim
RUN apt-get update && apt-get install -y --no-install-recommends ca-certificates && \
rm -rf /var/lib/apt/lists/*
COPY --from=builder /tmp/aether-hub /usr/local/bin/aether-hub
EXPOSE 8085
ENTRYPOINT ["/usr/local/bin/aether-hub"]
CMD ["--bind", "0.0.0.0:8085"]
+38
View File
@@ -0,0 +1,38 @@
# aether-hub
`aether-hub` 是 Tunnel Hub 服务,负责在 proxy 与 worker 之间路由帧。
已集成在Docker镜像中, 无需单独部署。
## 部署端指定 Hub 版本并构建
```bash
cd /path/to/Aether
./deploy.sh --hub-tag hub-v0.1.0
```
不指定 `--hub-tag` 时,`./deploy.sh` 会自动解析最新 `hub-v*` release,并在构建 app 镜像时从 GitHub Release 下载对应架构的 Hub 二进制。
## build.sh 模式说明
- 默认是 `binary` 模式(`cross` 构建二进制)。
- `--upload <hub-vX.Y.Z>` 会把构建产物上传到 GitHub Release。
- 加 `--image` 后进入镜像模式(`docker buildx`,可选)。
常用参数:
- `--tag <tag>`: 镜像 tag
- `--image-name <name>`: 镜像名(默认 `ghcr.io/fawney19/aether-hub`)
- `--platforms <list>`: 例如 `linux/amd64,linux/arm64`
- `--push`: 推送镜像
- `--load`: 加载到本地 Docker(单平台)
- `--latest`: 额外打 `latest` tag
## 运行时参数
- `TUNNEL_HUB_WORKER_IDLE_TIMEOUT`:worker 心跳空闲超时,默认 `60` 秒
- `TUNNEL_HUB_OUTBOUND_QUEUE_CAPACITY`:单连接出站队列容量,默认 `128`;队列打满时会把连接视为拥塞并主动关闭,避免 Hub 内存无限增长
## 与部署脚本关系
- `./deploy.sh`: 本地构建部署(会本地构建 app/base,并在构建 app 时从 GitHub Release 下载 Hub,可用 `--hub-tag` 固定版本)。
+272
View File
@@ -0,0 +1,272 @@
#!/bin/bash
# aether-hub 构建脚本
#
# 支持两种模式:
# 1) binary 模式(默认): 构建多架构二进制并可上传 GitHub Release
# 2) image 模式: 构建并推送/加载 Docker 镜像(推荐生产发布用)
#
# 示例:
# # binary 模式(兼容旧行为)
# ./build.sh
# ./build.sh amd64
# ./build.sh --upload hub-v0.1.0
#
# # image 模式(多架构推送)
# ./build.sh --image --tag v0.2.5 --push --latest
# ./build.sh --image --tag sha-abc123 --image-name ghcr.io/fawney19/aether-hub --push
# ./build.sh --image --tag local-test --platforms linux/amd64 --load
set -euo pipefail
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)"
PROJECT_DIR="$(cd "$SCRIPT_DIR/.." && pwd)"
DIST_DIR="$SCRIPT_DIR/dist"
# -------------------------------
# Defaults
# -------------------------------
MODE="binary" # binary | image
# binary mode options
UPLOAD=false
UPLOAD_TAG=""
BINARY_TARGETS=""
# image mode options
IMAGE_NAME="${IMAGE_NAME:-ghcr.io/fawney19/aether-hub}"
IMAGE_TAG=""
IMAGE_PLATFORMS="linux/amd64,linux/arm64"
IMAGE_PUSH=false
IMAGE_LOAD=false
IMAGE_LATEST=false
usage() {
cat <<'EOF'
用法:
./build.sh [binary-args]
./build.sh --image [image-args]
binary 模式(默认):
amd64|arm64 仅构建指定架构(可重复)
--upload <hub-vX.Y.Z> 上传到 GitHub Release(需要 gh CLI)
image 模式:
--image 启用镜像模式
--tag <tag> 镜像 tag(默认自动从 git describe 推导)
--image-name <name> 镜像名(默认 ghcr.io/fawney19/aether-hub)
--platforms <list> 平台列表,逗号分隔(默认 linux/amd64,linux/arm64)
--push 推送镜像到仓库
--load 加载到本地 Docker(仅单平台)
--latest 额外打 latest tag
通用:
-h, --help 显示帮助
EOF
}
while [ $# -gt 0 ]; do
case "$1" in
--image)
MODE="image"
shift
;;
--tag)
IMAGE_TAG="${2:-}"
shift 2
;;
--image-name)
IMAGE_NAME="${2:-}"
shift 2
;;
--platforms)
IMAGE_PLATFORMS="${2:-}"
shift 2
;;
--push)
IMAGE_PUSH=true
shift
;;
--load)
IMAGE_LOAD=true
shift
;;
--latest)
IMAGE_LATEST=true
shift
;;
--upload)
UPLOAD=true
UPLOAD_TAG="${2:-}"
shift 2
;;
amd64|arm64)
BINARY_TARGETS="$BINARY_TARGETS $1"
shift
;;
-h|--help)
usage
exit 0
;;
*)
echo "❌ 未知参数: $1"
usage
exit 1
;;
esac
done
build_binary() {
if [ -z "$BINARY_TARGETS" ]; then
BINARY_TARGETS="amd64 arm64"
fi
if ! command -v cross >/dev/null 2>&1; then
echo "❌ 需要安装 cross: cargo install cross --git https://github.com/cross-rs/cross"
exit 1
fi
mkdir -p "$DIST_DIR"
echo "🔨 开始构建 aether-hub 二进制..."
echo " 目标平台: $BINARY_TARGETS"
echo ""
ARTIFACTS=""
for arch in $BINARY_TARGETS; do
case "$arch" in
amd64) target="x86_64-unknown-linux-gnu" ;;
arm64) target="aarch64-unknown-linux-gnu" ;;
*) echo "❌ 未知架构: $arch"; exit 1 ;;
esac
echo ">>> 构建 $arch ($target)..."
cd "$SCRIPT_DIR"
cross build --release --target "$target" --locked
BIN="target/$target/release/aether-hub"
if [ ! -f "$BIN" ]; then
echo "❌ 未找到二进制文件: $BIN"
exit 1
fi
ARCHIVE="$DIST_DIR/aether-hub-linux-$arch.tar.gz"
tar czf "$ARCHIVE" -C "target/$target/release" aether-hub
ARTIFACTS="$ARTIFACTS $ARCHIVE"
SIZE=$(du -h "$ARCHIVE" | cut -f1)
echo "✅ $arch 构建完成: $ARCHIVE ($SIZE)"
echo ""
done
cd "$DIST_DIR"
shasum -a 256 aether-hub-*.tar.gz > SHA256SUMS.txt
echo "📋 SHA256 校验和:"
cat SHA256SUMS.txt
echo ""
if [ "$UPLOAD" = true ]; then
if [ -z "$UPLOAD_TAG" ]; then
echo "❌ --upload 需要指定 tag,例如: ./build.sh --upload hub-v0.1.0"
exit 1
fi
if ! command -v gh >/dev/null 2>&1; then
echo "❌ 需要安装 GitHub CLI: brew install gh"
exit 1
fi
echo "📦 上传到 GitHub Release: $UPLOAD_TAG"
cd "$PROJECT_DIR"
if ! git rev-parse "$UPLOAD_TAG" >/dev/null 2>&1; then
git tag "$UPLOAD_TAG"
git push origin "$UPLOAD_TAG"
fi
gh release create "$UPLOAD_TAG" \
--title "aether-hub ${UPLOAD_TAG#hub-}" \
--generate-notes \
$ARTIFACTS \
"$DIST_DIR/SHA256SUMS.txt"
echo "✅ 上传完成!"
fi
echo "🎉 binary 模式完成!"
}
build_image() {
if ! command -v docker >/dev/null 2>&1; then
echo "❌ 未找到 docker,请先安装 Docker"
exit 1
fi
if ! docker buildx version >/dev/null 2>&1; then
echo "❌ 未找到 docker buildx,请先启用 buildx"
exit 1
fi
if [ "$IMAGE_PUSH" = true ] && [ "$IMAGE_LOAD" = true ]; then
echo "❌ --push 与 --load 不能同时使用"
exit 1
fi
if [ "$IMAGE_PUSH" = false ] && [ "$IMAGE_LOAD" = false ]; then
# image 模式默认走 push,符合发布场景
IMAGE_PUSH=true
fi
if [ -z "$IMAGE_TAG" ]; then
IMAGE_TAG=$(git -C "$PROJECT_DIR" describe --tags --always 2>/dev/null | sed 's/^v//')
if [ -z "$IMAGE_TAG" ]; then
IMAGE_TAG=$(date +%Y%m%d%H%M%S)
fi
fi
if [ "$IMAGE_LOAD" = true ] && [[ "$IMAGE_PLATFORMS" == *,* ]]; then
echo "❌ --load 仅支持单平台,请用 --platforms linux/amd64(或 arm64)"
exit 1
fi
local ref="${IMAGE_NAME}:${IMAGE_TAG}"
local cmd=(docker buildx build
--platform "$IMAGE_PLATFORMS"
-f "$SCRIPT_DIR/Dockerfile"
-t "$ref"
)
if [ "$IMAGE_LATEST" = true ]; then
cmd+=(-t "${IMAGE_NAME}:latest")
fi
if [ "$IMAGE_PUSH" = true ]; then
cmd+=(--push)
else
cmd+=(--load)
fi
cmd+=("$SCRIPT_DIR")
echo "🔨 开始构建 aether-hub 镜像..."
echo " image: $ref"
echo " platforms: $IMAGE_PLATFORMS"
echo " mode: $([ "$IMAGE_PUSH" = true ] && echo push || echo load)"
echo ""
"${cmd[@]}"
if [ "$IMAGE_PUSH" = true ]; then
echo "✅ 镜像已推送: $ref"
if [ "$IMAGE_LATEST" = true ]; then
echo "✅ 镜像已推送: ${IMAGE_NAME}:latest"
fi
else
echo "✅ 镜像已加载到本地: $ref"
fi
echo "🎉 image 模式完成!"
}
if [ "$MODE" = "image" ]; then
build_image
else
build_binary
fi
+85
View File
@@ -0,0 +1,85 @@
use reqwest::Client;
#[derive(Clone)]
pub struct ControlPlaneClient {
client: Option<Client>,
base_url: String,
}
impl ControlPlaneClient {
pub fn new(base_url: String) -> Self {
let client = Client::builder()
.timeout(std::time::Duration::from_secs(10))
.build()
.ok();
Self { client, base_url }
}
pub fn disabled() -> Self {
Self {
client: None,
base_url: String::new(),
}
}
pub async fn heartbeat_ack(&self, payload: &[u8]) -> Result<Vec<u8>, String> {
let Some(client) = &self.client else {
return Ok(b"{}".to_vec());
};
let url = format!(
"{}/api/internal/hub/heartbeat",
self.base_url.trim_end_matches('/')
);
let response = client
.post(&url)
.header("content-type", "application/json")
.body(payload.to_vec())
.send()
.await
.map_err(|e| format!("heartbeat callback request failed: {e}"))?;
if !response.status().is_success() {
return Err(format!(
"heartbeat callback failed with status {}",
response.status()
));
}
response
.bytes()
.await
.map(|bytes| bytes.to_vec())
.map_err(|e| format!("heartbeat callback body read failed: {e}"))
}
pub async fn push_node_status(
&self,
node_id: &str,
connected: bool,
conn_count: usize,
) -> Result<(), String> {
let Some(client) = &self.client else {
return Ok(());
};
let url = format!(
"{}/api/internal/hub/node-status",
self.base_url.trim_end_matches('/')
);
let response = client
.post(&url)
.json(&serde_json::json!({
"node_id": node_id,
"connected": connected,
"conn_count": conn_count,
}))
.send()
.await
.map_err(|e| format!("node-status callback request failed: {e}"))?;
if response.status().is_success() {
Ok(())
} else {
Err(format!(
"node-status callback failed with status {}",
response.status()
))
}
}
}
+878
View File
@@ -0,0 +1,878 @@
use std::collections::HashMap;
use std::sync::atomic::{AtomicBool, AtomicU32, AtomicU64, AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::Duration;
use axum::extract::ws::Message;
use bytes::Bytes;
use dashmap::DashMap;
use parking_lot::{Mutex, RwLock};
use tokio::sync::mpsc;
use tokio::sync::mpsc::error::TrySendError;
use tokio::sync::{watch, Notify};
use tracing::{debug, info, warn};
use crate::control_plane::ControlPlaneClient;
use crate::protocol;
const MAX_REQUEST_BODY_FRAME_SIZE: usize = 32 * 1024;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SendStatus {
Queued,
Closed,
Congested,
}
#[derive(Debug, Clone, Copy)]
pub struct ConnConfig {
pub ping_interval: Duration,
pub idle_timeout: Duration,
pub outbound_queue_capacity: usize,
}
pub struct BoundedOutbound {
tx: mpsc::Sender<Message>,
close_tx: watch::Sender<bool>,
closing: AtomicBool,
}
impl BoundedOutbound {
pub fn new(tx: mpsc::Sender<Message>, close_tx: watch::Sender<bool>) -> Self {
Self {
tx,
close_tx,
closing: AtomicBool::new(false),
}
}
pub fn send(&self, msg: Message) -> SendStatus {
if self.is_closing() {
return SendStatus::Closed;
}
match self.tx.try_send(msg) {
Ok(()) => SendStatus::Queued,
Err(TrySendError::Closed(_)) => {
self.mark_closing();
SendStatus::Closed
}
Err(TrySendError::Full(_)) => {
self.mark_closing();
SendStatus::Congested
}
}
}
pub fn is_closing(&self) -> bool {
self.closing.load(Ordering::Acquire)
}
pub fn mark_closing(&self) -> bool {
if self.closing.swap(true, Ordering::AcqRel) {
return false;
}
let _ = self.close_tx.send(true);
true
}
}
pub struct ProxyConn {
pub id: u64,
pub node_id: String,
pub node_name: String,
pub outbound: BoundedOutbound,
next_stream_id: AtomicU32,
pub stream_count: AtomicUsize,
pub max_streams: usize,
}
impl ProxyConn {
pub fn new(
id: u64,
node_id: String,
node_name: String,
tx: mpsc::Sender<Message>,
close_tx: watch::Sender<bool>,
max_streams: usize,
) -> Self {
Self {
id,
node_id,
node_name,
outbound: BoundedOutbound::new(tx, close_tx),
next_stream_id: AtomicU32::new(2),
stream_count: AtomicUsize::new(0),
max_streams,
}
}
pub fn alloc_stream_id(&self) -> Option<u32> {
let mut current = self.stream_count.load(Ordering::Relaxed);
loop {
if current >= self.max_streams || !self.is_available() {
return None;
}
match self.stream_count.compare_exchange_weak(
current,
current + 1,
Ordering::AcqRel,
Ordering::Relaxed,
) {
Ok(_) => break,
Err(observed) => current = observed,
}
}
let sid = loop {
let current_sid = self.next_stream_id.load(Ordering::Relaxed);
let next_sid = if current_sid >= 0xFFFF_FFFE {
2
} else {
current_sid + 2
};
if self
.next_stream_id
.compare_exchange_weak(current_sid, next_sid, Ordering::AcqRel, Ordering::Relaxed)
.is_ok()
{
break current_sid;
}
};
Some(sid)
}
pub fn release_stream(&self) {
let mut current = self.stream_count.load(Ordering::Relaxed);
while current > 0 {
match self.stream_count.compare_exchange_weak(
current,
current - 1,
Ordering::AcqRel,
Ordering::Relaxed,
) {
Ok(_) => return,
Err(observed) => current = observed,
}
}
}
pub fn is_available(&self) -> bool {
!self.outbound.is_closing()
}
pub fn request_close(&self) {
self.outbound.mark_closing();
}
pub fn send(&self, msg: Message) -> SendStatus {
let was_closing = self.outbound.is_closing();
let status = self.outbound.send(msg);
if status == SendStatus::Congested && !was_closing {
warn!(
conn_id = self.id,
node_id = %self.node_id,
node_name = %self.node_name,
queued_streams = self.stream_count.load(Ordering::Relaxed),
"proxy outbound queue full, closing congested connection"
);
}
status
}
}
#[derive(Debug, Clone)]
pub struct LocalResponseHead {
pub status: u16,
pub headers: Vec<(String, String)>,
}
#[derive(Debug)]
pub enum LocalBodyEvent {
Chunk(Bytes),
End,
Error(String),
}
#[derive(Debug, Default)]
struct LocalWaitState {
response: Option<LocalResponseHead>,
error: Option<String>,
}
pub struct LocalStream {
pub id: u64,
proxy_conn_id: u64,
proxy_stream_id: u32,
wait_state: Mutex<LocalWaitState>,
headers_notify: Notify,
body_tx: mpsc::Sender<LocalBodyEvent>,
body_rx: Mutex<Option<mpsc::Receiver<LocalBodyEvent>>>,
terminal: AtomicBool,
}
impl LocalStream {
fn new(id: u64, proxy_conn_id: u64, proxy_stream_id: u32) -> Self {
let (body_tx, body_rx) = mpsc::channel(128);
Self {
id,
proxy_conn_id,
proxy_stream_id,
wait_state: Mutex::new(LocalWaitState::default()),
headers_notify: Notify::new(),
body_tx,
body_rx: Mutex::new(Some(body_rx)),
terminal: AtomicBool::new(false),
}
}
pub async fn wait_headers(&self, timeout: Duration) -> Result<LocalResponseHead, String> {
tokio::time::timeout(timeout, async {
loop {
let outcome = {
let state = self.wait_state.lock();
if let Some(response) = &state.response {
return Ok(response.clone());
}
state.error.clone()
};
if let Some(error) = outcome {
return Err(error);
}
self.headers_notify.notified().await;
}
})
.await
.map_err(|_| "timed out waiting for response headers".to_string())?
}
pub fn take_body_receiver(&self) -> Option<mpsc::Receiver<LocalBodyEvent>> {
self.body_rx.lock().take()
}
fn set_response_headers(&self, meta: protocol::ResponseMeta) {
let mut notify = false;
{
let mut state = self.wait_state.lock();
if state.response.is_none() && state.error.is_none() {
state.response = Some(LocalResponseHead {
status: meta.status,
headers: meta.headers,
});
notify = true;
}
}
if notify {
self.headers_notify.notify_waiters();
}
}
fn push_body_chunk(&self, payload: Bytes) -> bool {
if self.terminal.load(Ordering::Acquire) {
return false;
}
self.body_tx
.try_send(LocalBodyEvent::Chunk(payload))
.is_ok()
}
fn finish(&self) {
if self.terminal.swap(true, Ordering::AcqRel) {
return;
}
let mut notify = false;
{
let mut state = self.wait_state.lock();
if state.response.is_none() && state.error.is_none() {
state.error = Some("stream ended before response headers".to_string());
notify = true;
}
}
if notify {
self.headers_notify.notify_waiters();
}
let _ = self.body_tx.try_send(LocalBodyEvent::End);
}
fn fail(&self, error: impl Into<String>) {
if self.terminal.swap(true, Ordering::AcqRel) {
return;
}
let error = error.into();
let mut notify = false;
{
let mut state = self.wait_state.lock();
if state.response.is_none() && state.error.is_none() {
state.error = Some(error.clone());
notify = true;
}
}
if notify {
self.headers_notify.notify_waiters();
}
let _ = self.body_tx.try_send(LocalBodyEvent::Error(error));
}
}
pub struct HubRouter {
proxy_conns: RwLock<HashMap<String, Vec<Arc<ProxyConn>>>>,
proxy_conns_by_id: DashMap<u64, Arc<ProxyConn>>,
local_streams: DashMap<u64, Arc<LocalStream>>,
proxy_to_local: DashMap<(u64, u32), u64>,
next_conn_id: AtomicU64,
next_local_stream_id: AtomicU64,
control_plane: ControlPlaneClient,
}
impl HubRouter {
pub fn new(control_plane: ControlPlaneClient) -> Arc<Self> {
Arc::new(Self {
proxy_conns: RwLock::new(HashMap::new()),
proxy_conns_by_id: DashMap::new(),
local_streams: DashMap::new(),
proxy_to_local: DashMap::new(),
next_conn_id: AtomicU64::new(1),
next_local_stream_id: AtomicU64::new(1),
control_plane,
})
}
pub fn alloc_conn_id(&self) -> u64 {
self.next_conn_id.fetch_add(1, Ordering::Relaxed)
}
pub fn register_proxy(&self, conn: Arc<ProxyConn>) {
let node_id = conn.node_id.clone();
let node_name = conn.node_name.clone();
let conn_id = conn.id;
self.proxy_conns_by_id.insert(conn_id, conn.clone());
let pool_size = {
let mut map = self.proxy_conns.write();
map.entry(node_id.clone()).or_default().push(conn);
map.get(&node_id).map(|v| v.len()).unwrap_or(0)
};
info!(
node_id = %node_id,
node_name = %node_name,
conn_id = conn_id,
pool_size = pool_size,
"proxy connected"
);
self.notify_node_status(node_id, true, pool_size);
}
pub fn unregister_proxy(&self, conn_id: u64, node_id: &str) {
self.proxy_conns_by_id.remove(&conn_id);
let pool_size = {
let mut map = self.proxy_conns.write();
if let Some(conns) = map.get_mut(node_id) {
conns.retain(|c| c.id != conn_id);
if conns.is_empty() {
map.remove(node_id);
}
}
map.get(node_id).map(|v| v.len()).unwrap_or(0)
};
info!(
node_id = %node_id,
conn_id = conn_id,
remaining = pool_size,
"proxy disconnected"
);
self.cancel_streams_for_proxy(conn_id);
self.notify_node_status(node_id.to_string(), pool_size > 0, pool_size);
}
fn notify_node_status(&self, node_id: String, connected: bool, conn_count: usize) {
let control_plane = self.control_plane.clone();
tokio::spawn(async move {
if let Err(error) = control_plane
.push_node_status(&node_id, connected, conn_count)
.await
{
warn!(
node_id = %node_id,
connected = connected,
conn_count = conn_count,
error = %error,
"failed to push node status to app control plane"
);
}
});
}
fn get_proxy_conn(&self, node_id: &str) -> Option<Arc<ProxyConn>> {
let map = self.proxy_conns.read();
let conns = map.get(node_id)?;
conns
.iter()
.filter(|c| c.is_available())
.min_by_key(|c| c.stream_count.load(Ordering::Relaxed))
.cloned()
}
pub fn open_local_stream(
&self,
node_id: &str,
meta: &protocol::RequestMeta,
) -> Result<Arc<LocalStream>, String> {
let proxy_conn = self
.get_proxy_conn(node_id)
.ok_or_else(|| format!("no proxy connection for node {node_id}"))?;
let proxy_stream_id = proxy_conn
.alloc_stream_id()
.ok_or_else(|| format!("stream limit reached for node {node_id}"))?;
// Encode frames before registering the stream so that encoding failures
// (practically impossible but theoretically possible) don't leak a stream
// slot or orphan map entries.
let meta_json = match serde_json::to_vec(meta) {
Ok(json) => json,
Err(e) => {
proxy_conn.release_stream();
return Err(format!("failed to encode request metadata: {e}"));
}
};
let (meta_payload, meta_flags) = match protocol::compress_payload(&meta_json) {
Ok(result) => result,
Err(e) => {
proxy_conn.release_stream();
return Err(format!("failed to compress request metadata: {e}"));
}
};
let header_frame = protocol::encode_frame(
proxy_stream_id,
protocol::REQUEST_HEADERS,
meta_flags,
&meta_payload,
);
// Frames encoded successfully -- now register the stream.
let local_stream_id = self.next_local_stream_id.fetch_add(1, Ordering::Relaxed);
let local_stream = Arc::new(LocalStream::new(
local_stream_id,
proxy_conn.id,
proxy_stream_id,
));
self.local_streams
.insert(local_stream_id, local_stream.clone());
self.proxy_to_local
.insert((proxy_conn.id, proxy_stream_id), local_stream_id);
match proxy_conn.send(Message::Binary(header_frame.into())) {
SendStatus::Queued => Ok(local_stream),
SendStatus::Closed | SendStatus::Congested => {
self.cleanup_local_stream(local_stream_id);
proxy_conn.release_stream();
Err("proxy connection congested".to_string())
}
}
}
pub fn push_local_request_body(
&self,
local_stream_id: u64,
payload: Bytes,
end_stream: bool,
) -> Result<(), String> {
let stream = self
.local_streams
.get(&local_stream_id)
.map(|entry| entry.value().clone())
.ok_or_else(|| "local stream not found".to_string())?;
let proxy_conn = self
.proxy_conns_by_id
.get(&stream.proxy_conn_id)
.map(|entry| entry.value().clone())
.ok_or_else(|| "proxy connection unavailable".to_string())?;
let total_chunks = payload.len().div_ceil(MAX_REQUEST_BODY_FRAME_SIZE);
if total_chunks == 0 {
if end_stream {
self.send_request_body_frame(&proxy_conn, stream.proxy_stream_id, &[], true)?;
}
} else {
for (index, chunk) in payload.chunks(MAX_REQUEST_BODY_FRAME_SIZE).enumerate() {
let is_last_chunk = index + 1 == total_chunks;
self.send_request_body_frame(
&proxy_conn,
stream.proxy_stream_id,
chunk,
end_stream && is_last_chunk,
)?;
}
}
Ok(())
}
fn send_request_body_frame(
&self,
proxy_conn: &Arc<ProxyConn>,
proxy_stream_id: u32,
payload: &[u8],
end_stream: bool,
) -> Result<(), String> {
let (body_payload, body_flags) = protocol::compress_payload(payload)
.map_err(|e| format!("failed to compress request body: {e}"))?;
let body_frame = protocol::encode_frame(
proxy_stream_id,
protocol::REQUEST_BODY,
body_flags
| if end_stream {
protocol::FLAG_END_STREAM
} else {
0
},
&body_payload,
);
match proxy_conn.send(Message::Binary(body_frame.into())) {
SendStatus::Queued => Ok(()),
SendStatus::Closed | SendStatus::Congested => {
Err("proxy connection congested".to_string())
}
}
}
pub fn cancel_local_stream(&self, local_stream_id: u64, reason: &str) {
let Some((_, stream)) = self.local_streams.remove(&local_stream_id) else {
return;
};
self.proxy_to_local
.remove(&(stream.proxy_conn_id, stream.proxy_stream_id));
if let Some(pc) = self.proxy_conns_by_id.get(&stream.proxy_conn_id) {
pc.release_stream();
let frame = protocol::encode_stream_error(stream.proxy_stream_id, reason);
let _ = pc.send(Message::Binary(frame.into()));
}
stream.fail(reason.to_string());
}
fn cleanup_local_stream(&self, local_stream_id: u64) {
let Some((_, stream)) = self.local_streams.remove(&local_stream_id) else {
return;
};
self.proxy_to_local
.remove(&(stream.proxy_conn_id, stream.proxy_stream_id));
}
pub async fn handle_proxy_frame(&self, proxy_conn_id: u64, data: &mut [u8]) {
let header = match protocol::FrameHeader::parse(data) {
Some(h) => h,
None => return,
};
let expected_len = protocol::HEADER_SIZE + header.payload_len as usize;
if data.len() < expected_len {
return;
}
match header.msg_type {
protocol::RESPONSE_HEADERS => {
self.route_response_headers(proxy_conn_id, header, data);
}
protocol::RESPONSE_BODY => {
self.route_response_body(proxy_conn_id, header, data);
}
protocol::STREAM_END => {
self.finish_proxy_stream(proxy_conn_id, header.stream_id);
}
protocol::STREAM_ERROR => {
let message = protocol::decode_payload(data, &header)
.ok()
.and_then(|payload| String::from_utf8(payload).ok())
.unwrap_or_else(|| "stream error".to_string());
self.fail_proxy_stream(proxy_conn_id, header.stream_id, message);
}
protocol::HEARTBEAT_DATA => {
self.handle_heartbeat(proxy_conn_id, header.stream_id, data, &header)
.await;
}
protocol::PING => {
let payload = protocol::frame_payload_by_header(data, &header).unwrap_or(&[]);
let pong = protocol::encode_pong(payload);
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
let _ = pc.send(Message::Binary(pong.into()));
}
}
protocol::PONG => {}
protocol::GOAWAY => {
warn!(
proxy_conn_id = proxy_conn_id,
"received GOAWAY from proxy connection"
);
}
_ => {
debug!(
msg_type = header.msg_type,
proxy_conn_id = proxy_conn_id,
"unexpected frame type from proxy"
);
}
}
}
fn route_response_headers(
&self,
proxy_conn_id: u64,
header: protocol::FrameHeader,
data: &[u8],
) {
let Some(local_id) = self.lookup_local_stream(proxy_conn_id, header.stream_id) else {
return;
};
let Ok(payload) = protocol::decode_payload(data, &header) else {
self.fail_proxy_stream(
proxy_conn_id,
header.stream_id,
"failed to decode response headers",
);
return;
};
let Ok(meta) = serde_json::from_slice::<protocol::ResponseMeta>(&payload) else {
self.fail_proxy_stream(
proxy_conn_id,
header.stream_id,
"invalid response headers payload",
);
return;
};
if let Some(entry) = self.local_streams.get(&local_id) {
entry.value().set_response_headers(meta);
}
}
fn route_response_body(&self, proxy_conn_id: u64, header: protocol::FrameHeader, data: &[u8]) {
let Some(local_id) = self.lookup_local_stream(proxy_conn_id, header.stream_id) else {
return;
};
let Ok(payload) = protocol::decode_payload(data, &header) else {
self.fail_proxy_stream(
proxy_conn_id,
header.stream_id,
"failed to decode response body",
);
return;
};
let stream = match self.local_streams.get(&local_id) {
Some(entry) => entry.value().clone(),
None => return,
};
if !stream.push_body_chunk(Bytes::from(payload)) {
self.cancel_local_stream(local_id, "local relay response congested");
}
}
fn handle_stream_cleanup(
&self,
proxy_conn_id: u64,
proxy_stream_id: u32,
) -> Option<Arc<LocalStream>> {
let local_id = self
.proxy_to_local
.remove(&(proxy_conn_id, proxy_stream_id))
.map(|(_, local_id)| local_id)?;
let stream = self
.local_streams
.remove(&local_id)
.map(|(_, stream)| stream)?;
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
pc.release_stream();
}
Some(stream)
}
fn finish_proxy_stream(&self, proxy_conn_id: u64, proxy_stream_id: u32) {
if let Some(stream) = self.handle_stream_cleanup(proxy_conn_id, proxy_stream_id) {
stream.finish();
}
}
fn fail_proxy_stream(
&self,
proxy_conn_id: u64,
proxy_stream_id: u32,
error: impl Into<String>,
) {
if let Some(stream) = self.handle_stream_cleanup(proxy_conn_id, proxy_stream_id) {
stream.fail(error.into());
}
}
fn lookup_local_stream(&self, proxy_conn_id: u64, proxy_stream_id: u32) -> Option<u64> {
self.proxy_to_local
.get(&(proxy_conn_id, proxy_stream_id))
.map(|entry| *entry.value())
}
async fn handle_heartbeat(
&self,
proxy_conn_id: u64,
stream_id: u32,
data: &[u8],
header: &protocol::FrameHeader,
) {
let payload = match protocol::decode_payload(data, header) {
Ok(payload) => payload,
Err(error) => {
warn!(proxy_conn_id = proxy_conn_id, error = %error, "failed to decode heartbeat payload");
return;
}
};
let ack_payload = match self.control_plane.heartbeat_ack(&payload).await {
Ok(payload) => payload,
Err(error) => {
warn!(proxy_conn_id = proxy_conn_id, error = %error, "control-plane heartbeat callback failed");
b"{}".to_vec()
}
};
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
let frame = protocol::encode_frame(stream_id, protocol::HEARTBEAT_ACK, 0, &ack_payload);
let _ = pc.send(Message::Binary(frame.into()));
}
}
fn cancel_streams_for_proxy(&self, proxy_conn_id: u64) {
let mut cancelled = 0usize;
self.proxy_to_local.retain(|key, local_id| {
if key.0 != proxy_conn_id {
return true;
}
if let Some((_, stream)) = self.local_streams.remove(local_id) {
stream.fail("proxy disconnected".to_string());
}
cancelled += 1;
false
});
if cancelled > 0 {
warn!(
proxy_conn_id = proxy_conn_id,
streams_cancelled = cancelled,
"cancelled in-flight streams due to proxy disconnect"
);
}
}
pub fn stats(&self) -> HubStats {
let proxy_conns = self.proxy_conns.read();
let total_proxy = proxy_conns.values().map(|v| v.len()).sum();
let nodes = proxy_conns.len();
drop(proxy_conns);
HubStats {
proxy_connections: total_proxy,
nodes,
active_streams: self.local_streams.len(),
}
}
}
#[derive(serde::Serialize)]
pub struct HubStats {
pub proxy_connections: usize,
pub nodes: usize,
pub active_streams: usize,
}
#[cfg(test)]
mod tests {
use super::*;
fn build_meta() -> protocol::RequestMeta {
protocol::RequestMeta {
method: "GET".to_string(),
url: "https://example.com".to_string(),
headers: HashMap::new(),
timeout: 30,
}
}
#[tokio::test]
async fn cancel_local_stream_notifies_proxy() {
let hub = HubRouter::new(ControlPlaneClient::disabled());
let (proxy_tx, mut proxy_rx) = mpsc::channel(8);
let (proxy_close_tx, _) = watch::channel(false);
let proxy = Arc::new(ProxyConn::new(
100,
"node-1".to_string(),
"Node 1".to_string(),
proxy_tx,
proxy_close_tx,
16,
));
hub.register_proxy(proxy);
let stream = hub
.open_local_stream("node-1", &build_meta())
.expect("open local stream");
let _ = proxy_rx.try_recv().expect("headers frame");
hub.push_local_request_body(stream.id, Bytes::new(), true)
.expect("finish empty body");
let _ = proxy_rx.try_recv().expect("body frame");
hub.cancel_local_stream(stream.id, "client dropped");
let cancelled = proxy_rx.try_recv().expect("cancel frame");
let cancelled_data = match cancelled {
Message::Binary(data) => data.to_vec(),
other => panic!("unexpected message: {other:?}"),
};
let header = protocol::FrameHeader::parse(&cancelled_data).expect("cancel frame header");
assert_eq!(header.msg_type, protocol::STREAM_ERROR);
}
#[tokio::test]
async fn push_local_request_body_splits_large_payload_and_marks_end() {
let hub = HubRouter::new(ControlPlaneClient::disabled());
let (proxy_tx, mut proxy_rx) = mpsc::channel(8);
let (proxy_close_tx, _) = watch::channel(false);
let proxy = Arc::new(ProxyConn::new(
200,
"node-2".to_string(),
"Node 2".to_string(),
proxy_tx,
proxy_close_tx,
16,
));
hub.register_proxy(proxy);
let stream = hub
.open_local_stream("node-2", &build_meta())
.expect("open local stream");
let _ = proxy_rx.try_recv().expect("headers frame");
let payload = Bytes::from(vec![b'x'; MAX_REQUEST_BODY_FRAME_SIZE + 17]);
hub.push_local_request_body(stream.id, payload, true)
.expect("push request body");
let first = match proxy_rx.try_recv().expect("first body frame") {
Message::Binary(data) => data.to_vec(),
other => panic!("unexpected message: {other:?}"),
};
let first_header = protocol::FrameHeader::parse(&first).expect("first body header");
assert_eq!(first_header.msg_type, protocol::REQUEST_BODY);
assert_eq!(first_header.flags & protocol::FLAG_END_STREAM, 0);
let second = match proxy_rx.try_recv().expect("second body frame") {
Message::Binary(data) => data.to_vec(),
other => panic!("unexpected message: {other:?}"),
};
let second_header = protocol::FrameHeader::parse(&second).expect("second body header");
assert_eq!(second_header.msg_type, protocol::REQUEST_BODY);
assert_ne!(second_header.flags & protocol::FLAG_END_STREAM, 0);
}
}
+262
View File
@@ -0,0 +1,262 @@
use std::io;
use std::net::SocketAddr;
use std::time::Duration;
use async_stream::stream;
use axum::body::{Body, Bytes};
use axum::extract::{ConnectInfo, Path, Request, State};
use axum::http::{HeaderMap, HeaderName, HeaderValue, Response, StatusCode};
use axum::response::IntoResponse;
use bytes::BytesMut;
use futures_util::StreamExt;
use tracing::warn;
use crate::hub::{LocalBodyEvent, LocalStream};
use crate::protocol;
use crate::AppState;
pub const TUNNEL_ERROR_HEADER: &str = "x-aether-tunnel-error";
const MAX_RELAY_META_LEN: usize = 256 * 1024;
struct StreamGuard {
hub: std::sync::Arc<crate::hub::HubRouter>,
stream_id: u64,
finished: bool,
}
impl Drop for StreamGuard {
fn drop(&mut self) {
if !self.finished {
self.hub
.cancel_local_stream(self.stream_id, "local relay client dropped");
}
}
}
pub async fn relay_request(
Path(node_id): Path<String>,
State(state): State<AppState>,
ConnectInfo(addr): ConnectInfo<SocketAddr>,
request: Request,
) -> impl IntoResponse {
if !addr.ip().is_loopback() {
return tunnel_error_response(
StatusCode::FORBIDDEN,
"forbidden",
"local relay only accepts loopback requests",
);
}
let mut body_stream = request.into_body().into_data_stream();
let mut envelope_buf = BytesMut::new();
let mut meta: Option<protocol::RequestMeta> = None;
let mut stream: Option<std::sync::Arc<LocalStream>> = None;
while let Some(chunk_result) = body_stream.next().await {
let chunk = match chunk_result {
Ok(chunk) => chunk,
Err(error) => {
if let Some(active_stream) = &stream {
state
.hub
.cancel_local_stream(active_stream.id, "failed to read relay request body");
}
warn!(error = %error, "failed to read local relay request body");
return tunnel_error_response(
StatusCode::BAD_GATEWAY,
"relay",
"failed to read relay request body",
);
}
};
if stream.is_none() {
envelope_buf.extend_from_slice(&chunk);
let Some((parsed_meta, body_offset)) = (match try_decode_envelope_meta(&envelope_buf) {
Ok(result) => result,
Err(error) => {
return tunnel_error_response(StatusCode::BAD_REQUEST, "bad_request", &error);
}
}) else {
continue;
};
let opened_stream = match state.hub.open_local_stream(&node_id, &parsed_meta) {
Ok(stream) => stream,
Err(error) => {
return tunnel_error_response(
StatusCode::SERVICE_UNAVAILABLE,
"connect",
&error,
);
}
};
if envelope_buf.len() > body_offset {
let first_body_chunk = Bytes::copy_from_slice(&envelope_buf[body_offset..]);
if let Err(error) =
state
.hub
.push_local_request_body(opened_stream.id, first_body_chunk, false)
{
state.hub.cancel_local_stream(opened_stream.id, &error);
return tunnel_error_response(
StatusCode::SERVICE_UNAVAILABLE,
"connect",
&error,
);
}
}
envelope_buf.clear();
meta = Some(parsed_meta);
stream = Some(opened_stream);
continue;
}
let Some(active_stream) = &stream else {
continue;
};
if let Err(error) = state
.hub
.push_local_request_body(active_stream.id, chunk, false)
{
state.hub.cancel_local_stream(active_stream.id, &error);
return tunnel_error_response(StatusCode::SERVICE_UNAVAILABLE, "connect", &error);
}
}
let (meta, stream) = match (meta, stream) {
(Some(meta), Some(stream)) => (meta, stream),
_ => {
return tunnel_error_response(
StatusCode::BAD_REQUEST,
"bad_request",
"relay envelope metadata truncated",
);
}
};
if let Err(error) = state
.hub
.push_local_request_body(stream.id, Bytes::new(), true)
{
state.hub.cancel_local_stream(stream.id, &error);
return tunnel_error_response(StatusCode::SERVICE_UNAVAILABLE, "connect", &error);
}
let request_guard = StreamGuard {
hub: state.hub.clone(),
stream_id: stream.id,
finished: false,
};
let wait_timeout = Duration::from_secs(meta.timeout.clamp(5, 300));
let response_head = match stream.wait_headers(wait_timeout).await {
Ok(response) => response,
Err(error) => {
state.hub.cancel_local_stream(stream.id, &error);
return tunnel_error_response(StatusCode::GATEWAY_TIMEOUT, "timeout", &error);
}
};
let Some(mut body_rx) = stream.take_body_receiver() else {
state
.hub
.cancel_local_stream(stream.id, "missing relay response body receiver");
return tunnel_error_response(
StatusCode::BAD_GATEWAY,
"relay",
"missing relay response body receiver",
);
};
let hub = state.hub.clone();
let stream_id = stream.id;
let body_stream = stream! {
let mut guard = request_guard;
guard.hub = hub;
guard.stream_id = stream_id;
while let Some(event) = body_rx.recv().await {
match event {
LocalBodyEvent::Chunk(chunk) => yield Ok::<Bytes, io::Error>(chunk),
LocalBodyEvent::End => {
guard.finished = true;
break;
}
LocalBodyEvent::Error(error) => {
guard.finished = true;
yield Err(io::Error::other(error));
break;
}
}
}
guard.finished = true;
};
let mut builder = Response::builder().status(response_head.status);
if let Some(headers) = builder.headers_mut() {
append_headers(headers, &response_head.headers);
}
match builder.body(Body::from_stream(body_stream)) {
Ok(response) => response,
Err(error) => {
warn!(error = %error, "failed to build relay response");
tunnel_error_response(
StatusCode::BAD_GATEWAY,
"relay",
"failed to build relay response",
)
}
}
}
fn try_decode_envelope_meta(
buffer: &BytesMut,
) -> Result<Option<(protocol::RequestMeta, usize)>, String> {
if buffer.len() < 4 {
return Ok(None);
}
let meta_len = u32::from_be_bytes([buffer[0], buffer[1], buffer[2], buffer[3]]) as usize;
if meta_len > MAX_RELAY_META_LEN {
return Err("relay metadata too large".to_string());
}
let meta_end = 4usize
.checked_add(meta_len)
.ok_or_else(|| "relay envelope length overflow".to_string())?;
if buffer.len() < meta_end {
return Ok(None);
}
let meta = serde_json::from_slice::<protocol::RequestMeta>(&buffer[4..meta_end])
.map_err(|e| format!("invalid relay metadata: {e}"))?;
Ok(Some((meta, meta_end)))
}
fn append_headers(target: &mut HeaderMap, headers: &[(String, String)]) {
for (name, value) in headers {
let Ok(name) = HeaderName::from_bytes(name.as_bytes()) else {
continue;
};
let Ok(value) = HeaderValue::from_str(value) else {
continue;
};
target.append(name, value);
}
}
fn tunnel_error_response(status: StatusCode, kind: &str, message: &str) -> Response<Body> {
let mut builder = Response::builder().status(status);
if let Some(headers) = builder.headers_mut() {
headers.insert(
HeaderName::from_static(TUNNEL_ERROR_HEADER),
HeaderValue::from_str(kind).unwrap_or_else(|_| HeaderValue::from_static("relay")),
);
headers.insert(
axum::http::header::CONTENT_TYPE,
HeaderValue::from_static("text/plain; charset=utf-8"),
);
}
builder
.body(Body::from(message.to_string()))
.unwrap_or_else(|_| Response::new(Body::from("relay error")))
}
+167
View File
@@ -0,0 +1,167 @@
mod control_plane;
mod hub;
mod local_relay;
mod protocol;
mod proxy_conn;
use std::net::SocketAddr;
use std::time::Duration;
use axum::extract::ws::WebSocketUpgrade;
use axum::extract::State;
use axum::response::{IntoResponse, Json};
use axum::routing::{get, post};
use axum::Router;
use clap::Parser;
use tracing::{info, warn};
use crate::control_plane::ControlPlaneClient;
use crate::hub::{ConnConfig, HubRouter};
use crate::local_relay::relay_request;
#[derive(Parser, Debug)]
#[command(name = "aether-hub", about = "Tunnel Hub for Aether")]
struct Args {
/// Bind address
#[arg(long, default_value = "0.0.0.0:8085", env = "TUNNEL_HUB_BIND")]
bind: String,
/// Proxy-side idle timeout in seconds (0 to disable)
#[arg(long, default_value_t = 0, env = "TUNNEL_HUB_PROXY_IDLE_TIMEOUT")]
proxy_idle_timeout: u64,
/// Ping interval in seconds (for both sides)
#[arg(long, default_value_t = 15, env = "TUNNEL_HUB_PING_INTERVAL")]
ping_interval: u64,
/// Max concurrent streams per proxy connection
#[arg(long, default_value_t = 2048, env = "TUNNEL_HUB_MAX_STREAMS")]
max_streams: usize,
/// Per-connection outbound queue capacity before treating the socket as congested
#[arg(
long,
default_value_t = 128,
env = "TUNNEL_HUB_OUTBOUND_QUEUE_CAPACITY"
)]
outbound_queue_capacity: usize,
/// Local Aether app base URL for control-plane callbacks
#[arg(
long,
default_value = "http://127.0.0.1:8084",
env = "TUNNEL_HUB_APP_BASE_URL"
)]
app_base_url: String,
}
#[derive(Clone)]
pub struct AppState {
pub hub: std::sync::Arc<HubRouter>,
pub proxy_conn_cfg: ConnConfig,
pub max_streams: usize,
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
// Initialize tracing
tracing_subscriber::fmt()
.with_env_filter(
tracing_subscriber::EnvFilter::try_from_default_env()
.unwrap_or_else(|_| "aether_hub=info".into()),
)
.init();
let args = Args::parse();
let hub = HubRouter::new(ControlPlaneClient::new(args.app_base_url));
let outbound_queue_capacity = args.outbound_queue_capacity.clamp(8, 4096);
let ping_interval = Duration::from_secs(args.ping_interval);
let state = AppState {
hub,
proxy_conn_cfg: ConnConfig {
ping_interval,
idle_timeout: Duration::from_secs(args.proxy_idle_timeout),
outbound_queue_capacity,
},
max_streams: args.max_streams,
};
let app = Router::new()
.route("/health", get(health))
.route("/stats", get(stats))
.route("/proxy", get(ws_proxy))
.route("/local/relay/{node_id}", post(relay_request))
.with_state(state);
let listener = tokio::net::TcpListener::bind(&args.bind).await?;
info!(bind = %args.bind, "aether-hub started");
axum::serve(
listener,
app.into_make_service_with_connect_info::<SocketAddr>(),
)
.await?;
Ok(())
}
// ---------------------------------------------------------------------------
// HTTP endpoints
// ---------------------------------------------------------------------------
async fn health() -> impl IntoResponse {
Json(serde_json::json!({"status": "ok"}))
}
async fn stats(State(state): State<AppState>) -> impl IntoResponse {
Json(state.hub.stats())
}
// ---------------------------------------------------------------------------
// WebSocket endpoints
// ---------------------------------------------------------------------------
async fn ws_proxy(
ws: WebSocketUpgrade,
State(state): State<AppState>,
headers: axum::http::HeaderMap,
) -> impl IntoResponse {
let node_id = headers
.get("x-node-id")
.and_then(|v| v.to_str().ok())
.unwrap_or("")
.trim()
.to_string();
let node_name = headers
.get("x-node-name")
.and_then(|v| v.to_str().ok())
.unwrap_or(&node_id)
.trim()
.to_string();
let max_streams: usize = headers
.get("x-tunnel-max-streams")
.and_then(|v| v.to_str().ok())
.and_then(|v| v.parse().ok())
.unwrap_or(state.max_streams)
.clamp(64, 2048);
if node_id.is_empty() {
warn!("proxy connection rejected: missing X-Node-ID header");
return axum::http::StatusCode::BAD_REQUEST.into_response();
}
ws.max_frame_size(64 * 1024 * 1024)
.on_upgrade(move |socket| {
proxy_conn::handle_proxy_connection(
socket,
state.hub,
node_id,
node_name,
max_streams,
state.proxy_conn_cfg,
)
})
.into_response()
}
+176
View File
@@ -0,0 +1,176 @@
/// Tunnel binary frame protocol
///
/// Frame format (10-byte header + payload):
/// | stream_id (4B) | msg_type (1B) | flags (1B) | payload_len (4B) | payload (NB) |
use std::io::Read;
use flate2::read::GzDecoder;
use flate2::write::GzEncoder;
use flate2::Compression;
pub const HEADER_SIZE: usize = 10;
// Message types
pub const REQUEST_HEADERS: u8 = 0x01;
pub const REQUEST_BODY: u8 = 0x02;
pub const RESPONSE_HEADERS: u8 = 0x03;
pub const RESPONSE_BODY: u8 = 0x04;
pub const STREAM_END: u8 = 0x05;
pub const STREAM_ERROR: u8 = 0x06;
pub const PING: u8 = 0x10;
pub const PONG: u8 = 0x11;
pub const GOAWAY: u8 = 0x12;
pub const HEARTBEAT_DATA: u8 = 0x13;
pub const HEARTBEAT_ACK: u8 = 0x14;
// Flags
pub const FLAG_END_STREAM: u8 = 0x01;
pub const FLAG_GZIP_COMPRESSED: u8 = 0x02;
#[derive(Debug, Clone, Copy)]
pub struct FrameHeader {
pub stream_id: u32,
pub msg_type: u8,
pub flags: u8,
pub payload_len: u32,
}
impl FrameHeader {
/// Parse frame header from raw bytes (must be >= HEADER_SIZE)
#[inline]
pub fn parse(data: &[u8]) -> Option<Self> {
if data.len() < HEADER_SIZE {
return None;
}
Some(Self {
stream_id: u32::from_be_bytes([data[0], data[1], data[2], data[3]]),
msg_type: data[4],
flags: data[5],
payload_len: u32::from_be_bytes([data[6], data[7], data[8], data[9]]),
})
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct RequestMeta {
pub method: String,
pub url: String,
pub headers: std::collections::HashMap<String, String>,
#[serde(default = "default_timeout", deserialize_with = "deserialize_timeout")]
pub timeout: u64,
}
fn default_timeout() -> u64 {
60
}
fn deserialize_timeout<'de, D>(deserializer: D) -> Result<u64, D::Error>
where
D: serde::Deserializer<'de>,
{
#[derive(serde::Deserialize)]
#[serde(untagged)]
enum TimeoutValue {
Int(u64),
Float(f64),
}
match <TimeoutValue as serde::Deserialize>::deserialize(deserializer)? {
TimeoutValue::Int(v) => Ok(v),
TimeoutValue::Float(v) => {
if !v.is_finite() || v < 0.0 {
return Err(serde::de::Error::custom(
"timeout must be a non-negative finite number",
));
}
if v.fract() != 0.0 {
return Err(serde::de::Error::custom("timeout must be integer seconds"));
}
if v > (u64::MAX as f64) {
return Err(serde::de::Error::custom("timeout is too large"));
}
Ok(v as u64)
}
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct ResponseMeta {
pub status: u16,
pub headers: Vec<(String, String)>,
}
pub fn encode_frame(stream_id: u32, msg_type: u8, flags: u8, payload: &[u8]) -> Vec<u8> {
let mut buf = Vec::with_capacity(HEADER_SIZE + payload.len());
buf.extend_from_slice(&stream_id.to_be_bytes());
buf.push(msg_type);
buf.push(flags);
buf.extend_from_slice(&(payload.len() as u32).to_be_bytes());
buf.extend_from_slice(payload);
buf
}
/// Encode a STREAM_ERROR frame for a given stream_id with an error message
pub fn encode_stream_error(stream_id: u32, msg: &str) -> Vec<u8> {
encode_frame(stream_id, STREAM_ERROR, 0, msg.as_bytes())
}
/// Encode a PING frame (stream_id=0)
pub fn encode_ping() -> Vec<u8> {
encode_frame(0, PING, 0, &[])
}
/// Encode a PONG frame (stream_id=0, echo payload)
pub fn encode_pong(payload: &[u8]) -> Vec<u8> {
encode_frame(0, PONG, 0, payload)
}
/// Encode a GOAWAY frame (stream_id=0)
pub fn encode_goaway() -> Vec<u8> {
encode_frame(0, GOAWAY, 0, &[])
}
#[inline]
pub fn frame_payload_by_header<'a>(data: &'a [u8], header: &FrameHeader) -> Option<&'a [u8]> {
let payload_len = header.payload_len as usize;
let end = HEADER_SIZE.checked_add(payload_len)?;
if data.len() < end {
return None;
}
Some(&data[HEADER_SIZE..end])
}
pub fn decode_payload(data: &[u8], header: &FrameHeader) -> Result<Vec<u8>, String> {
let payload = frame_payload_by_header(data, header)
.ok_or_else(|| "incomplete frame payload".to_string())?;
if header.flags & FLAG_GZIP_COMPRESSED != 0 {
let mut decoder = GzDecoder::new(payload);
let mut decoded = Vec::new();
decoder
.read_to_end(&mut decoded)
.map_err(|e| format!("failed to decompress payload: {e}"))?;
Ok(decoded)
} else {
Ok(payload.to_vec())
}
}
pub fn compress_payload(payload: &[u8]) -> Result<(Vec<u8>, u8), std::io::Error> {
maybe_recompress_payload(payload, true)
}
fn maybe_recompress_payload(
payload: &[u8],
prefer_gzip: bool,
) -> Result<(Vec<u8>, u8), std::io::Error> {
if !prefer_gzip {
return Ok((payload.to_vec(), 0));
}
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
std::io::Write::write_all(&mut encoder, payload)?;
let compressed = encoder.finish()?;
if compressed.len() < payload.len() {
Ok((compressed, FLAG_GZIP_COMPRESSED))
} else {
Ok((payload.to_vec(), 0))
}
}
+156
View File
@@ -0,0 +1,156 @@
/// Proxy-side WebSocket connection handler
///
/// Handles the lifecycle of a single aether-proxy connection:
/// accept -> authenticate (headers) -> read loop -> cleanup
use std::sync::Arc;
use std::time::Duration;
use axum::extract::ws::{Message, WebSocket};
use futures_util::{SinkExt, StreamExt};
use tokio::sync::{mpsc, watch};
use tracing::{debug, info, warn};
use crate::hub::{ConnConfig, HubRouter, ProxyConn, SendStatus};
use crate::protocol;
/// Maximum single frame size: 64 MB
const MAX_FRAME_SIZE: usize = 64 * 1024 * 1024;
pub async fn handle_proxy_connection(
ws: WebSocket,
hub: Arc<HubRouter>,
node_id: String,
node_name: String,
max_streams: usize,
cfg: ConnConfig,
) {
let conn_id = hub.alloc_conn_id();
let (mut ws_tx, ws_rx) = ws.split();
let (tx, mut rx) = mpsc::channel::<Message>(cfg.outbound_queue_capacity);
let (close_tx, mut close_rx) = watch::channel(false);
let conn = Arc::new(ProxyConn::new(
conn_id,
node_id.clone(),
node_name.clone(),
tx,
close_tx,
max_streams,
));
hub.register_proxy(conn.clone());
let writer = tokio::spawn(async move {
loop {
tokio::select! {
msg = rx.recv() => match msg {
Some(msg) => {
if ws_tx.send(msg).await.is_err() {
break;
}
}
None => break,
},
changed = close_rx.changed() => {
if changed.is_err() || *close_rx.borrow() {
break;
}
}
}
}
let _ = ws_tx.close().await;
});
let ping_conn = conn.clone();
let ping_interval = cfg.ping_interval;
let ping_task = tokio::spawn(async move {
loop {
tokio::time::sleep(ping_interval).await;
let ping = protocol::encode_ping();
if !matches!(
ping_conn.send(Message::Binary(ping.into())),
SendStatus::Queued
) {
break;
}
}
});
let reader_hub = hub.clone();
let reader_conn = conn.clone();
let reader = tokio::spawn(async move {
run_proxy_reader(ws_rx, reader_hub, reader_conn, cfg.idle_timeout).await;
});
let _ = reader.await;
ping_task.abort();
conn.request_close();
hub.unregister_proxy(conn_id, &node_id);
drop(conn);
tokio::time::sleep(Duration::from_millis(100)).await;
writer.abort();
let _ = writer.await;
}
async fn run_proxy_reader(
mut ws_rx: futures_util::stream::SplitStream<WebSocket>,
hub: Arc<HubRouter>,
conn: Arc<ProxyConn>,
idle_timeout: Duration,
) {
let idle_enabled = !idle_timeout.is_zero();
let mut oversized_count = 0u32;
loop {
let msg = if idle_enabled {
tokio::select! {
msg = ws_rx.next() => msg,
_ = tokio::time::sleep(idle_timeout) => {
warn!(conn_id = conn.id, node_id = %conn.node_id, "proxy idle timeout");
let _ = conn.send(Message::Binary(protocol::encode_goaway().into()));
conn.request_close();
break;
}
}
} else {
ws_rx.next().await
};
match msg {
Some(Ok(Message::Binary(data))) => {
let mut data = data.to_vec();
if data.len() > MAX_FRAME_SIZE {
oversized_count += 1;
warn!(
conn_id = conn.id,
size = data.len(),
"oversized frame from proxy"
);
if oversized_count >= 5 {
warn!(conn_id = conn.id, "too many oversized frames, closing");
conn.request_close();
break;
}
continue;
}
oversized_count = 0;
if data.len() < protocol::HEADER_SIZE {
debug!(conn_id = conn.id, "frame too small, skipping");
continue;
}
hub.handle_proxy_frame(conn.id, &mut data).await;
}
Some(Ok(Message::Close(_))) | None => {
info!(conn_id = conn.id, node_id = %conn.node_id, "proxy WebSocket closed");
break;
}
Some(Err(e)) => {
warn!(conn_id = conn.id, error = %e, "proxy WebSocket error");
break;
}
_ => {}
}
}
}
+8
View File
@@ -0,0 +1,8 @@
# Aether server URL
AETHER_PROXY_AETHER_URL=https://aether.example.com
# Management Token (ae_xxx, must belong to an ADMIN user)
AETHER_PROXY_MANAGEMENT_TOKEN=ae_xxxxx
# Node identification
AETHER_PROXY_NODE_NAME=proxy-01
+3403
View File
File diff suppressed because it is too large Load Diff
+44
View File
@@ -0,0 +1,44 @@
[package]
name = "aether-proxy"
version = "0.2.5"
edition = "2021"
description = "Tunnel proxy for Aether"
[dependencies]
tokio = { version = "1", features = ["full"] }
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls", "stream", "http2"] }
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"
tracing-subscriber = { version = "0.3", features = ["env-filter", "json"] }
serde = { version = "1", features = ["derive"] }
serde_json = "1"
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"
[profile.release]
lto = true
strip = true
codegen-units = 1
+10
View File
@@ -0,0 +1,10 @@
FROM debian:bookworm-slim
ARG TARGETARCH
RUN apt-get update && apt-get install -y --no-install-recommends ca-certificates \
&& rm -rf /var/lib/apt/lists/*
COPY build/linux-${TARGETARCH}/aether-proxy /usr/local/bin/aether-proxy
ENTRYPOINT ["aether-proxy"]
+154
View File
@@ -0,0 +1,154 @@
# aether-proxy
Aether Tunnel 代理节点,部署在海外 VPS 上,通过 WebSocket 隧道为 Aether 实例中转 API 流量。
Tunnel 模式下代理节点**无需对外监听端口**,仅需出站连接到 Aether 服务器。
## 安装
### Docker Compose 部署
```bash
cp .env.example .env
# 编辑 .env 填入 AETHER_PROXY_AETHER_URL 和 AETHER_PROXY_MANAGEMENT_TOKEN
docker compose up -d
```
### 下载预编译二进制
<!-- DOWNLOAD_TABLE_START -->
| Platform | Download |
|----------|----------|
| Linux x86_64 | [aether-proxy-linux-amd64.tar.gz](https://github.com/fawney19/Aether/releases/download/proxy-v0.2.5/aether-proxy-linux-amd64.tar.gz) |
| Linux ARM64 | [aether-proxy-linux-arm64.tar.gz](https://github.com/fawney19/Aether/releases/download/proxy-v0.2.5/aether-proxy-linux-arm64.tar.gz) |
| macOS x86_64 | [aether-proxy-macos-amd64.tar.gz](https://github.com/fawney19/Aether/releases/download/proxy-v0.2.5/aether-proxy-macos-amd64.tar.gz) |
| macOS ARM64 | [aether-proxy-macos-arm64.tar.gz](https://github.com/fawney19/Aether/releases/download/proxy-v0.2.5/aether-proxy-macos-arm64.tar.gz) |
| Windows x86_64 | [aether-proxy-windows-amd64.zip](https://github.com/fawney19/Aether/releases/download/proxy-v0.2.5/aether-proxy-windows-amd64.zip) |
<!-- DOWNLOAD_TABLE_END -->
## 快速开始
```bash
# 1. 首次安装配置(TUI 向导,勾选 Install Service 随系统启动服务)
sudo ./aether-proxy setup
# 2. 日常管理 (勾选 Install Service 作为系统服务的情况下)
aether-proxy status # 看状态
aether-proxy logs # 看日志
sudo aether-proxy start # 启动服务
sudo aether-proxy stop # 停止服务
sudo aether-proxy restart # 重启服务
# 3. 重新配置(改完自动重启服务)
sudo aether-proxy setup
# 4. 彻底卸载
sudo aether-proxy uninstall
```
完成向导后, 配置自动保存到 `aether-proxy.toml`,如果启用了 Install Service,将自动注册并启动 systemd 服务。
### 直接运行
如果不需要安装为系统服务,可以直接运行。缺少必填参数时会自动进入 setup 向导:
```bash
./aether-proxy
```
## 配置
配置按以下优先级加载(高优先级覆盖低优先级):
1. CLI 参数
2. 环境变量(`AETHER_PROXY_*`)
3. 配置文件(`aether-proxy.toml`,或通过 `AETHER_PROXY_CONFIG` 指定路径)
### 参数一览
#### 基础配置
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--aether-url` | `AETHER_PROXY_AETHER_URL` | **必填** | Aether 服务器地址 |
| `--management-token` | `AETHER_PROXY_MANAGEMENT_TOKEN` | **必填** | 管理员 Token(`ae_xxx` 格式) |
| `--public-ip` | `AETHER_PROXY_PUBLIC_IP` | 自动检测 | 公网 IP |
| `--node-name` | `AETHER_PROXY_NODE_NAME` | `proxy-01` | 节点名称标识 |
| `--node-region` | `AETHER_PROXY_NODE_REGION` | 自动检测 | 地区标识 |
| `--heartbeat-interval` | `AETHER_PROXY_HEARTBEAT_INTERVAL` | `30` | 心跳间隔(秒) |
| `--allowed-ports` | `AETHER_PROXY_ALLOWED_PORTS` | `80,443,8080,8443` | 允许代理的目标端口 |
#### Tunnel 连接
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--tunnel-connections` | `AETHER_PROXY_TUNNEL_CONNECTIONS` | `3` | 到 Aether 的连接池大小 |
| `--tunnel-max-streams` | `AETHER_PROXY_TUNNEL_MAX_STREAMS` | 自动(硬件估算) | 单连接最大并发 stream 数 |
| `--tunnel-connect-timeout-secs` | `AETHER_PROXY_TUNNEL_CONNECT_TIMEOUT_SECS` | `15` | TCP + TLS 握手超时(秒) |
| `--tunnel-tcp-keepalive-secs` | `AETHER_PROXY_TUNNEL_TCP_KEEPALIVE_SECS` | `30` | TCP keepalive 初始延迟(秒) |
| `--tunnel-tcp-nodelay` | `AETHER_PROXY_TUNNEL_TCP_NODELAY` | `true` | 禁用 Nagle 算法 |
| `--tunnel-ping-interval-secs` | `AETHER_PROXY_TUNNEL_PING_INTERVAL_SECS` | `15` | WebSocket Ping 频率(秒) |
| `--tunnel-stale-timeout-secs` | `AETHER_PROXY_TUNNEL_STALE_TIMEOUT_SECS` | `45` | 无数据断连阈值(秒) |
| `--tunnel-reconnect-base-ms` | `AETHER_PROXY_TUNNEL_RECONNECT_BASE_MS` | `500` | 指数退避基础延迟(毫秒) |
| `--tunnel-reconnect-max-ms` | `AETHER_PROXY_TUNNEL_RECONNECT_MAX_MS` | `30000` | 指数退避上限(毫秒) |
#### 上游 HTTP 请求
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--upstream-connect-timeout-secs` | `AETHER_PROXY_UPSTREAM_CONNECT_TIMEOUT_SECS` | `30` | 上游建连超时(秒) |
| `--upstream-pool-max-idle-per-host` | `AETHER_PROXY_UPSTREAM_POOL_MAX_IDLE_PER_HOST` | `64` | 每 Host 最大空闲连接数 |
| `--upstream-pool-idle-timeout-secs` | `AETHER_PROXY_UPSTREAM_POOL_IDLE_TIMEOUT_SECS` | `300` | 连接池空闲超时(秒) |
| `--upstream-tcp-keepalive-secs` | `AETHER_PROXY_UPSTREAM_TCP_KEEPALIVE_SECS` | `60` | TCP keepalive(秒,0 关闭) |
| `--upstream-tcp-nodelay` | `AETHER_PROXY_UPSTREAM_TCP_NODELAY` | `true` | 启用 TCP_NODELAY |
#### Aether API 客户端
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--aether-request-timeout-secs` | `AETHER_PROXY_AETHER_REQUEST_TIMEOUT_SECS` | `10` | 请求总超时(秒) |
| `--aether-connect-timeout-secs` | `AETHER_PROXY_AETHER_CONNECT_TIMEOUT_SECS` | `10` | 建连超时(秒) |
| `--aether-retry-max-attempts` | `AETHER_PROXY_AETHER_RETRY_MAX_ATTEMPTS` | `3` | 最大重试次数 |
#### DNS 与安全
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--dns-cache-ttl-secs` | `AETHER_PROXY_DNS_CACHE_TTL_SECS` | `60` | DNS 缓存 TTL(秒) |
| `--dns-cache-capacity` | `AETHER_PROXY_DNS_CACHE_CAPACITY` | `1024` | DNS 缓存容量(条目数) |
#### 日志
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--log-level` | `AETHER_PROXY_LOG_LEVEL` | `info` | 日志级别 |
| `--log-json` | `AETHER_PROXY_LOG_JSON` | `false` | JSON 格式日志 |
### 多服务器配置
在 `aether-proxy.toml` 中使用 `[[servers]]` 配置多个 Aether 服务器:
```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"
```
## 发布新版本
推送 `proxy-v*` 格式的 tag,GitHub Actions 会自动:
- 编译所有平台二进制并发布到 Releases
- 构建 Docker 镜像并推送到 GHCR 和 Docker Hub
- 更新 README 中的下载链接表格
```bash
git tag proxy-v0.2.0
git push origin proxy-v0.2.0
```
+14
View File
@@ -0,0 +1,14 @@
services:
aether-proxy:
image: ghcr.io/fawney19/aether-proxy:latest
container_name: aether-proxy
restart: unless-stopped
env_file:
- .env
environment:
AETHER_PROXY_LOG_JSON: "true"
logging:
driver: json-file
options:
max-size: "50m"
max-file: "3"
+360
View File
@@ -0,0 +1,360 @@
//! Application lifecycle: initialization, task orchestration, and shutdown.
use std::sync::atomic::AtomicU64;
use std::sync::{Arc, RwLock};
use std::time::Duration;
use arc_swap::ArcSwap;
use tokio::signal;
use tokio::sync::{watch, Mutex};
use tracing::{error, info, warn};
use crate::config::{Config, ServerEntry};
use crate::net;
use crate::registration::client::AetherClient;
use crate::runtime::{self, DynamicConfig};
use crate::state::{AppState, ProxyMetrics, ServerContext};
use crate::upstream_client;
use crate::{hardware, target_filter, tunnel};
/// Run the full application lifecycle after config has been parsed.
pub async fn run(mut config: Config, servers: Vec<ServerEntry>) -> anyhow::Result<()> {
config.validate()?;
init_tracing(&config);
info!(
version = env!("CARGO_PKG_VERSION"),
node_name = %config.node_name,
server_count = servers.len(),
"aether-proxy starting (tunnel mode)"
);
// Resolve public IP (best-effort for region info)
let public_ip = match &config.public_ip {
Some(ip) => ip.clone(),
None => net::detect_public_ip()
.await
.unwrap_or_else(|_| "0.0.0.0".to_string()),
};
// Auto-detect region if not configured
if config.node_region.is_none() {
if let Some(region) = net::detect_region(&public_ip).await {
config.node_region = Some(region);
}
}
// Collect hardware info (once at startup, sent during registration)
let hw_info = hardware::collect();
// Auto-detect tunnel_max_streams from hardware if not explicitly set
if config.tunnel_max_streams.is_none() {
let auto = (hw_info.estimated_max_concurrency / 10).clamp(64, 1024) as u32;
config.tunnel_max_streams = Some(auto);
info!(
tunnel_max_streams = auto,
"auto-detected tunnel_max_streams from hardware"
);
}
info!(
max_concurrency = hw_info.estimated_max_concurrency,
"hardware info collected"
);
let dns_cache = Arc::new(target_filter::DnsCache::new(
Duration::from_secs(config.dns_cache_ttl_secs),
config.dns_cache_capacity,
));
// Build Hyper client for tunnel upstream requests (shared).
// DNS still flows through validated addresses from DnsCache, while the
// custom connector exposes per-request connect/TLS timing when available.
let upstream_client = upstream_client::build_upstream_client(&config, Arc::clone(&dns_cache));
// Register with each Aether server and build per-server contexts.
// Wrapped in Arc<Mutex> so retry_failed_registrations can append later.
let server_contexts: Arc<Mutex<Vec<Arc<ServerContext>>>> = Arc::new(Mutex::new(Vec::new()));
let mut failed_entries: Vec<(String, ServerEntry)> = Vec::new();
for (i, entry) in servers.iter().enumerate() {
let label = if servers.len() == 1 {
"server".to_string()
} else {
format!("server-{}", i)
};
let node_name = entry
.node_name
.clone()
.unwrap_or_else(|| config.node_name.clone());
let client = Arc::new(AetherClient::new(
&config,
&entry.aether_url,
&entry.management_token,
));
match client
.register(&config, &node_name, &public_ip, Some(&hw_info))
.await
{
Ok(node_id) => {
info!(server = %label, node_id = %node_id, url = %entry.aether_url, node_name = %node_name, "registered");
// Initialize dynamic config with per-server node_name (not global),
// so that the heartbeat and reconnect use the correct name.
let mut dynamic = DynamicConfig::from_config(&config);
dynamic.node_name = node_name.clone();
server_contexts.lock().await.push(Arc::new(ServerContext {
server_label: label,
aether_url: entry.aether_url.clone(),
management_token: entry.management_token.clone(),
node_name,
node_id: Arc::new(RwLock::new(node_id)),
aether_client: client,
dynamic: Arc::new(ArcSwap::from_pointee(dynamic)),
active_connections: Arc::new(AtomicU64::new(0)),
metrics: Arc::new(ProxyMetrics::new()),
}));
}
Err(e) => {
warn!(
server = %label,
url = %entry.aether_url,
error = %e,
"registration failed, will retry in background"
);
failed_entries.push((label, entry.clone()));
}
}
}
{
let ctx_count = server_contexts.lock().await.len();
if ctx_count == 0 && failed_entries.is_empty() {
anyhow::bail!("no servers configured");
}
if ctx_count == 0 {
anyhow::bail!(
"no servers registered successfully (all {} failed)",
failed_entries.len()
);
}
}
// Build shared application state
let tunnel_tls_config = Arc::new(crate::tunnel::client::build_tls_config());
let state = Arc::new(AppState {
config: Arc::new(config),
dns_cache,
upstream_client,
tunnel_tls_config,
});
// Shutdown signal channel
let (shutdown_tx, shutdown_rx) = watch::channel(false);
info!(
active_servers = server_contexts.lock().await.len(),
"running in tunnel mode"
);
// Spawn tunnel connections per server (pool_size connections each)
let pool_size = state.config.tunnel_connections.max(1) as usize;
let mut tunnel_handles = Vec::new();
for server in server_contexts.lock().await.iter() {
for conn_idx in 0..pool_size {
let s = Arc::clone(&state);
let srv = Arc::clone(server);
let rx = shutdown_rx.clone();
tunnel_handles.push(tokio::spawn(async move {
tunnel::run(&s, &srv, conn_idx, rx).await;
}));
}
}
// Spawn background retry for failed server registrations
if !failed_entries.is_empty() {
let retry_state = Arc::clone(&state);
let retry_contexts = Arc::clone(&server_contexts);
let retry_public_ip = public_ip.clone();
let retry_hw_info = hw_info.clone();
let retry_shutdown = shutdown_rx.clone();
let retry_pool_size = pool_size;
tokio::spawn(async move {
retry_failed_registrations(
retry_state,
retry_contexts,
failed_entries,
retry_public_ip,
retry_hw_info,
retry_pool_size,
retry_shutdown,
)
.await;
});
}
// Wait for shutdown signal
wait_for_shutdown().await;
info!("shutdown signal received, cleaning up...");
let _ = shutdown_tx.send(true);
// Graceful unregister from all servers (including retry-registered ones)
for server in server_contexts.lock().await.iter() {
let node_id = server.node_id.read().unwrap().clone();
if let Err(e) = server.aether_client.unregister(&node_id).await {
error!(
server = %server.server_label,
error = %e,
"unregister failed during shutdown"
);
}
}
// Wait for all tunnel tasks
for h in tunnel_handles {
let _ = h.await;
}
info!("aether-proxy stopped");
Ok(())
}
/// Retry interval for failed server registrations (5 minutes).
const REGISTRATION_RETRY_INTERVAL: Duration = Duration::from_secs(300);
/// Max registration retry attempts before giving up.
const REGISTRATION_RETRY_MAX: u32 = 12;
/// Background task that retries registration for servers that failed at startup.
async fn retry_failed_registrations(
state: Arc<AppState>,
server_contexts: Arc<Mutex<Vec<Arc<ServerContext>>>>,
failed: Vec<(String, ServerEntry)>,
public_ip: String,
hw_info: crate::hardware::HardwareInfo,
pool_size: usize,
mut shutdown: watch::Receiver<bool>,
) {
for (label, entry) in &failed {
let node_name = entry
.node_name
.clone()
.unwrap_or_else(|| state.config.node_name.clone());
let client = Arc::new(AetherClient::new(
&state.config,
&entry.aether_url,
&entry.management_token,
));
let mut attempt = 0u32;
let node_id = loop {
attempt += 1;
tokio::select! {
_ = tokio::time::sleep(REGISTRATION_RETRY_INTERVAL) => {}
_ = shutdown.changed() => {
info!(server = %label, "shutdown during registration retry");
return;
}
}
match client
.register(&state.config, &node_name, &public_ip, Some(&hw_info))
.await
{
Ok(id) => {
info!(server = %label, node_id = %id, attempt, "registration retry succeeded");
break id;
}
Err(e) => {
warn!(
server = %label,
attempt,
max = REGISTRATION_RETRY_MAX,
error = %e,
"registration retry failed"
);
if attempt >= REGISTRATION_RETRY_MAX {
error!(server = %label, "giving up registration after {} attempts", attempt);
return;
}
}
}
};
// Build server context and spawn tunnels
let mut dynamic = DynamicConfig::from_config(&state.config);
dynamic.node_name = node_name.clone();
let server = Arc::new(ServerContext {
server_label: label.clone(),
aether_url: entry.aether_url.clone(),
management_token: entry.management_token.clone(),
node_name,
node_id: Arc::new(RwLock::new(node_id)),
aether_client: client,
dynamic: Arc::new(ArcSwap::from_pointee(dynamic)),
active_connections: Arc::new(AtomicU64::new(0)),
metrics: Arc::new(ProxyMetrics::new()),
});
// Add to shared list so shutdown can unregister this server
server_contexts.lock().await.push(Arc::clone(&server));
for conn_idx in 0..pool_size {
let s = Arc::clone(&state);
let srv = Arc::clone(&server);
let rx = shutdown.clone();
tokio::spawn(async move {
tunnel::run(&s, &srv, conn_idx, rx).await;
});
}
}
}
fn init_tracing(config: &Config) {
use tracing_subscriber::prelude::*;
use tracing_subscriber::{reload, EnvFilter};
let filter = EnvFilter::try_new(&config.log_level).unwrap_or_else(|_| EnvFilter::new("info"));
let (filter_layer, reload_handle) = reload::Layer::new(filter);
runtime::set_log_reloader(Box::new(move |level: &str| {
if let Ok(new_filter) = EnvFilter::try_new(level) {
let _ = reload_handle.modify(|f| *f = new_filter);
}
}));
if config.log_json {
tracing_subscriber::registry()
.with(filter_layer)
.with(tracing_subscriber::fmt::layer().json())
.init();
} else {
tracing_subscriber::registry()
.with(filter_layer)
.with(tracing_subscriber::fmt::layer())
.init();
}
}
async fn wait_for_shutdown() {
let ctrl_c = async {
signal::ctrl_c()
.await
.expect("failed to install Ctrl+C handler");
};
#[cfg(unix)]
let terminate = async {
signal::unix::signal(signal::unix::SignalKind::terminate())
.expect("failed to install SIGTERM handler")
.recv()
.await;
};
#[cfg(not(unix))]
let terminate = std::future::pending::<()>();
tokio::select! {
_ = ctrl_c => {},
_ = terminate => {},
}
}
+656
View File
@@ -0,0 +1,656 @@
use std::path::Path;
use clap::Parser;
use serde::{Deserialize, Serialize};
/// Fields that existed in 0.1.x but were removed in 0.2.0.
const LEGACY_ONLY_KEYS: &[&str] = &[
"hmac_key",
"listen_port",
"timestamp_tolerance",
"connect_timeout_secs",
"tls_handshake_timeout_secs",
"enable_tls",
"tls_cert",
"tls_key",
];
/// Fields renamed from 0.1.x `delegate_*` to 0.2.0 `upstream_*`.
const DELEGATE_TO_UPSTREAM: &[(&str, &str)] = &[
(
"delegate_connect_timeout_secs",
"upstream_connect_timeout_secs",
),
(
"delegate_pool_max_idle_per_host",
"upstream_pool_max_idle_per_host",
),
(
"delegate_pool_idle_timeout_secs",
"upstream_pool_idle_timeout_secs",
),
("delegate_tcp_keepalive_secs", "upstream_tcp_keepalive_secs"),
("delegate_tcp_nodelay", "upstream_tcp_nodelay"),
];
/// Aether tunnel proxy.
///
/// Deployed on overseas VPS to relay API traffic for Aether instances
/// behind the GFW. Connects to Aether via WebSocket tunnel, registers
/// with Aether, and relays upstream requests.
#[derive(Parser, Debug, Clone)]
#[command(version, about)]
pub struct Config {
/// Aether server URL (e.g. https://aether.example.com)
#[arg(long, env = "AETHER_PROXY_AETHER_URL")]
pub aether_url: String,
/// Management Token for Aether admin API (ae_xxx)
#[arg(long, env = "AETHER_PROXY_MANAGEMENT_TOKEN")]
pub management_token: String,
/// Public IP address of this node (auto-detected if omitted)
#[arg(long, env = "AETHER_PROXY_PUBLIC_IP")]
pub public_ip: Option<String>,
/// Human-readable node name
#[arg(long, env = "AETHER_PROXY_NODE_NAME", default_value = "proxy-01")]
pub node_name: String,
/// Region label (e.g. ap-northeast-1)
#[arg(long, env = "AETHER_PROXY_NODE_REGION")]
pub node_region: Option<String>,
/// Heartbeat interval in seconds
#[arg(long, env = "AETHER_PROXY_HEARTBEAT_INTERVAL", default_value_t = 30)]
pub heartbeat_interval: u64,
/// Allowed destination ports (default: 80,443,8080,8443)
#[arg(
long,
env = "AETHER_PROXY_ALLOWED_PORTS",
value_delimiter = ',',
default_values_t = vec![80, 443, 8080, 8443]
)]
pub allowed_ports: Vec<u16>,
/// Aether API request timeout in seconds
#[arg(
long,
env = "AETHER_PROXY_AETHER_REQUEST_TIMEOUT",
default_value_t = 10
)]
pub aether_request_timeout_secs: u64,
/// Aether API connect timeout in seconds
#[arg(
long,
env = "AETHER_PROXY_AETHER_CONNECT_TIMEOUT",
default_value_t = 10
)]
pub aether_connect_timeout_secs: u64,
/// Aether API max idle connections per host
#[arg(
long,
env = "AETHER_PROXY_AETHER_POOL_MAX_IDLE_PER_HOST",
default_value_t = 8
)]
pub aether_pool_max_idle_per_host: usize,
/// Aether API idle timeout in seconds
#[arg(
long,
env = "AETHER_PROXY_AETHER_POOL_IDLE_TIMEOUT",
default_value_t = 90
)]
pub aether_pool_idle_timeout_secs: u64,
/// Aether API TCP keepalive in seconds (0 disables)
#[arg(long, env = "AETHER_PROXY_AETHER_TCP_KEEPALIVE", default_value_t = 60)]
pub aether_tcp_keepalive_secs: u64,
/// Aether API TCP_NODELAY
#[arg(long, env = "AETHER_PROXY_AETHER_TCP_NODELAY", default_value_t = true)]
pub aether_tcp_nodelay: bool,
/// Enable HTTP/2 when talking to Aether API
#[arg(long, env = "AETHER_PROXY_AETHER_HTTP2", default_value_t = true)]
pub aether_http2: bool,
/// Aether API retry attempts (including initial)
#[arg(
long,
env = "AETHER_PROXY_AETHER_RETRY_MAX_ATTEMPTS",
default_value_t = 3
)]
pub aether_retry_max_attempts: u32,
/// Aether API retry base delay in milliseconds
#[arg(
long,
env = "AETHER_PROXY_AETHER_RETRY_BASE_DELAY_MS",
default_value_t = 200
)]
pub aether_retry_base_delay_ms: u64,
/// Aether API retry max delay in milliseconds
#[arg(
long,
env = "AETHER_PROXY_AETHER_RETRY_MAX_DELAY_MS",
default_value_t = 2000
)]
pub aether_retry_max_delay_ms: u64,
/// Maximum concurrent TCP connections (defaults to hardware estimate)
#[arg(long, env = "AETHER_PROXY_MAX_CONCURRENT_CONNECTIONS")]
pub max_concurrent_connections: Option<u64>,
/// DNS cache TTL in seconds
#[arg(long, env = "AETHER_PROXY_DNS_CACHE_TTL", default_value_t = 60)]
pub dns_cache_ttl_secs: u64,
/// DNS cache capacity (entries)
#[arg(long, env = "AETHER_PROXY_DNS_CACHE_CAPACITY", default_value_t = 1024)]
pub dns_cache_capacity: usize,
/// Upstream HTTP client connect timeout in seconds
#[arg(
long,
env = "AETHER_PROXY_UPSTREAM_CONNECT_TIMEOUT",
default_value_t = 30
)]
pub upstream_connect_timeout_secs: u64,
/// Upstream HTTP client max idle connections per host
#[arg(
long,
env = "AETHER_PROXY_UPSTREAM_POOL_MAX_IDLE_PER_HOST",
default_value_t = 64
)]
pub upstream_pool_max_idle_per_host: usize,
/// Upstream HTTP client idle timeout in seconds
#[arg(
long,
env = "AETHER_PROXY_UPSTREAM_POOL_IDLE_TIMEOUT",
default_value_t = 300
)]
pub upstream_pool_idle_timeout_secs: u64,
/// Upstream TCP keepalive in seconds (0 disables)
#[arg(
long,
env = "AETHER_PROXY_UPSTREAM_TCP_KEEPALIVE",
default_value_t = 60
)]
pub upstream_tcp_keepalive_secs: u64,
/// Upstream TCP_NODELAY
#[arg(
long,
env = "AETHER_PROXY_UPSTREAM_TCP_NODELAY",
default_value_t = true
)]
pub upstream_tcp_nodelay: bool,
/// Log level (trace, debug, info, warn, error)
#[arg(long, env = "AETHER_PROXY_LOG_LEVEL", default_value = "info")]
pub log_level: String,
/// Output logs as JSON
#[arg(long, env = "AETHER_PROXY_LOG_JSON", default_value_t = false)]
pub log_json: bool,
/// Tunnel reconnect base delay in milliseconds (used by exponential backoff)
#[arg(
long,
env = "AETHER_PROXY_TUNNEL_RECONNECT_BASE_MS",
default_value_t = 500
)]
pub tunnel_reconnect_base_ms: u64,
/// Tunnel reconnect max delay in milliseconds (cap for exponential backoff)
#[arg(
long,
env = "AETHER_PROXY_TUNNEL_RECONNECT_MAX_MS",
default_value_t = 30000
)]
pub tunnel_reconnect_max_ms: u64,
/// WebSocket tunnel ping interval in seconds
#[arg(long, env = "AETHER_PROXY_TUNNEL_PING_INTERVAL", default_value_t = 15)]
pub tunnel_ping_interval_secs: u64,
/// Maximum concurrent streams over tunnel (auto-detected from hardware if omitted)
#[arg(long, env = "AETHER_PROXY_TUNNEL_MAX_STREAMS")]
pub tunnel_max_streams: Option<u32>,
/// WebSocket tunnel TCP connect timeout in seconds
#[arg(
long,
env = "AETHER_PROXY_TUNNEL_CONNECT_TIMEOUT",
default_value_t = 15
)]
pub tunnel_connect_timeout_secs: u64,
/// WebSocket tunnel TCP keepalive in seconds (0 disables)
#[arg(long, env = "AETHER_PROXY_TUNNEL_TCP_KEEPALIVE", default_value_t = 30)]
pub tunnel_tcp_keepalive_secs: u64,
/// WebSocket tunnel TCP_NODELAY
#[arg(long, env = "AETHER_PROXY_TUNNEL_TCP_NODELAY", default_value_t = true)]
pub tunnel_tcp_nodelay: bool,
/// Tunnel connection staleness timeout in seconds (triggers reconnect if no data received)
#[arg(long, env = "AETHER_PROXY_TUNNEL_STALE_TIMEOUT", default_value_t = 45)]
pub tunnel_stale_timeout_secs: u64,
/// Number of parallel WebSocket tunnel connections per server (connection pool)
#[arg(long, env = "AETHER_PROXY_TUNNEL_CONNECTIONS", default_value_t = 3)]
pub tunnel_connections: u32,
}
impl Config {
/// Validate configuration values are within sane ranges.
/// Called after parsing to catch misconfigurations early.
pub fn validate(&self) -> anyhow::Result<()> {
if self.heartbeat_interval == 0 {
anyhow::bail!("heartbeat_interval must be > 0");
}
if self.heartbeat_interval > 3600 {
anyhow::bail!("heartbeat_interval must be <= 3600");
}
if self.allowed_ports.is_empty() {
anyhow::bail!("allowed_ports must not be empty");
}
for &port in &self.allowed_ports {
if port == 0 {
anyhow::bail!("allowed_ports: port 0 is not valid");
}
}
if self.tunnel_connect_timeout_secs == 0 {
anyhow::bail!("tunnel_connect_timeout_secs must be > 0");
}
if self.tunnel_ping_interval_secs == 0 {
anyhow::bail!("tunnel_ping_interval_secs must be > 0");
}
if self.tunnel_stale_timeout_secs <= self.tunnel_ping_interval_secs {
anyhow::bail!(
"tunnel_stale_timeout_secs ({}) must be > tunnel_ping_interval_secs ({})",
self.tunnel_stale_timeout_secs,
self.tunnel_ping_interval_secs
);
}
if self.tunnel_connections == 0 {
anyhow::bail!("tunnel_connections must be > 0");
}
if self.aether_retry_max_attempts == 0 {
anyhow::bail!("aether_retry_max_attempts must be >= 1");
}
if self.upstream_connect_timeout_secs == 0 {
anyhow::bail!("upstream_connect_timeout_secs must be > 0");
}
Ok(())
}
}
/// Per-server connection config (used in multi-server TOML `[[servers]]`).
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ServerEntry {
pub aether_url: String,
pub management_token: String,
/// Per-server node name override. Falls back to the global `node_name`.
pub node_name: Option<String>,
}
// ---------------------------------------------------------------------------
// TOML config file support
// ---------------------------------------------------------------------------
/// Serializable config for TOML file persistence.
/// All fields are optional -- only populated values are written.
#[derive(Debug, Default, Serialize, Deserialize)]
pub struct ConfigFile {
#[serde(skip_serializing_if = "Option::is_none")]
pub aether_url: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub management_token: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub public_ip: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub node_name: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub node_region: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub heartbeat_interval: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub allowed_ports: Option<Vec<u16>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub aether_request_timeout_secs: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub aether_connect_timeout_secs: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub aether_pool_max_idle_per_host: Option<usize>,
#[serde(skip_serializing_if = "Option::is_none")]
pub aether_pool_idle_timeout_secs: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub aether_tcp_keepalive_secs: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub aether_tcp_nodelay: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub aether_http2: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub aether_retry_max_attempts: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub aether_retry_base_delay_ms: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub aether_retry_max_delay_ms: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_concurrent_connections: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub dns_cache_ttl_secs: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub dns_cache_capacity: Option<usize>,
#[serde(skip_serializing_if = "Option::is_none")]
pub upstream_connect_timeout_secs: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub upstream_pool_max_idle_per_host: Option<usize>,
#[serde(skip_serializing_if = "Option::is_none")]
pub upstream_pool_idle_timeout_secs: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub upstream_tcp_keepalive_secs: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub upstream_tcp_nodelay: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub log_level: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub log_json: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tunnel_reconnect_base_ms: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tunnel_reconnect_max_ms: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tunnel_ping_interval_secs: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tunnel_max_streams: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tunnel_connect_timeout_secs: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tunnel_tcp_keepalive_secs: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tunnel_tcp_nodelay: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tunnel_stale_timeout_secs: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tunnel_connections: Option<u32>,
/// Multi-server config: each entry connects to a separate Aether instance.
/// When present, top-level aether_url/management_token are ignored for
/// tunnel connections (but still injected as env for clap compatibility).
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub servers: Vec<ServerEntry>,
}
impl ConfigFile {
/// Load from a TOML file.
pub fn load(path: &Path) -> anyhow::Result<Self> {
let content = std::fs::read_to_string(path)?;
Ok(toml::from_str(&content)?)
}
/// Save to a TOML file.
pub fn save(&self, path: &Path) -> anyhow::Result<()> {
let content = toml::to_string_pretty(self)?;
std::fs::write(path, content)?;
Ok(())
}
/// Detect and migrate a 0.1.x config file to 0.2.0 format in-place.
///
/// Returns `true` if migration was performed, `false` if already current.
/// The original file is backed up as `<name>.v1.bak` before rewriting.
pub fn migrate_legacy(path: &Path) -> anyhow::Result<bool> {
let content = match std::fs::read_to_string(path) {
Ok(c) => c,
Err(_) => return Ok(false),
};
let mut table: toml::map::Map<String, toml::Value> = toml::from_str(&content)?;
// Detect legacy format: presence of any 0.1.x-only key.
let is_legacy = LEGACY_ONLY_KEYS.iter().any(|k| table.contains_key(*k))
|| DELEGATE_TO_UPSTREAM
.iter()
.any(|(old, _)| table.contains_key(*old));
if !is_legacy {
return Ok(false);
}
// 1. Rename delegate_* -> upstream_* (carry over user-customized values)
for &(old, new) in DELEGATE_TO_UPSTREAM {
if let Some(val) = table.remove(old) {
table.entry(new.to_string()).or_insert(val);
}
}
// 2. Build [[servers]] from top-level aether_url + management_token + node_name
if !table.contains_key("servers") {
let aether_url = table.get("aether_url").and_then(|v| v.as_str());
let management_token = table.get("management_token").and_then(|v| v.as_str());
if let (Some(url), Some(token)) = (aether_url, management_token) {
let mut entry = toml::map::Map::new();
entry.insert("aether_url".into(), toml::Value::String(url.to_string()));
entry.insert(
"management_token".into(),
toml::Value::String(token.to_string()),
);
if let Some(name) = table.get("node_name").and_then(|v| v.as_str()) {
entry.insert("node_name".into(), toml::Value::String(name.to_string()));
}
table.insert(
"servers".into(),
toml::Value::Array(vec![toml::Value::Table(entry)]),
);
}
}
// 3. Remove top-level fields that are now in [[servers]] or obsolete
table.remove("aether_url");
table.remove("management_token");
table.remove("node_name");
for &key in LEGACY_ONLY_KEYS {
table.remove(key);
}
// 4. Backup original file (abort migration if backup fails)
let backup_path = path.with_extension("v1.bak");
std::fs::copy(path, &backup_path).map_err(|e| {
anyhow::anyhow!(
"failed to backup config before migration: {} -> {}: {}",
path.display(),
backup_path.display(),
e
)
})?;
// 5. Write migrated config
let new_content = toml::to_string_pretty(&table)?;
std::fs::write(path, &new_content)?;
eprintln!(" Config migrated from 0.1.x to 0.2.0 format.");
eprintln!(" Backup saved: {}", backup_path.display());
Ok(true)
}
/// Resolve the effective server list.
///
/// If `[[servers]]` is present, use it. Otherwise fall back to the
/// top-level `aether_url` + `management_token` as a single server.
pub fn effective_servers(&self) -> Vec<ServerEntry> {
if !self.servers.is_empty() {
return self.servers.clone();
}
match (&self.aether_url, &self.management_token) {
(Some(url), Some(token)) => vec![ServerEntry {
aether_url: url.clone(),
management_token: token.clone(),
node_name: None,
}],
_ => vec![],
}
}
/// Inject values as environment variables so clap picks them up.
///
/// Only sets variables that are **not** already present in the
/// environment, preserving the precedence: CLI > env > config file.
pub fn inject_env(&self) {
self.inject_env_inner(false);
}
/// Inject values as environment variables, **overriding** any existing
/// values. Used after setup to ensure the freshly-saved config takes
/// effect before re-parsing.
pub fn inject_env_override(&self) {
self.inject_env_inner(true);
}
fn inject_env_inner(&self, force: bool) {
macro_rules! set {
($env:expr, $val:expr) => {
if let Some(ref v) = $val {
if force || std::env::var($env).is_err() {
std::env::set_var($env, v.to_string());
}
}
};
}
// When top-level fields are absent, fall back to the first [[servers]]
// entry so that clap's required `aether_url` / `management_token` are
// satisfied even with the new config format.
let first_server = self.servers.first();
let aether_url = self
.aether_url
.as_deref()
.or(first_server.map(|s| s.aether_url.as_str()));
let management_token = self
.management_token
.as_deref()
.or(first_server.map(|s| s.management_token.as_str()));
let node_name = self
.node_name
.as_deref()
.or(first_server.and_then(|s| s.node_name.as_deref()));
set!("AETHER_PROXY_AETHER_URL", aether_url);
set!("AETHER_PROXY_MANAGEMENT_TOKEN", management_token);
set!("AETHER_PROXY_PUBLIC_IP", self.public_ip);
set!("AETHER_PROXY_NODE_NAME", node_name);
set!("AETHER_PROXY_NODE_REGION", self.node_region);
set!("AETHER_PROXY_HEARTBEAT_INTERVAL", self.heartbeat_interval);
set!(
"AETHER_PROXY_AETHER_REQUEST_TIMEOUT",
self.aether_request_timeout_secs
);
set!(
"AETHER_PROXY_AETHER_CONNECT_TIMEOUT",
self.aether_connect_timeout_secs
);
set!(
"AETHER_PROXY_AETHER_POOL_MAX_IDLE_PER_HOST",
self.aether_pool_max_idle_per_host
);
set!(
"AETHER_PROXY_AETHER_POOL_IDLE_TIMEOUT",
self.aether_pool_idle_timeout_secs
);
set!(
"AETHER_PROXY_AETHER_TCP_KEEPALIVE",
self.aether_tcp_keepalive_secs
);
set!("AETHER_PROXY_AETHER_TCP_NODELAY", self.aether_tcp_nodelay);
set!("AETHER_PROXY_AETHER_HTTP2", self.aether_http2);
set!(
"AETHER_PROXY_AETHER_RETRY_MAX_ATTEMPTS",
self.aether_retry_max_attempts
);
set!(
"AETHER_PROXY_AETHER_RETRY_BASE_DELAY_MS",
self.aether_retry_base_delay_ms
);
set!(
"AETHER_PROXY_AETHER_RETRY_MAX_DELAY_MS",
self.aether_retry_max_delay_ms
);
set!(
"AETHER_PROXY_MAX_CONCURRENT_CONNECTIONS",
self.max_concurrent_connections
);
set!("AETHER_PROXY_DNS_CACHE_TTL", self.dns_cache_ttl_secs);
set!("AETHER_PROXY_DNS_CACHE_CAPACITY", self.dns_cache_capacity);
set!(
"AETHER_PROXY_UPSTREAM_CONNECT_TIMEOUT",
self.upstream_connect_timeout_secs
);
set!(
"AETHER_PROXY_UPSTREAM_POOL_MAX_IDLE_PER_HOST",
self.upstream_pool_max_idle_per_host
);
set!(
"AETHER_PROXY_UPSTREAM_POOL_IDLE_TIMEOUT",
self.upstream_pool_idle_timeout_secs
);
set!(
"AETHER_PROXY_UPSTREAM_TCP_KEEPALIVE",
self.upstream_tcp_keepalive_secs
);
set!(
"AETHER_PROXY_UPSTREAM_TCP_NODELAY",
self.upstream_tcp_nodelay
);
set!("AETHER_PROXY_LOG_LEVEL", self.log_level);
set!("AETHER_PROXY_LOG_JSON", self.log_json);
set!(
"AETHER_PROXY_TUNNEL_RECONNECT_BASE_MS",
self.tunnel_reconnect_base_ms
);
set!(
"AETHER_PROXY_TUNNEL_RECONNECT_MAX_MS",
self.tunnel_reconnect_max_ms
);
set!(
"AETHER_PROXY_TUNNEL_PING_INTERVAL",
self.tunnel_ping_interval_secs
);
set!("AETHER_PROXY_TUNNEL_MAX_STREAMS", self.tunnel_max_streams);
set!(
"AETHER_PROXY_TUNNEL_CONNECT_TIMEOUT",
self.tunnel_connect_timeout_secs
);
set!(
"AETHER_PROXY_TUNNEL_TCP_KEEPALIVE",
self.tunnel_tcp_keepalive_secs
);
set!("AETHER_PROXY_TUNNEL_TCP_NODELAY", self.tunnel_tcp_nodelay);
set!(
"AETHER_PROXY_TUNNEL_STALE_TIMEOUT",
self.tunnel_stale_timeout_secs
);
set!("AETHER_PROXY_TUNNEL_CONNECTIONS", self.tunnel_connections);
// allowed_ports needs special handling (comma-separated)
if let Some(ref ports) = self.allowed_ports {
if force || std::env::var("AETHER_PROXY_ALLOWED_PORTS").is_err() {
let s: String = ports
.iter()
.map(|p| p.to_string())
.collect::<Vec<_>>()
.join(",");
std::env::set_var("AETHER_PROXY_ALLOWED_PORTS", s);
}
}
}
}
+79
View File
@@ -0,0 +1,79 @@
use serde::Serialize;
use sysinfo::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,
}
/// 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
}
+168
View File
@@ -0,0 +1,168 @@
mod app;
mod config;
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-proxy.toml";
/// Build the full clap command: Config args + discoverable subcommands.
///
/// `subcommand_negates_reqs` lets subcommands bypass the required Config
/// flags so that e.g. `aether-proxy 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 systemd 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 systemd service"))
.subcommand(clap::Command::new("stop").about("Stop the systemd service"))
.subcommand(clap::Command::new("uninstall").about("Uninstall the systemd 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"))?;
// Load config file as env-var defaults (before clap parsing)
let config_file_path =
std::env::var("AETHER_PROXY_CONFIG").unwrap_or_else(|_| DEFAULT_CONFIG.to_string());
let config_path = std::path::Path::new(&config_file_path);
if config_path.exists() {
// Migrate legacy 0.1.x config to 0.2.0 format if needed
if let Err(e) = config::ConfigFile::migrate_legacy(config_path) {
eprintln!(" WARNING: config migration failed: {}", e);
}
if let Ok(file_cfg) = config::ConfigFile::load(config_path) {
file_cfg.inject_env();
}
}
// 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 proxy with parsed config.
let config = Config::from_arg_matches(&matches)?;
run_proxy(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();
}
}
}
}
/// 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-proxy"])
.map_err(|e| anyhow::anyhow!("config invalid after setup: {}", e))?;
eprintln!(" Starting proxy...\n");
run_proxy(config).await
}
setup::SetupOutcome::Cancelled => {
eprintln!(" Setup cancelled.");
Ok(())
}
}
}
/// Start the proxy server, checking for systemd conflicts first.
async fn run_proxy(config: Config) -> anyhow::Result<()> {
// Warn if systemd service is already running (would cause port conflict).
// Skip this check when we ARE the systemd service (INVOCATION_ID is set by systemd).
if std::env::var_os("INVOCATION_ID").is_none() && setup::service::is_service_active() {
eprintln!("Warning: systemd service is already running.");
eprintln!("Use `./aether-proxy stop` to stop it first, or manage via subcommands:");
eprintln!(" ./aether-proxy status / logs / restart / stop");
std::process::exit(1);
}
// Resolve server list: prefer [[servers]] from TOML, fall back to CLI/env single server.
let config_path =
std::env::var("AETHER_PROXY_CONFIG").unwrap_or_else(|_| DEFAULT_CONFIG.to_string());
let servers = if std::path::Path::new(&config_path).exists() {
config::ConfigFile::load(std::path::Path::new(&config_path))
.ok()
.map(|f| f.effective_servers())
.filter(|s| !s.is_empty())
.unwrap_or_else(|| {
vec![config::ServerEntry {
aether_url: config.aether_url.clone(),
management_token: config.management_token.clone(),
node_name: None,
}]
})
} else {
vec![config::ServerEntry {
aether_url: config.aether_url.clone(),
management_token: config.management_token.clone(),
node_name: None,
}]
};
app::run(config, servers).await
}
@@ -2,7 +2,7 @@
//!
//! These are standalone helpers not tied to any specific client or service.
use aether_http::{build_http_client, HttpClientConfig};
use reqwest::Client;
use tracing::{debug, info};
/// Auto-detect public IP by querying external services.
@@ -13,11 +13,9 @@ pub async fn detect_public_ip() -> anyhow::Result<String> {
"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()
})?;
let client = Client::builder()
.timeout(std::time::Duration::from_secs(5))
.build()?;
for endpoint in &endpoints {
match client.get(*endpoint).send().await {
@@ -50,11 +48,9 @@ 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()
})
let client = Client::builder()
.timeout(std::time::Duration::from_secs(5))
.build()
.ok()?;
// Try ipinfo.io (HTTPS, returns plain text country code)
+257
View File
@@ -0,0 +1,257 @@
use std::time::{Duration, SystemTime, UNIX_EPOCH};
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 proxy node lifecycle management.
pub struct AetherClient {
http: Client,
base_url: String,
token: String,
retry_max_attempts: u32,
retry_base_delay: Duration,
retry_max_delay: Duration,
}
impl AetherClient {
pub fn new(config: &Config, aether_url: &str, management_token: &str) -> Self {
let mut builder = Client::builder()
.timeout(Duration::from_secs(config.aether_request_timeout_secs))
.connect_timeout(Duration::from_secs(config.aether_connect_timeout_secs))
.pool_max_idle_per_host(config.aether_pool_max_idle_per_host)
.pool_idle_timeout(Duration::from_secs(config.aether_pool_idle_timeout_secs))
.tcp_nodelay(config.aether_tcp_nodelay);
if config.aether_tcp_keepalive_secs > 0 {
builder =
builder.tcp_keepalive(Some(Duration::from_secs(config.aether_tcp_keepalive_secs)));
} else {
builder = builder.tcp_keepalive(None);
}
if config.aether_http2 {
builder = builder.http2_adaptive_window(true);
}
let http = builder.build().expect("failed to create HTTP client");
let retry_base_delay = Duration::from_millis(config.aether_retry_base_delay_ms);
let retry_max_delay =
Duration::from_millis(config.aether_retry_max_delay_ms).max(retry_base_delay);
Self {
http,
base_url: aether_url.trim_end_matches('/').to_string(),
token: management_token.to_string(),
retry_max_attempts: config.aether_retry_max_attempts.max(1),
retry_base_delay,
retry_max_delay,
}
}
/// 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();
anyhow::bail!("register failed (HTTP {}): {}", status, text);
}
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 text = r.text().await.unwrap_or_default();
error!(body = %text, "unregister failed");
anyhow::bail!("unregister failed: {}", text);
}
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;
let mut delay = self.retry_base_delay;
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 = jitter_delay(delay);
debug!(
attempt,
status = %resp.status(),
sleep_ms = sleep_for.as_millis(),
label,
"Aether request retrying"
);
sleep(sleep_for).await;
let next_delay = delay.checked_mul(2).unwrap_or(self.retry_max_delay);
delay = std::cmp::min(next_delay, self.retry_max_delay);
continue;
}
return Ok(resp);
}
Err(e) => {
if attempt < self.retry_max_attempts {
let sleep_for = jitter_delay(delay);
debug!(
attempt,
error = %e,
sleep_ms = sleep_for.as_millis(),
label,
"Aether request retrying"
);
sleep(sleep_for).await;
let next_delay = delay.checked_mul(2).unwrap_or(self.retry_max_delay);
delay = std::cmp::min(next_delay, self.retry_max_delay);
continue;
}
return Err(e);
}
}
}
}
}
fn should_retry_status(status: StatusCode) -> bool {
status.is_server_error()
|| status == StatusCode::TOO_MANY_REQUESTS
|| status == StatusCode::REQUEST_TIMEOUT
}
fn jitter_delay(base: Duration) -> Duration {
if base.is_zero() {
return base;
}
let nanos = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.subsec_nanos() as u64)
.unwrap_or(0);
let jitter_ms = nanos % 100;
base + Duration::from_millis(jitter_ms)
}
+256
View File
@@ -0,0 +1,256 @@
//! Systemd service installation for aether-proxy.
//!
//! Called from the setup TUI when the user enables "Install Service".
//! The unit file points to the binary and config at their current
//! absolute paths -- no files are copied.
use std::path::Path;
use std::process::Command;
const UNIT_PATH: &str = "/etc/systemd/system/aether-proxy.service";
const SERVICE_NAME: &str = "aether-proxy";
/// Whether systemd service installation is possible (systemd present + root).
pub fn is_available() -> bool {
is_systemd_available() && is_root()
}
/// Install aether-proxy as a systemd service. Must be run as root.
pub fn install_service(config_path: &Path) -> anyhow::Result<()> {
if !is_systemd_available() {
anyhow::bail!("systemd not available");
}
if !is_root() {
anyhow::bail!("root required, use: sudo ./aether-proxy setup");
}
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("/");
// Stop existing service if running (ignore errors)
if Path::new(UNIT_PATH).exists() {
eprintln!(" Stopping existing service...");
let _ = Command::new("systemctl")
.args(["stop", SERVICE_NAME])
.status();
}
// Write unit file
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 Proxy\n\
After=network.target\n\
\n\
[Service]\n\
Type=simple\n\
WorkingDirectory={working_dir}\n\
Environment=AETHER_PROXY_CONFIG={config_str}\n\
ExecStart={exe_str}\n\
Restart=on-failure\n\
RestartSec=5\n\
LimitNOFILE=65535\n\
UMask=0077\n\
\n\
[Install]\n\
WantedBy=multi-user.target\n",
);
std::fs::write(UNIT_PATH, &unit_content)?;
// Reload and enable
eprintln!(" Enabling and starting service...");
run_cmd("systemctl", &["daemon-reload"])?;
run_cmd("systemctl", &["enable", "--now", SERVICE_NAME])?;
// Verify
eprintln!();
let output = Command::new("systemctl")
.args(["is-active", SERVICE_NAME])
.output()?;
let state = String::from_utf8_lossy(&output.stdout).trim().to_string();
if state == "active" {
eprintln!(" Service started successfully!");
} else {
eprintln!(" Service state: {} (check logs)", state);
}
eprintln!();
eprintln!(" Commands:");
eprintln!(" ./aether-proxy status # service status");
eprintln!(" ./aether-proxy logs # tail logs");
eprintln!(" sudo ./aether-proxy restart # restart");
eprintln!(" sudo ./aether-proxy stop # stop");
eprintln!(" sudo ./aether-proxy uninstall # remove service");
eprintln!();
Ok(())
}
fn is_systemd_available() -> bool {
Command::new("systemctl")
.arg("--version")
.stdout(std::process::Stdio::null())
.stderr(std::process::Stdio::null())
.status()
.map(|s| s.success())
.unwrap_or(false)
}
pub(crate) fn is_root() -> bool {
#[cfg(unix)]
{
unsafe { libc::geteuid() == 0 }
}
#[cfg(not(unix))]
{
false
}
}
/// Whether a systemd unit file is currently installed.
pub fn is_installed() -> bool {
Path::new(UNIT_PATH).exists()
}
/// Remove the systemd service (called from setup TUI when Install Service is toggled off).
pub fn uninstall_service() -> anyhow::Result<()> {
if !Path::new(UNIT_PATH).exists() {
return Ok(());
}
eprintln!(" Stopping and removing existing service...");
let _ = Command::new("systemctl")
.args(["disable", "--now", SERVICE_NAME])
.status();
std::fs::remove_file(UNIT_PATH)?;
eprintln!(" Removed {}", UNIT_PATH);
run_cmd("systemctl", &["daemon-reload"])?;
eprintln!(" Service uninstalled.");
eprintln!();
Ok(())
}
/// Check if the systemd service is currently active.
pub fn is_service_active() -> bool {
std::path::Path::new(UNIT_PATH).exists()
&& Command::new("systemctl")
.args(["is-active", "--quiet", SERVICE_NAME])
.stdout(std::process::Stdio::null())
.stderr(std::process::Stdio::null())
.status()
.map(|s| s.success())
.unwrap_or(false)
}
// ── CLI subcommands (systemd wrappers) ──────────────────────────────────────
fn ensure_service_installed() -> anyhow::Result<()> {
if !std::path::Path::new(UNIT_PATH).exists() {
anyhow::bail!("service not installed, run `sudo ./aether-proxy setup` first");
}
Ok(())
}
fn ensure_root_and_service() -> anyhow::Result<()> {
ensure_service_installed()?;
if !is_root() {
anyhow::bail!("root required, use: sudo ./aether-proxy <command>");
}
Ok(())
}
/// `aether-proxy status` -- show service status.
pub fn cmd_status() -> anyhow::Result<()> {
ensure_service_installed()?;
let status = Command::new("systemctl")
.args(["status", SERVICE_NAME])
.status()?;
// systemctl status returns non-zero when inactive; that's fine
std::process::exit(status.code().unwrap_or(1));
}
/// `aether-proxy logs` -- tail service logs.
pub fn cmd_logs() -> anyhow::Result<()> {
ensure_service_installed()?;
let status = Command::new("journalctl")
.args(["-u", SERVICE_NAME, "-f", "--no-pager", "-n", "100"])
.status()?;
std::process::exit(status.code().unwrap_or(1));
}
/// `aether-proxy start` -- start the service.
pub fn cmd_start() -> anyhow::Result<()> {
ensure_root_and_service()?;
run_cmd("systemctl", &["start", SERVICE_NAME])?;
eprintln!(" Service started.");
Ok(())
}
/// `aether-proxy restart` -- restart the service.
pub fn cmd_restart() -> anyhow::Result<()> {
ensure_root_and_service()?;
run_cmd("systemctl", &["restart", SERVICE_NAME])?;
eprintln!(" Service restarted.");
Ok(())
}
/// `aether-proxy stop` -- stop the service.
pub fn cmd_stop() -> anyhow::Result<()> {
ensure_root_and_service()?;
run_cmd("systemctl", &["stop", SERVICE_NAME])?;
eprintln!(" Service stopped.");
Ok(())
}
/// `aether-proxy uninstall` -- disable and remove the systemd service.
pub fn cmd_uninstall() -> anyhow::Result<()> {
ensure_root_and_service()?;
eprintln!(" Stopping and disabling service...");
let _ = Command::new("systemctl")
.args(["disable", "--now", SERVICE_NAME])
.status();
if std::path::Path::new(UNIT_PATH).exists() {
std::fs::remove_file(UNIT_PATH)?;
eprintln!(" Removed {}", UNIT_PATH);
}
run_cmd("systemctl", &["daemon-reload"])?;
eprintln!(" Service uninstalled.");
eprintln!();
eprintln!(" Config file and TLS certs 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(())
}
+868
View File
@@ -0,0 +1,868 @@
//! Interactive TUI for configuring aether-proxy.
//!
//! Launched via `aether-proxy setup [path]`. Presents a full-screen form
//! backed by ratatui where the user can navigate fields, edit values, and
//! save to a TOML config file. Supports multi-server configuration via
//! a tabbed interface.
use std::io;
use std::path::PathBuf;
use std::time::{Duration, Instant};
use crossterm::event::{self, Event, KeyCode, KeyEvent, KeyEventKind, KeyModifiers};
use crossterm::execute;
use crossterm::terminal::{self, EnterAlternateScreen, LeaveAlternateScreen};
use ratatui::backend::CrosstermBackend;
use ratatui::layout::{Constraint, Layout, Rect};
use ratatui::style::{Color, Modifier, Style};
use ratatui::text::{Line, Span};
use ratatui::widgets::{Block, Borders, Paragraph};
use ratatui::Frame;
use ratatui::Terminal;
use crate::config::{ConfigFile, ServerEntry};
/// Outcome of the setup wizard, returned to the caller.
pub enum SetupOutcome {
/// Config saved; systemd service installed and started.
ServiceInstalled,
/// Config saved; no service -- caller should start the proxy directly.
ReadyToRun(PathBuf),
/// User quit without saving.
Cancelled,
}
/// Column width reserved for the field label (chars).
const LABEL_WIDTH: usize = 22;
// -- Field types --------------------------------------------------------------
#[derive(Clone, Copy, PartialEq)]
enum FieldKind {
Text,
Secret,
Bool,
LogLevel,
}
struct Field {
label: &'static str,
key: &'static str,
value: String,
kind: FieldKind,
required: bool,
help: &'static str,
}
// -- Server tab ---------------------------------------------------------------
/// A single server tab's editable fields.
struct ServerTab {
fields: Vec<Field>,
}
impl ServerTab {
fn new() -> Self {
Self {
fields: vec![
Field {
label: "Aether URL",
key: "aether_url",
value: String::new(),
kind: FieldKind::Text,
required: true,
help: "Aether URL (e.g. https://aether.example.com)",
},
Field {
label: "Management Token",
key: "management_token",
value: String::new(),
kind: FieldKind::Secret,
required: true,
help: "Aether Management Token (ae_xxx)",
},
Field {
label: "Node Name",
key: "node_name",
value: "proxy-01".into(),
kind: FieldKind::Text,
required: true,
help: "Node name for identification in Aether dashboard",
},
],
}
}
fn from_entry(entry: &ServerEntry) -> Self {
let mut tab = Self::new();
tab.fields[0].value = entry.aether_url.clone();
tab.fields[1].value = entry.management_token.clone();
if let Some(ref name) = entry.node_name {
tab.fields[2].value = name.clone();
}
tab
}
}
// -- App state ----------------------------------------------------------------
#[derive(PartialEq)]
enum Mode {
Normal,
Editing,
}
struct App {
server_tabs: Vec<ServerTab>,
active_tab: usize,
global_fields: Vec<Field>,
selected: usize,
mode: Mode,
edit_buffer: String,
edit_cursor: usize,
config_path: PathBuf,
modified: bool,
message: Option<(String, Instant, bool)>,
scroll_offset: usize,
saved_once: bool,
pending_quit: bool,
confirm_delete: bool,
}
impl App {
fn new(config_path: PathBuf) -> Self {
Self {
server_tabs: vec![ServerTab::new()],
active_tab: 0,
global_fields: vec![
Field {
label: "Log Level",
key: "log_level",
value: "info".into(),
kind: FieldKind::LogLevel,
required: true,
help: "Log level -- Enter to cycle: trace / debug / info / warn / error",
},
Field {
label: "Log JSON",
key: "log_json",
value: "false".into(),
kind: FieldKind::Bool,
required: true,
help: "Output logs as JSON -- Enter to toggle",
},
Field {
label: "Install Service",
key: "install_service",
value: if super::service::is_available() {
"true"
} else {
"false"
}
.into(),
kind: FieldKind::Bool,
required: true,
help: "Install as systemd service (requires root) -- Enter to toggle",
},
],
selected: 0,
mode: Mode::Normal,
edit_buffer: String::new(),
edit_cursor: 0,
config_path,
modified: false,
message: None,
scroll_offset: 0,
saved_once: false,
pending_quit: false,
confirm_delete: false,
}
}
// -- Field accessors (unified index across server + global) ---------------
fn server_field_count(&self) -> usize {
self.server_tabs[self.active_tab].fields.len()
}
fn total_field_count(&self) -> usize {
self.server_field_count() + self.global_fields.len()
}
fn selected_field(&self) -> &Field {
let sc = self.server_field_count();
if self.selected < sc {
&self.server_tabs[self.active_tab].fields[self.selected]
} else {
&self.global_fields[self.selected - sc]
}
}
fn selected_field_mut(&mut self) -> &mut Field {
let sc = self.server_field_count();
if self.selected < sc {
&mut self.server_tabs[self.active_tab].fields[self.selected]
} else {
&mut self.global_fields[self.selected - sc]
}
}
fn clamp_selection(&mut self) {
let max = self.total_field_count();
if self.selected >= max {
self.selected = max.saturating_sub(1);
}
self.scroll_offset = 0;
self.confirm_delete = false;
}
// -- Config <-> fields -----------------------------------------------------
fn load_from_file(&mut self) {
if let Ok(cfg) = ConfigFile::load(&self.config_path) {
self.apply_config(&cfg);
}
}
fn apply_config(&mut self, cfg: &ConfigFile) {
// Global fields
for field in &mut self.global_fields {
let val: Option<String> = match field.key {
"log_level" => cfg.log_level.clone(),
"log_json" => cfg.log_json.map(|v| v.to_string()),
_ => None,
};
if let Some(v) = val {
field.value = v;
}
}
// Server tabs
let servers = cfg.effective_servers();
if servers.is_empty() {
let mut tab = ServerTab::new();
// Single-server fallback: use top-level node_name
if let Some(ref name) = cfg.node_name {
tab.fields[2].value = name.clone();
}
self.server_tabs = vec![tab];
} else {
self.server_tabs = servers.iter().map(ServerTab::from_entry).collect();
// For single-server mode, node_name might be in top-level only
if self.server_tabs.len() == 1 && self.server_tabs[0].fields[2].value.is_empty() {
if let Some(ref name) = cfg.node_name {
self.server_tabs[0].fields[2].value = name.clone();
}
}
}
self.active_tab = 0;
self.selected = 0;
self.scroll_offset = 0;
}
fn to_config(&self) -> ConfigFile {
let get_global = |key: &str| -> Option<String> {
self.global_fields
.iter()
.find(|f| f.key == key)
.map(|f| f.value.clone())
.filter(|v| !v.is_empty())
};
let get_tab = |tab: &ServerTab, key: &str| -> Option<String> {
tab.fields
.iter()
.find(|f| f.key == key)
.map(|f| f.value.clone())
.filter(|v| !v.is_empty())
};
let mut cfg = ConfigFile {
log_level: get_global("log_level"),
log_json: get_global("log_json").and_then(|v| v.parse().ok()),
..ConfigFile::default()
};
// Always write [[servers]] format; old top-level fields are read-only compat
cfg.servers = self
.server_tabs
.iter()
.map(|tab| ServerEntry {
aether_url: get_tab(tab, "aether_url").unwrap_or_default(),
management_token: get_tab(tab, "management_token").unwrap_or_default(),
node_name: get_tab(tab, "node_name"),
})
.collect();
cfg
}
fn save(&mut self) -> anyhow::Result<()> {
let cfg = self.to_config();
cfg.save(&self.config_path)?;
// Restrict config file permissions to owner-only (contains management token).
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
let _ =
std::fs::set_permissions(&self.config_path, std::fs::Permissions::from_mode(0o600));
}
self.modified = false;
self.saved_once = true;
self.message = Some((
format!("saved to {}", self.config_path.display()),
Instant::now(),
false,
));
Ok(())
}
// -- Scrolling ---------------------------------------------------------------
fn ensure_visible(&mut self, visible_rows: usize) {
if visible_rows == 0 {
return;
}
// Account for separator line between server and global fields
let display_row = if self.selected >= self.server_field_count() {
self.selected + 1
} else {
self.selected
};
if display_row < self.scroll_offset {
self.scroll_offset = display_row;
} else if display_row >= self.scroll_offset + visible_rows {
self.scroll_offset = display_row - visible_rows + 1;
}
}
// -- Key handling -------------------------------------------------------------
/// Returns `true` when the app should exit.
fn handle_key(&mut self, key: KeyEvent) -> bool {
// Expire old messages (but keep quit-confirmation messages alive)
if let Some((_, when, _)) = &self.message {
if !self.pending_quit && !self.confirm_delete && when.elapsed() > Duration::from_secs(4)
{
self.message = None;
}
}
match self.mode {
Mode::Normal => self.handle_normal(key),
Mode::Editing => {
self.handle_edit(key);
false
}
}
}
fn handle_normal(&mut self, key: KeyEvent) -> bool {
// -- Quit handling (with unsaved-changes confirmation) -----------------
let is_quit_key = matches!(key.code, KeyCode::Char('q') | KeyCode::Esc);
if is_quit_key {
if !self.modified || self.pending_quit {
return true;
}
self.pending_quit = true;
self.confirm_delete = false;
self.message = Some((
"unsaved changes! q again to discard, ^S to save".into(),
Instant::now(),
true,
));
return false;
}
// Any other key cancels pending quit / pending delete
if self.pending_quit {
self.pending_quit = false;
self.message = None;
}
if self.confirm_delete && !matches!(key.code, KeyCode::Delete | KeyCode::Char('x')) {
self.confirm_delete = false;
self.message = None;
}
match key.code {
KeyCode::Char('s')
if key.modifiers.contains(KeyModifiers::CONTROL)
|| key.modifiers.contains(KeyModifiers::SUPER) =>
{
if let Err(e) = self.save() {
self.message = Some((format!("error: {}", e), Instant::now(), true));
}
}
KeyCode::Up | KeyCode::Char('k') => {
self.selected = self.selected.saturating_sub(1);
}
KeyCode::Down | KeyCode::Char('j') => {
if self.selected + 1 < self.total_field_count() {
self.selected += 1;
}
}
KeyCode::Home => self.selected = 0,
KeyCode::End => self.selected = self.total_field_count() - 1,
KeyCode::Enter | KeyCode::Char(' ') => {
let kind = self.selected_field().kind;
let key_str = self.selected_field().key;
let value = self.selected_field().value.clone();
match kind {
FieldKind::Bool => {
let toggled = if value == "true" { "false" } else { "true" };
if key_str == "install_service"
&& toggled == "true"
&& !super::service::is_available()
{
self.message = Some((
"requires root with systemd, use: sudo aether-proxy setup".into(),
Instant::now(),
true,
));
} else {
self.selected_field_mut().value = toggled.into();
self.modified = true;
}
}
FieldKind::LogLevel => {
const LEVELS: &[&str] = &["trace", "debug", "info", "warn", "error"];
let idx = LEVELS.iter().position(|l| *l == value).unwrap_or(2);
self.selected_field_mut().value = LEVELS[(idx + 1) % LEVELS.len()].into();
self.modified = true;
}
_ => {
self.edit_buffer = value;
self.edit_cursor = self.edit_buffer.chars().count();
self.mode = Mode::Editing;
}
}
}
// -- Tab navigation --
KeyCode::Tab => {
if self.server_tabs.len() > 1 {
self.active_tab = (self.active_tab + 1) % self.server_tabs.len();
self.clamp_selection();
}
}
KeyCode::BackTab => {
if self.server_tabs.len() > 1 {
self.active_tab = if self.active_tab == 0 {
self.server_tabs.len() - 1
} else {
self.active_tab - 1
};
self.clamp_selection();
}
}
KeyCode::Char(c @ '1'..='9') if !key.modifiers.contains(KeyModifiers::CONTROL) => {
let idx = (c as usize) - ('1' as usize);
if idx < self.server_tabs.len() && idx != self.active_tab {
self.active_tab = idx;
self.clamp_selection();
}
}
// -- Add / remove server --
KeyCode::Char('+') | KeyCode::Char('a') => {
self.server_tabs.push(ServerTab::new());
self.active_tab = self.server_tabs.len() - 1;
self.selected = 0;
self.scroll_offset = 0;
self.modified = true;
self.message = Some((
format!("added server {}", self.server_tabs.len()),
Instant::now(),
false,
));
}
KeyCode::Delete | KeyCode::Char('x') => {
if self.server_tabs.len() <= 1 {
self.message =
Some(("cannot remove the last server".into(), Instant::now(), true));
} else if self.confirm_delete {
let removed = self.active_tab + 1;
self.server_tabs.remove(self.active_tab);
self.active_tab = self.active_tab.min(self.server_tabs.len() - 1);
self.clamp_selection();
self.modified = true;
self.message =
Some((format!("server {} removed", removed), Instant::now(), false));
} else {
self.confirm_delete = true;
self.message = Some((
"press Delete/x again to remove this server".into(),
Instant::now(),
true,
));
}
}
_ => {}
}
false
}
fn handle_edit(&mut self, key: KeyEvent) {
match key.code {
KeyCode::Esc => {
self.mode = Mode::Normal;
}
KeyCode::Enter => {
if self.validate_edit() {
self.selected_field_mut().value = self.edit_buffer.clone();
self.modified = true;
self.mode = Mode::Normal;
} else {
self.message = Some(("invalid format".into(), Instant::now(), true));
}
}
KeyCode::Backspace => {
if self.edit_cursor > 0 {
self.edit_cursor -= 1;
let byte = self.char_byte_pos(self.edit_cursor);
self.edit_buffer.remove(byte);
}
}
KeyCode::Delete => {
if self.edit_cursor < self.edit_buffer.chars().count() {
let byte = self.char_byte_pos(self.edit_cursor);
self.edit_buffer.remove(byte);
}
}
KeyCode::Left => {
self.edit_cursor = self.edit_cursor.saturating_sub(1);
}
KeyCode::Right => {
let len = self.edit_buffer.chars().count();
if self.edit_cursor < len {
self.edit_cursor += 1;
}
}
KeyCode::Home => self.edit_cursor = 0,
KeyCode::End => self.edit_cursor = self.edit_buffer.chars().count(),
KeyCode::Char(c) => {
let byte = self.char_byte_pos(self.edit_cursor);
self.edit_buffer.insert(byte, c);
self.edit_cursor += 1;
}
_ => {}
}
}
fn validate_edit(&self) -> bool {
true
}
/// Byte offset of the char at `char_idx`.
fn char_byte_pos(&self, char_idx: usize) -> usize {
self.edit_buffer
.char_indices()
.nth(char_idx)
.map(|(i, _)| i)
.unwrap_or(self.edit_buffer.len())
}
}
// -- Rendering ----------------------------------------------------------------
fn ui(f: &mut Frame, app: &mut App) {
let area = f.area();
let title = if app.modified {
" Aether Proxy Setup [*] "
} else {
" Aether Proxy Setup "
};
let outer = Block::default()
.borders(Borders::ALL)
.title(title)
.title_alignment(ratatui::layout::Alignment::Center)
.border_style(Style::default().fg(Color::Cyan));
let inner = outer.inner(area);
f.render_widget(outer, area);
// Split: fields | tab bar | footer
let chunks = Layout::vertical([
Constraint::Min(1),
Constraint::Length(1),
Constraint::Length(4),
])
.split(inner);
render_fields(f, app, chunks[0]);
render_tab_bar(f, app, chunks[1]);
render_footer(f, app, chunks[2]);
}
fn render_fields(f: &mut Frame, app: &mut App, area: Rect) {
let visible = area.height as usize;
app.ensure_visible(visible);
let server_count = app.server_field_count();
let mut lines: Vec<Line> = Vec::new();
// display_row tracks the actual row index (including separator)
let mut display_row: usize = 0;
// Server fields
for i in 0..server_count {
if display_row >= app.scroll_offset && display_row < app.scroll_offset + visible {
lines.push(build_field_line(app, i, display_row));
}
display_row += 1;
}
// Separator line
if display_row >= app.scroll_offset && display_row < app.scroll_offset + visible {
lines.push(Line::from(Span::styled(
" ----------------------------------------",
Style::default().fg(Color::DarkGray),
)));
}
display_row += 1;
// Global fields
for i in 0..app.global_fields.len() {
let field_idx = server_count + i;
if display_row >= app.scroll_offset && display_row < app.scroll_offset + visible {
lines.push(build_field_line(app, field_idx, display_row));
}
display_row += 1;
}
let paragraph = Paragraph::new(lines);
f.render_widget(paragraph, area);
// Cursor position while editing
if app.mode == Mode::Editing {
let sel_display_row = if app.selected >= server_count {
app.selected + 1
} else {
app.selected
};
let row_in_view = sel_display_row.saturating_sub(app.scroll_offset);
let prefix: u16 = 3 + LABEL_WIDTH as u16 + 2;
let cx = area.x + prefix + app.edit_cursor as u16;
let cy = area.y + row_in_view as u16;
if cx < area.x + area.width && cy < area.y + area.height {
f.set_cursor_position((cx, cy));
}
}
}
fn build_field_line(app: &App, field_idx: usize, _display_row: usize) -> Line<'static> {
let sc = app.server_field_count();
let field = if field_idx < sc {
&app.server_tabs[app.active_tab].fields[field_idx]
} else {
&app.global_fields[field_idx - sc]
};
let selected = field_idx == app.selected;
let indicator = if selected { " > " } else { " " };
let label_style = if selected {
Style::default()
.fg(Color::Cyan)
.add_modifier(Modifier::BOLD)
} else {
Style::default().fg(Color::DarkGray)
};
let padded_label = format!("{:<width$}", field.label, width = LABEL_WIDTH);
let (value_text, value_style) = if app.mode == Mode::Editing && selected {
(app.edit_buffer.clone(), Style::default().fg(Color::Yellow))
} else {
field_display(field)
};
Line::from(vec![
Span::styled(indicator.to_string(), label_style),
Span::styled(padded_label, label_style),
Span::raw(" "),
Span::styled(value_text, value_style),
])
}
/// Returns (display_text, style) for a field in normal mode.
fn field_display(field: &Field) -> (String, Style) {
if field.value.is_empty() {
let text = if field.required {
"(required)".into()
} else {
"-".into()
};
let color = if field.required {
Color::Red
} else {
Color::DarkGray
};
return (text, Style::default().fg(color));
}
match field.kind {
FieldKind::Secret => (
"*".repeat(field.value.len().min(20)),
Style::default().fg(Color::White),
),
FieldKind::Bool => {
if field.value == "true" {
("[x] on".into(), Style::default().fg(Color::Green))
} else {
("[ ] off".into(), Style::default().fg(Color::DarkGray))
}
}
FieldKind::LogLevel => {
let color = match field.value.as_str() {
"trace" => Color::Magenta,
"debug" => Color::Blue,
"info" => Color::Green,
"warn" => Color::Yellow,
"error" => Color::Red,
_ => Color::White,
};
(field.value.clone(), Style::default().fg(color))
}
_ => (field.value.clone(), Style::default().fg(Color::White)),
}
}
fn render_tab_bar(f: &mut Frame, app: &App, area: Rect) {
let mut spans: Vec<Span> = Vec::new();
spans.push(Span::raw(" "));
for (i, tab) in app.server_tabs.iter().enumerate() {
let num = i + 1;
let name = tab
.fields
.iter()
.find(|f| f.key == "node_name")
.filter(|f| !f.value.is_empty())
.map(|f| f.value.clone())
.unwrap_or_else(|| format!("Server {}", num));
let label = format!(" {} {} ", num, name);
if i == app.active_tab {
spans.push(Span::styled(
label,
Style::default()
.fg(Color::Black)
.bg(Color::Cyan)
.add_modifier(Modifier::BOLD),
));
} else {
spans.push(Span::styled(label, Style::default().fg(Color::DarkGray)));
}
spans.push(Span::raw(" "));
}
spans.push(Span::styled(" + Add ", Style::default().fg(Color::Green)));
f.render_widget(Paragraph::new(Line::from(spans)), area);
}
fn render_footer(f: &mut Frame, app: &App, area: Rect) {
let help = app.selected_field().help;
let keybindings = if app.mode == Mode::Editing {
"Enter confirm Esc cancel"
} else if app.server_tabs.len() > 1 {
"j/k select Enter edit Tab switch + add x remove ^S save q quit"
} else {
"j/k select Enter edit + add server ^S save q quit"
};
let mut status_spans: Vec<Span> = vec![Span::styled(
format!(" {}", keybindings),
Style::default().fg(Color::DarkGray),
)];
if let Some((msg, _, is_err)) = &app.message {
let color = if *is_err { Color::Red } else { Color::Green };
status_spans.push(Span::raw(" "));
status_spans.push(Span::styled(msg.clone(), Style::default().fg(color)));
}
let footer_text = vec![
Line::raw(""),
Line::from(Span::styled(
format!(" {}", help),
Style::default().fg(Color::DarkGray),
)),
Line::from(status_spans),
];
let footer = Paragraph::new(footer_text).block(
Block::default()
.borders(Borders::TOP)
.border_style(Style::default().fg(Color::DarkGray)),
);
f.render_widget(footer, area);
}
// -- Entry point --------------------------------------------------------------
pub fn run(config_path: PathBuf) -> anyhow::Result<SetupOutcome> {
terminal::enable_raw_mode()?;
let mut stdout = io::stdout();
execute!(stdout, EnterAlternateScreen)?;
let backend = CrosstermBackend::new(stdout);
let mut terminal = Terminal::new(backend)?;
let mut app = App::new(config_path.clone());
app.load_from_file();
let result = event_loop(&mut terminal, &mut app);
terminal::disable_raw_mode()?;
execute!(terminal.backend_mut(), LeaveAlternateScreen)?;
terminal.show_cursor()?;
result?;
// -- Post-TUI: decide outcome ---------------------------------------------
if !app.saved_once {
return Ok(SetupOutcome::Cancelled);
}
eprintln!();
eprintln!(" Config saved to {}", config_path.display());
eprintln!();
let wants_service = app
.global_fields
.iter()
.find(|f| f.key == "install_service")
.map(|f| f.value == "true")
.unwrap_or(false);
if wants_service {
match super::service::install_service(&config_path) {
Ok(()) => return Ok(SetupOutcome::ServiceInstalled),
Err(e) => {
eprintln!(" Service install failed: {}", e);
eprintln!(" Starting proxy directly instead.\n");
}
}
} else if super::service::is_installed() {
if let Err(e) = super::service::uninstall_service() {
eprintln!(" Service uninstall failed: {}", e);
eprintln!();
}
}
Ok(SetupOutcome::ReadyToRun(config_path))
}
fn event_loop(
terminal: &mut Terminal<CrosstermBackend<io::Stdout>>,
app: &mut App,
) -> anyhow::Result<()> {
loop {
terminal.draw(|f| ui(f, app))?;
if event::poll(Duration::from_millis(200))? {
if let Event::Key(key) = event::read()? {
if key.kind == KeyEventKind::Press && app.handle_key(key) {
break;
}
}
}
}
Ok(())
}
@@ -1,12 +1,10 @@
//! Self-upgrade support for `aether-tunnel`.
//! Self-upgrade for aether-proxy.
//!
//! Downloads a release from GitHub, verifies the SHA256 checksum, replaces the
//! running binary atomically, and restarts the active managed service when
//! applicable.
//! Downloads a release from GitHub, verifies SHA256 checksum, and atomically
//! replaces the running binary. Restarts the systemd service if active.
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";
@@ -24,14 +22,7 @@ struct GithubRelease {
// ── 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") {
if cfg!(target_os = "linux") && cfg!(target_arch = "x86_64") {
"linux-amd64"
} else if cfg!(target_os = "linux") && cfg!(target_arch = "aarch64") {
"linux-arm64"
@@ -65,14 +56,10 @@ fn build_github_client() -> anyhow::Result<reqwest::Client> {
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()
},
)
Ok(reqwest::Client::builder()
.timeout(std::time::Duration::from_secs(300))
.user_agent(format!("aether-proxy/{}", CURRENT_VERSION))
.default_headers(headers)
.build()?)
}
@@ -84,11 +71,11 @@ async fn fetch_release(
) -> 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") {
// Accept both "proxy-v0.2.0" and bare "0.2.0"
let tag = if ver.starts_with("proxy-v") {
ver.to_string()
} else {
format!("tunnel-v{}", ver)
format!("proxy-v{}", ver)
};
let url = format!(
"{}/repos/{}/releases/tags/{}",
@@ -103,7 +90,7 @@ async fn fetch_release(
Ok(resp.json().await?)
}
None => {
// List releases and find the latest tunnel-v* tag
// List releases and find the latest proxy-v* tag
let url = format!(
"{}/repos/{}/releases?per_page=20",
GITHUB_API_BASE, GITHUB_REPO
@@ -117,8 +104,8 @@ async fn fetch_release(
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"))
.find(|r| r.tag_name.starts_with("proxy-v"))
.ok_or_else(|| anyhow::anyhow!("no proxy-v* release found"))
}
}
}
@@ -171,7 +158,7 @@ async fn download_and_verify(
platform: &str,
dest: &Path,
) -> anyhow::Result<()> {
let archive_name = format!("aether-tunnel-{}.tar.gz", platform);
let archive_name = format!("aether-proxy-{}.tar.gz", platform);
eprintln!(" Downloading {}...", archive_name);
let (archive_bytes, checksum_bytes) = tokio::try_join!(
@@ -220,9 +207,9 @@ fn extract_binary(archive_bytes: &[u8], dest: &Path) -> anyhow::Result<()> {
let mut archive = Archive::new(decoder);
let binary_name = if cfg!(target_os = "windows") {
"aether-tunnel.exe"
"aether-proxy.exe"
} else {
"aether-tunnel"
"aether-proxy"
};
for entry in archive.entries()? {
@@ -310,7 +297,7 @@ async fn execute_upgrade(
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");
let temp_path = exe_dir.join(".aether-proxy.upgrade.tmp");
if require_root {
if !super::service::is_root() {
@@ -318,14 +305,14 @@ async fn execute_upgrade(
}
} 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");
let test_path = exe_dir.join(".aether-proxy.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",
"no write access to {}. Use: sudo aether-proxy upgrade",
exe_dir.display()
);
}
@@ -339,10 +326,7 @@ async fn execute_upgrade(
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);
let target_semver = target_tag.strip_prefix("proxy-v").unwrap_or(target_tag);
eprintln!(" Target version: {} ({})", target_tag, release.name);
@@ -372,33 +356,34 @@ async fn execute_upgrade(
match restart_mode {
RestartMode::BestEffort => {
// Restart systemd service if running.
// 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() {
eprintln!(" Restarting systemd service...");
match super::service::run_cmd("systemctl", &["restart", "aether-proxy"]) {
Ok(()) => eprintln!(" Service restarted."),
Err(e) => {
eprintln!(" WARNING: failed to restart service: {}", e);
eprintln!(" Run manually: sudo aether-tunnel restart");
eprintln!(" Run manually: sudo systemctl restart aether-proxy");
}
}
} else {
eprintln!(" Managed service is active, but restart requires root.");
eprintln!(" Run: sudo aether-tunnel restart");
eprintln!(" Systemd service is active, but restart requires root.");
eprintln!(" Run: sudo systemctl restart aether-proxy");
eprintln!(" Skipping restart.");
}
} else {
eprintln!(" No active service detected, skipping restart.");
eprintln!(" No active systemd 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!(" Restarting systemd service...");
super::service::run_cmd("systemctl", &["restart", "aether-proxy"])?;
eprintln!(" Service restarted.");
}
}
@@ -412,15 +397,15 @@ async fn execute_upgrade(
Ok(())
}
/// `aether-tunnel upgrade [version]` -- self-upgrade from GitHub releases.
/// `aether-proxy 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.
/// This path is designed for server-pushed upgrades in systemd/root scenarios:
/// it requires root and requires a successful `systemctl restart aether-proxy`.
pub async fn perform_upgrade(version: &str) -> anyhow::Result<()> {
execute_upgrade(Some(version), true, RestartMode::Required).await
}
+77
View File
@@ -0,0 +1,77 @@
//! Shared application state passed to all subsystems.
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, RwLock};
use std::time::Duration;
use crate::config::Config;
use crate::registration::client::AetherClient;
use crate::runtime::SharedDynamicConfig;
use crate::target_filter::DnsCache;
use crate::upstream_client::UpstreamClient;
/// 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>,
/// Hyper client for tunnel upstream requests with validated DNS and connection timing.
pub upstream_client: UpstreamClient,
/// Shared TLS config for tunnel WebSocket connections (avoids re-parsing root CAs on each reconnect).
pub tunnel_tls_config: Arc<rustls::ClientConfig>,
}
/// 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<ProxyMetrics>,
}
/// Aggregate metrics for reporting to Aether.
pub struct ProxyMetrics {
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,
}
impl ProxyMetrics {
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),
}
}
/// 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);
}
}
@@ -210,15 +210,13 @@ impl DnsCache {
}
}
/// Resolve a hostname to validated socket addresses.
/// Resolve a hostname to public (non-private) 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.
/// Results are cached in `dns_cache`. Private/reserved IPs are filtered out.
/// Returns an error if no public addresses remain after filtering.
pub async fn resolve_public_addrs(
host: &str,
port: u16,
allow_private: bool,
dns_cache: &DnsCache,
) -> Result<Vec<SocketAddr>, FilterError> {
// Cache hit
@@ -237,15 +235,11 @@ pub async fn resolve_public_addrs(
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
// Filter out private/reserved addresses
let public: Vec<SocketAddr> = resolved
.into_iter()
.filter(|addr| !is_private_ip(&addr.ip()))
.collect()
};
.collect();
if public.is_empty() {
return Err(FilterError::NoPublicAddrs(host.to_string()));
@@ -266,7 +260,6 @@ 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
@@ -276,14 +269,14 @@ pub async fn validate_target(
// Try parsing as IP directly (no DNS needed)
if let Ok(ip) = host.parse::<IpAddr>() {
if !allow_private && is_private_ip(&ip) {
if 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
resolve_public_addrs(host, port, dns_cache).await
}
#[cfg(test)]
@@ -340,54 +333,27 @@ mod tests {
#[tokio::test]
async fn test_port_not_allowed() {
let cache = cache();
let result = validate_target("8.8.8.8", 22, &ports(), false, &cache).await;
let result = validate_target("8.8.8.8", 22, &ports(), &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;
let result = validate_target("127.0.0.1", 80, &ports(), &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;
let result = validate_target("8.8.8.8", 443, &ports(), &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();
+238
View File
@@ -0,0 +1,238 @@
//! WebSocket tunnel client: connect, authenticate, and run the tunnel.
use std::sync::Arc;
use std::time::Duration;
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::state::{AppState, ServerContext};
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>,
) -> Result<TunnelOutcome, anyhow::Error> {
let ws_url = build_tunnel_url(server);
info!(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))?,
);
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 (backward-compatible: old backends
// ignore this header).
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 = Duration::from_secs(state.config.tunnel_connect_timeout_secs);
let tcp_stream = tokio::time::timeout(connect_timeout, TcpStream::connect((host, port)))
.await
.map_err(|_| {
anyhow::anyhow!(
"tunnel TCP connect timeout ({}s)",
connect_timeout.as_secs()
)
})??;
// 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 = Duration::from_secs(state.config.tunnel_connect_timeout_secs);
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 ({}s)",
handshake_timeout.as_secs()
)
})??;
info!(
conn = conn_idx,
tcp_keepalive_secs = state.config.tunnel_tcp_keepalive_secs,
tcp_nodelay = state.config.tunnel_tcp_nodelay,
connect_timeout_secs = state.config.tunnel_connect_timeout_secs,
stale_timeout_secs = state.config.tunnel_stale_timeout_secs,
"tunnel connected"
);
// 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 ping_interval = Duration::from_secs(state.config.tunnel_ping_interval_secs);
let (frame_tx, mut writer_handle) = writer::spawn_writer(ws_sink, ping_interval);
// 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.config),
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) => {
match result {
Ok(()) => TunnelOutcome::Disconnected,
Err(e) => return 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");
} else {
warn!(error = %e, "writer task cancelled, triggering reconnect");
}
}
}
TunnelOutcome::Disconnected
}
_ = shutdown.changed() => {
debug!("shutdown during tunnel dispatch");
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);
// 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;
}
info!("tunnel disconnected");
Ok(outcome)
}
/// 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 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)
}
+249
View File
@@ -0,0 +1,249 @@
//! 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::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;
/// 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,
) -> 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 = Duration::from_secs(state.config.tunnel_stale_timeout_secs);
// Track last time we received any data to detect stale connections
let mut last_data_at = tokio::time::Instant::now();
let read_err = loop {
let msg_result = tokio::select! {
msg = ws_stream.next() => {
match msg {
Some(r) => r,
None => break None,
}
}
_ = tokio::time::sleep_until(last_data_at + stale_timeout) => {
warn!(
stale_secs = stale_timeout.as_secs(),
"tunnel connection stale, no data received"
);
break None;
}
};
let msg = match msg_result {
Ok(m) => m,
Err(e) => {
error!(error = %e, "WebSocket read error");
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) => Bytes::from(data),
Message::Ping(_) => continue,
Message::Pong(_) => continue,
Message::Close(_) => {
info!("received WebSocket close");
break None;
}
_ => continue,
};
let frame = match Frame::decode(data) {
Ok(f) => f,
Err(e) => {
warn!(error = %e, "failed to decode frame");
continue;
}
};
match frame.msg_type {
MsgType::RequestHeaders => {
// 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) {
let is_end = frame.is_end_stream();
let sid = frame.stream_id;
let _ = tx.send(frame).await;
if is_end {
streams.remove(&sid);
}
}
}
MsgType::StreamEnd | MsgType::StreamError => {
// Client-side cancellation or end
if let Some(tx) = streams.remove(&frame.stream_id) {
let _ = tx.send(frame).await;
}
}
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(()),
}
}
/// 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;
}
+340
View File
@@ -0,0 +1,340 @@
//! 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 tracing::{debug, info, warn};
use crate::config::Config;
use crate::registration::client::RemoteConfig;
use crate::runtime;
use crate::state::ServerContext;
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: Option<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
/// resetting shared atomic metrics via `swap(0)`.
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 {
requests: u64,
latency_ns: u64,
failed: u64,
dns_failures: u64,
stream_errors: u64,
}
/// Spawn the heartbeat task. Returns a handle for forwarding ACKs.
pub fn spawn(
_config: Arc<Config>,
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.
// Snapshot is only cleared after receiving an ACK, which avoids losing
// interval counters when ACK/frame delivery is temporarily unstable.
let mut pending: Option<(u64, HeartbeatSnapshot)> = None;
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 (heartbeat_id, snapshot) = if let Some((id, snap)) = pending {
(id, snap)
} else {
let snap = collect_snapshot(&server);
let id = next_heartbeat_id;
next_heartbeat_id = next_heartbeat_id.wrapping_add(1);
if next_heartbeat_id == 0 {
next_heartbeat_id = 1;
}
pending = Some((id, snap));
(id, snap)
};
let payload = build_heartbeat_payload(
&server,
&heartbeat_session_id,
heartbeat_id,
snapshot
);
let frame = Frame::control(MsgType::HeartbeatData, payload);
if frame_tx.send(frame).await.is_err() {
if let Some((_, snap)) = pending.take() {
restore_snapshot(&server, snap);
}
break; // Writer closed
}
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((pending_id, _)) = pending {
match ack_id {
Some(id) if id == pending_id => {
pending = None;
}
None => {
// Backward-compatible with servers that don't echo
// heartbeat_id in ACK payload yet.
pending = None;
}
_ => {}
}
}
maybe_trigger_upgrade(upgrade_to);
}
AckDecision::Ignore => {}
}
}
_ = shutdown.changed() => {
debug!("heartbeat task shutting down");
if let Some((_, snap)) = pending.take() {
restore_snapshot(&server, snap);
}
break;
}
}
}
});
HeartbeatHandle { ack_tx }
}
fn collect_snapshot(server: &ServerContext) -> HeartbeatSnapshot {
HeartbeatSnapshot {
requests: server.metrics.total_requests.swap(0, Ordering::AcqRel),
latency_ns: server.metrics.total_latency_ns.swap(0, Ordering::AcqRel),
failed: server.metrics.failed_requests.swap(0, Ordering::AcqRel),
dns_failures: server.metrics.dns_failures.swap(0, Ordering::AcqRel),
stream_errors: server.metrics.stream_errors.swap(0, Ordering::AcqRel),
}
}
fn restore_snapshot(server: &ServerContext, snap: HeartbeatSnapshot) {
if snap.requests > 0 {
server
.metrics
.total_requests
.fetch_add(snap.requests, Ordering::Release);
}
if snap.latency_ns > 0 {
server
.metrics
.total_latency_ns
.fetch_add(snap.latency_ns, Ordering::Release);
}
if snap.failed > 0 {
server
.metrics
.failed_requests
.fetch_add(snap.failed, Ordering::Release);
}
if snap.dns_failures > 0 {
server
.metrics
.dns_failures
.fetch_add(snap.dns_failures, Ordering::Release);
}
if snap.stream_errors > 0 {
server
.metrics
.stream_errors
.fetch_add(snap.stream_errors, Ordering::Release);
}
}
fn build_heartbeat_payload(
server: &ServerContext,
heartbeat_session_id: &str,
heartbeat_id: u64,
snapshot: HeartbeatSnapshot,
) -> Bytes {
let node_id = server.node_id.read().unwrap().clone();
let avg_latency_ms = if snapshot.requests > 0 {
Some(snapshot.latency_ns as f64 / snapshot.requests as f64 / 1_000_000.0)
} else {
None
};
let payload = serde_json::json!({
"node_id": node_id,
"heartbeat_session_id": heartbeat_session_id,
"heartbeat_id": heartbeat_id,
"active_connections": server.active_connections.load(Ordering::Acquire),
"total_requests": snapshot.requests,
"avg_latency_ms": avg_latency_ms,
"failed_requests": snapshot.failed,
"dns_failures": snapshot.dns_failures,
"stream_errors": snapshot.stream_errors,
"proxy_metadata": {
"version": CURRENT_VERSION,
},
});
Bytes::from(serde_json::to_vec(&payload).unwrap_or_default())
}
fn handle_ack(server: &ServerContext, payload: &[u8]) -> AckDecision {
if payload.is_empty() {
return AckDecision::Accept {
heartbeat_id: None,
upgrade_to: None,
};
}
#[derive(serde::Deserialize)]
struct AckPayload {
#[serde(default)]
remote_config: Option<RemoteConfig>,
#[serde(default)]
config_version: u64,
#[serde(default)]
heartbeat_id: Option<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");
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("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);
}
}
});
}
+236
View File
@@ -0,0 +1,236 @@
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::{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>,
) {
info!(server = %server.server_label, conn = conn_idx, "starting tunnel");
let reconnect_salt = compute_connection_salt(server, conn_idx);
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;
}
}
}
let mut consecutive_failures: u32 = 0;
loop {
let started_at = Instant::now();
match client::connect_and_run(state, server, conn_idx, &mut shutdown).await {
Ok(client::TunnelOutcome::Shutdown) => {
info!(server = %server.server_label, conn = conn_idx, "tunnel shut down gracefully");
return;
}
Ok(client::TunnelOutcome::Disconnected) => {
info!(server = %server.server_label, conn = conn_idx, "tunnel disconnected, reconnecting");
}
Err(e) => {
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;
}
// 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,
);
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;
}
}
}
}
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::time::Duration;
use super::{
compute_reconnect_cap_ms, compute_reconnect_delay, compute_startup_stagger,
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));
}
}
+260
View File
@@ -0,0 +1,260 @@
//! Binary frame protocol for WebSocket tunnel multiplexing.
//!
//! Frame layout (10-byte header + variable payload):
//! ```text
//! | stream_id (4B) | msg_type (1B) | flags (1B) | payload_len (4B) | payload (NB) |
//! ```
use bytes::{Buf, BufMut, Bytes, BytesMut};
pub const HEADER_SIZE: usize = 10;
/// Frame flags.
pub mod flags {
pub const END_STREAM: u8 = 0x01;
pub const GZIP_COMPRESSED: u8 = 0x02;
}
/// Message types for the tunnel protocol.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub enum MsgType {
RequestHeaders = 0x01,
RequestBody = 0x02,
ResponseHeaders = 0x03,
ResponseBody = 0x04,
StreamEnd = 0x05,
StreamError = 0x06,
Ping = 0x10,
Pong = 0x11,
GoAway = 0x12,
HeartbeatData = 0x13,
HeartbeatAck = 0x14,
}
impl MsgType {
pub fn from_u8(v: u8) -> Option<Self> {
match v {
0x01 => Some(Self::RequestHeaders),
0x02 => Some(Self::RequestBody),
0x03 => Some(Self::ResponseHeaders),
0x04 => Some(Self::ResponseBody),
0x05 => Some(Self::StreamEnd),
0x06 => Some(Self::StreamError),
0x10 => Some(Self::Ping),
0x11 => Some(Self::Pong),
0x12 => Some(Self::GoAway),
0x13 => Some(Self::HeartbeatData),
0x14 => Some(Self::HeartbeatAck),
_ => None,
}
}
}
/// A single multiplexed frame.
#[derive(Debug, Clone)]
pub struct Frame {
pub stream_id: u32,
pub msg_type: MsgType,
pub flags: u8,
pub payload: Bytes,
}
impl Frame {
pub fn new(stream_id: u32, msg_type: MsgType, flags: u8, payload: impl Into<Bytes>) -> Self {
Self {
stream_id,
msg_type,
flags,
payload: payload.into(),
}
}
/// Control frame (stream_id = 0).
pub fn control(msg_type: MsgType, payload: impl Into<Bytes>) -> Self {
Self::new(0, msg_type, 0, payload)
}
pub fn is_end_stream(&self) -> bool {
self.flags & flags::END_STREAM != 0
}
pub fn is_gzip(&self) -> bool {
self.flags & flags::GZIP_COMPRESSED != 0
}
/// Encode into a binary buffer.
pub fn encode(&self) -> Bytes {
let mut buf = BytesMut::with_capacity(HEADER_SIZE + self.payload.len());
buf.put_u32(self.stream_id);
buf.put_u8(self.msg_type as u8);
buf.put_u8(self.flags);
buf.put_u32(self.payload.len() as u32);
buf.put(self.payload.clone());
buf.freeze()
}
/// Decode from a binary buffer.
pub fn decode(mut data: Bytes) -> Result<Self, ProtocolError> {
if data.len() < HEADER_SIZE {
return Err(ProtocolError::TooShort {
expected: HEADER_SIZE,
actual: data.len(),
});
}
let stream_id = data.get_u32();
let msg_type_raw = data.get_u8();
let frame_flags = data.get_u8();
let payload_len = data.get_u32() as usize;
if data.remaining() < payload_len {
return Err(ProtocolError::Incomplete {
expected: HEADER_SIZE + payload_len,
actual: HEADER_SIZE + data.remaining(),
});
}
let msg_type =
MsgType::from_u8(msg_type_raw).ok_or(ProtocolError::UnknownMsgType(msg_type_raw))?;
let payload = data.split_to(payload_len);
Ok(Self {
stream_id,
msg_type,
flags: frame_flags,
payload,
})
}
}
/// Protocol errors.
#[derive(Debug, thiserror::Error)]
pub enum ProtocolError {
#[error("frame too short: expected {expected} bytes, got {actual}")]
TooShort { expected: usize, actual: usize },
#[error("frame incomplete: expected {expected} bytes, got {actual}")]
Incomplete { expected: usize, actual: usize },
#[error("unknown message type: 0x{0:02x}")]
UnknownMsgType(u8),
}
/// JSON payload for REQUEST_HEADERS frames.
#[derive(Debug, serde::Deserialize)]
pub struct RequestMeta {
pub method: String,
pub url: String,
pub headers: std::collections::HashMap<String, String>,
#[serde(default = "default_timeout", deserialize_with = "deserialize_timeout")]
pub timeout: u64,
}
fn default_timeout() -> u64 {
60
}
fn deserialize_timeout<'de, D>(deserializer: D) -> Result<u64, D::Error>
where
D: serde::Deserializer<'de>,
{
#[derive(serde::Deserialize)]
#[serde(untagged)]
enum TimeoutValue {
Int(u64),
Float(f64),
}
match <TimeoutValue as serde::Deserialize>::deserialize(deserializer)? {
TimeoutValue::Int(v) => Ok(v),
TimeoutValue::Float(v) => {
if !v.is_finite() || v < 0.0 {
return Err(serde::de::Error::custom(
"timeout must be a non-negative finite number",
));
}
if v.fract() != 0.0 {
return Err(serde::de::Error::custom("timeout must be integer seconds"));
}
if v > (u64::MAX as f64) {
return Err(serde::de::Error::custom("timeout is too large"));
}
Ok(v as u64)
}
}
}
/// JSON payload for RESPONSE_HEADERS frames.
#[derive(Debug, serde::Serialize)]
pub struct ResponseMeta {
pub status: u16,
/// Header list preserving duplicates (e.g. multiple Set-Cookie).
pub headers: Vec<(String, String)>,
}
// ---------------------------------------------------------------------------
// Tunnel frame compression helpers
// ---------------------------------------------------------------------------
/// Minimum payload size to attempt gzip compression (bytes).
const COMPRESS_MIN_SIZE: usize = 512;
/// If the frame has the GZIP_COMPRESSED flag, decompress the payload; otherwise
/// return a clone of the raw payload bytes.
pub fn decompress_if_gzip(frame: &Frame) -> Result<Bytes, std::io::Error> {
if frame.is_gzip() {
decompress_gzip(&frame.payload)
} else {
Ok(frame.payload.clone())
}
}
/// Gzip-compress `data` if it is large enough and compression actually shrinks
/// the payload. Returns `(payload, extra_flags)` where `extra_flags` contains
/// `GZIP_COMPRESSED` when compression was applied.
pub fn compress_payload(data: Bytes) -> (Bytes, u8) {
if data.len() >= COMPRESS_MIN_SIZE {
if let Ok(compressed) = compress_gzip(&data) {
if compressed.len() < data.len() {
return (compressed, flags::GZIP_COMPRESSED);
}
}
}
(data, 0)
}
fn decompress_gzip(data: &[u8]) -> Result<Bytes, std::io::Error> {
use flate2::read::GzDecoder;
use std::io::Read;
let mut decoder = GzDecoder::new(data);
let mut buf = Vec::new();
decoder.read_to_end(&mut buf)?;
Ok(Bytes::from(buf))
}
fn compress_gzip(data: &[u8]) -> Result<Bytes, std::io::Error> {
use flate2::write::GzEncoder;
use flate2::Compression;
use std::io::Write;
let mut encoder = GzEncoder::new(Vec::new(), Compression::fast());
encoder.write_all(data)?;
let compressed = encoder.finish()?;
Ok(Bytes::from(compressed))
}
#[cfg(test)]
mod tests {
use super::RequestMeta;
#[test]
fn request_meta_accepts_integer_timeout() {
let raw = br#"{"method":"GET","url":"https://example.com","headers":{},"timeout":15}"#;
let meta: RequestMeta = serde_json::from_slice(raw).expect("parse request meta");
assert_eq!(meta.timeout, 15);
}
#[test]
fn request_meta_accepts_integer_like_float_timeout() {
let raw = br#"{"method":"GET","url":"https://example.com","headers":{},"timeout":15.0}"#;
let meta: RequestMeta = serde_json::from_slice(raw).expect("parse request meta");
assert_eq!(meta.timeout, 15);
}
}
+506
View File
@@ -0,0 +1,506 @@
//! Per-stream request handler.
//!
//! Receives request frames, executes the upstream HTTP request,
//! and sends response frames back through the writer channel.
use std::io;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
use std::sync::Arc;
use std::time::{Duration, Instant};
use bytes::Bytes;
use futures_util::stream;
use futures_util::StreamExt;
use http_body_util::BodyExt;
use hyper::body::Frame as BodyFrame;
use tokio::sync::mpsc;
use tracing::{debug, warn};
use crate::state::{AppState, ServerContext};
use crate::target_filter;
use crate::upstream_client;
use super::protocol::{
compress_payload, decompress_if_gzip, flags, Frame as TunnelFrame, MsgType, RequestMeta,
ResponseMeta,
};
use super::writer::FrameSender;
/// Maximum response body chunk size per frame (32 KB).
const MAX_CHUNK_SIZE: usize = 32 * 1024;
/// Timeout for sending a single frame to the writer channel.
/// If the writer is congested (TCP backpressure), we abandon the stream
/// rather than blocking indefinitely and exhausting the stream pool.
const FRAME_SEND_TIMEOUT: Duration = Duration::from_secs(30);
/// Minimum allowed upstream request timeout (seconds).
const MIN_TIMEOUT_SECS: u64 = 5;
/// Maximum allowed upstream request timeout (seconds).
const MAX_TIMEOUT_SECS: u64 = 300;
/// Headers that must not be forwarded to upstream (hop-by-hop or security-sensitive).
///
/// `host` and `content-length` are managed by the HTTP client (reqwest/hyper):
/// - `host` → translated to `:authority` pseudo-header in HTTP/2; forwarding
/// the original `host` alongside `:authority` triggers PROTOCOL_ERROR on
/// strict H2 implementations (e.g. Google APIs).
/// - `content-length` → recalculated by hyper from the actual body; a stale
/// value from the tunnel (body may have been re-compressed) causes H2
/// PROTOCOL_ERROR when it mismatches the real frame length.
const BLOCKED_HEADERS: &[&str] = &[
"connection",
"content-length",
"host",
"keep-alive",
"proxy-authenticate",
"proxy-authorization",
"proxy-connection",
"te",
"trailer",
"transfer-encoding",
"upgrade",
];
/// Handle a single stream: receive body, execute upstream, send response.
pub async fn handle_stream(
state: Arc<AppState>,
server: Arc<ServerContext>,
stream_id: u32,
meta: RequestMeta,
body_rx: mpsc::Receiver<TunnelFrame>,
frame_tx: FrameSender,
) {
server.active_connections.fetch_add(1, Ordering::Release);
let connect_elapsed =
handle_stream_inner(&state, &server, stream_id, meta, body_rx, &frame_tx).await;
server.active_connections.fetch_sub(1, Ordering::Release);
if let Some(d) = connect_elapsed {
server.metrics.record_request(d);
}
}
/// Send a frame to the writer with a timeout. Returns false if send failed.
async fn send_frame(tx: &FrameSender, frame: TunnelFrame) -> bool {
match tokio::time::timeout(FRAME_SEND_TIMEOUT, tx.send(frame)).await {
Ok(Ok(())) => true,
Ok(Err(_)) => {
// Channel closed (writer exited)
false
}
Err(_) => {
// Timeout — writer is congested
warn!("frame send timeout (writer congested), abandoning stream");
false
}
}
}
/// Returns the connection-establishment duration (DNS + TCP/TLS + TTFB) if the
/// upstream request succeeded, or `None` if the request never reached the
/// response-headers stage.
async fn handle_stream_inner(
state: &AppState,
server: &ServerContext,
stream_id: u32,
meta: RequestMeta,
body_rx: mpsc::Receiver<TunnelFrame>,
frame_tx: &FrameSender,
) -> Option<Duration> {
// Validate target
let target_url = match url::Url::parse(&meta.url) {
Ok(u) => u,
Err(e) => {
send_error(frame_tx, stream_id, &format!("invalid URL: {e}")).await;
return None;
}
};
// Only allow http/https schemes (block file://, data://, etc.)
match target_url.scheme() {
"http" | "https" => {}
other => {
send_error(
frame_tx,
stream_id,
&format!("unsupported URL scheme: {other}"),
)
.await;
return None;
}
}
let host = match target_url.host_str() {
Some(h) => h.to_string(),
None => {
send_error(frame_tx, stream_id, "missing host in URL").await;
return None;
}
};
let port = target_url.port_or_known_default().unwrap_or(443);
// DNS + target validation (populates dns_cache for SafeDnsResolver)
let connect_start = Instant::now();
{
let allowed_ports = Arc::clone(&server.dynamic.load().allowed_ports);
if let Err(e) =
target_filter::validate_target(&host, port, &allowed_ports, &state.dns_cache).await
{
server.metrics.dns_failures.fetch_add(1, Ordering::Release);
send_error(frame_tx, stream_id, &format!("target blocked: {e}")).await;
return None;
}
}
let dns_ms = connect_start.elapsed().as_millis() as u64;
// Execute upstream request
let client = &state.upstream_client;
let timeout = Duration::from_secs(meta.timeout.clamp(MIN_TIMEOUT_SECS, MAX_TIMEOUT_SECS));
let request_body_size = Arc::new(AtomicUsize::new(0));
let request_body = build_streaming_request_body(body_rx, Arc::clone(&request_body_size));
let method: hyper::Method = meta.method.parse().unwrap_or(hyper::Method::GET);
let mut request = match hyper::Request::builder()
.method(method)
.uri(meta.url.as_str())
.body(request_body)
{
Ok(request) => request,
Err(e) => {
send_error(
frame_tx,
stream_id,
&format!("invalid upstream request: {e}"),
)
.await;
return None;
}
};
let headers = request.headers_mut();
for (k, v) in &meta.headers {
let k_lower = k.to_ascii_lowercase();
if BLOCKED_HEADERS.contains(&k_lower.as_str()) {
continue;
}
if let (Ok(name), Ok(value)) = (
hyper::header::HeaderName::from_bytes(k.as_bytes()),
hyper::header::HeaderValue::from_str(v),
) {
headers.insert(name, value);
}
}
let mut captured_connection = upstream_client::capture_connection(&mut request);
let connection_start = Instant::now();
let connection_capture = tokio::spawn(async move {
let connected = captured_connection.wait_for_connection_metadata().await;
connected
.as_ref()
.map(|_| connection_start.elapsed().as_millis() as u64)
});
let upstream_start = Instant::now();
let response = match tokio::time::timeout(timeout, client.request(request)).await {
Ok(Ok(response)) => response,
Ok(Err(e)) => {
connection_capture.abort();
server
.metrics
.failed_requests
.fetch_add(1, Ordering::Release);
let msg = if e.is_connect() {
format!("upstream connect error: {e}")
} else {
format!("upstream error: {e}")
};
send_error(frame_tx, stream_id, &msg).await;
return None;
}
Err(_) => {
connection_capture.abort();
server
.metrics
.failed_requests
.fetch_add(1, Ordering::Release);
send_error(frame_tx, stream_id, "upstream timeout").await;
return None;
}
};
// Capture connection-establishment duration (DNS + TCP/TLS + TTFB)
// before proceeding to stream the response body.
let connect_elapsed = connect_start.elapsed();
// Send RESPONSE_HEADERS
let status = response.status().as_u16();
let ttfb_ms = upstream_start.elapsed().as_millis() as u64;
// Short timeout: on connection reuse hyper may never fire the connect
// callback, so avoid blocking indefinitely.
let connection_acquire_ms =
match tokio::time::timeout(Duration::from_millis(100), connection_capture).await {
Ok(Ok(ms)) => ms,
Ok(Err(_)) => None, // JoinError (task panicked / cancelled)
Err(_) => None, // timeout -- task is detached but lightweight
};
let request_timing =
upstream_client::resolve_request_timing(&response, connection_acquire_ms, ttfb_ms);
let mut resp_headers: Vec<(String, String)> = Vec::with_capacity(response.headers().len() + 1);
for (k, v) in response.headers() {
if let Ok(vs) = v.to_str() {
resp_headers.push((k.as_str().to_string(), vs.to_string()));
}
}
let timing = serde_json::json!({
"dns_ms": dns_ms,
"connection_acquire_ms": request_timing.connection_acquire_ms,
"connection_reused": request_timing.connection_reused,
"connect_ms": request_timing.connect_ms,
"tls_ms": request_timing.tls_ms,
"ttfb_ms": ttfb_ms,
"upstream_ms": ttfb_ms,
"response_wait_ms": request_timing.response_wait_ms,
"upstream_processing_ms": request_timing.response_wait_ms,
"timing_source": "instrumented_connector",
"total_ms": connect_elapsed.as_millis() as u64,
"body_size": request_body_size.load(Ordering::Relaxed),
"mode": "tunnel",
});
resp_headers.push(("x-proxy-timing".to_string(), timing.to_string()));
let resp_meta = ResponseMeta {
status,
headers: resp_headers,
};
let meta_json: Bytes = serde_json::to_vec(&resp_meta).unwrap_or_default().into();
let (meta_payload, meta_flags) = compress_payload(meta_json);
if !send_frame(
frame_tx,
TunnelFrame::new(
stream_id,
MsgType::ResponseHeaders,
meta_flags,
meta_payload,
),
)
.await
{
return Some(connect_elapsed);
}
// Stream response body — relay upstream bytes through the tunnel.
// Apply tunnel-level frame compression for chunks that benefit from it
// (e.g. uncompressed SSE text). Already-compressed data (gzip/br from
// upstream Content-Encoding) won't shrink further and will be sent as-is
// thanks to the size check in compress_payload().
let mut stream = response.into_body().into_data_stream();
while let Some(chunk_result) = stream.next().await {
match chunk_result {
Ok(chunk) => {
if chunk.len() <= MAX_CHUNK_SIZE {
let (payload, extra_flags) = compress_payload(chunk);
if !send_frame(
frame_tx,
TunnelFrame::new(stream_id, MsgType::ResponseBody, extra_flags, payload),
)
.await
{
return Some(connect_elapsed);
}
} else {
// Split oversized chunks, compress each slice
let mut offset = 0;
while offset < chunk.len() {
let end = (offset + MAX_CHUNK_SIZE).min(chunk.len());
let slice = chunk.slice(offset..end);
let (payload, extra_flags) = compress_payload(slice);
if !send_frame(
frame_tx,
TunnelFrame::new(
stream_id,
MsgType::ResponseBody,
extra_flags,
payload,
),
)
.await
{
return Some(connect_elapsed);
}
offset = end;
}
}
}
Err(e) => {
server.metrics.stream_errors.fetch_add(1, Ordering::Release);
warn!(stream_id, error = %e, "upstream body read error");
send_error(frame_tx, stream_id, &format!("body read error: {e}")).await;
return Some(connect_elapsed);
}
}
}
// Send STREAM_END
let _ = send_frame(
frame_tx,
TunnelFrame::new(
stream_id,
MsgType::StreamEnd,
flags::END_STREAM,
Bytes::new(),
),
)
.await;
debug!(stream_id, status, "stream completed");
Some(connect_elapsed)
}
async fn send_error(tx: &FrameSender, stream_id: u32, msg: &str) {
// Error frames use best-effort delivery — don't block if writer is congested
let _ = send_frame(
tx,
TunnelFrame::new(
stream_id,
MsgType::StreamError,
0,
Bytes::from(msg.to_string()),
),
)
.await;
}
fn build_streaming_request_body(
body_rx: mpsc::Receiver<TunnelFrame>,
body_size: Arc<AtomicUsize>,
) -> upstream_client::UpstreamRequestBody {
let body_stream = stream::unfold(
(body_rx, body_size, false),
|(mut body_rx, body_size, finished)| async move {
if finished {
return None;
}
loop {
let frame = match body_rx.recv().await {
Some(frame) => frame,
None => return None,
};
match frame.msg_type {
MsgType::RequestBody => {
let end_stream = frame.is_end_stream();
let payload = match decompress_if_gzip(&frame) {
Ok(payload) => payload,
Err(error) => {
let err =
io::Error::other(format!("gzip decompress failed: {error}"));
return Some((Err(err), (body_rx, body_size, true)));
}
};
if payload.is_empty() {
if end_stream {
return None;
}
continue;
}
body_size.fetch_add(payload.len(), Ordering::Relaxed);
return Some((
Ok(BodyFrame::data(payload)),
(body_rx, body_size, end_stream),
));
}
MsgType::StreamError => {
let message = String::from_utf8(frame.payload.to_vec())
.unwrap_or_else(|_| "client cancelled request body".to_string());
return Some((Err(io::Error::other(message)), (body_rx, body_size, true)));
}
MsgType::StreamEnd => return None,
_ => continue,
}
}
},
);
upstream_client::stream_request_body(body_stream)
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn streaming_request_body_yields_chunks_and_tracks_size() {
let (tx, rx) = mpsc::channel(4);
let body_size = Arc::new(AtomicUsize::new(0));
let mut body = build_streaming_request_body(rx, Arc::clone(&body_size));
tx.send(TunnelFrame::new(
1,
MsgType::RequestBody,
0,
Bytes::from_static(b"abc"),
))
.await
.expect("send first chunk");
tx.send(TunnelFrame::new(
1,
MsgType::RequestBody,
flags::END_STREAM,
Bytes::from_static(b"def"),
))
.await
.expect("send final chunk");
drop(tx);
let first = body
.frame()
.await
.expect("first frame")
.expect("first frame ok")
.into_data()
.expect("first data frame");
let second = body
.frame()
.await
.expect("second frame")
.expect("second frame ok")
.into_data()
.expect("second data frame");
assert_eq!(first, Bytes::from_static(b"abc"));
assert_eq!(second, Bytes::from_static(b"def"));
assert!(body.frame().await.is_none());
assert_eq!(body_size.load(Ordering::Relaxed), 6);
}
#[tokio::test]
async fn streaming_request_body_surfaces_client_cancel_as_error() {
let (tx, rx) = mpsc::channel(4);
let body_size = Arc::new(AtomicUsize::new(0));
let mut body = build_streaming_request_body(rx, Arc::clone(&body_size));
tx.send(TunnelFrame::new(
1,
MsgType::StreamError,
0,
Bytes::from_static(b"client cancelled"),
))
.await
.expect("send cancel frame");
drop(tx);
let err = body
.frame()
.await
.expect("error frame present")
.expect_err("body should surface cancellation error");
assert!(err.to_string().contains("client cancelled"));
assert!(body.frame().await.is_none());
assert_eq!(body_size.load(Ordering::Relaxed), 0);
}
}
+63
View File
@@ -0,0 +1,63 @@
//! 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::time::Duration;
use futures_util::SinkExt;
use tokio::sync::mpsc;
use tokio::task::JoinHandle;
use tokio_tungstenite::tungstenite::Message;
use tracing::{debug, error, trace};
use super::protocol::Frame;
/// Sender half — cloned by stream handlers and heartbeat.
pub type FrameSender = mpsc::Sender<Frame>;
/// 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.
pub fn spawn_writer<S>(mut sink: S, ping_interval: Duration) -> (FrameSender, JoinHandle<()>)
where
S: SinkExt<Message, Error = tokio_tungstenite::tungstenite::Error> + Unpin + Send + 'static,
{
let (tx, mut rx) = mpsc::channel::<Frame>(256);
let handle = tokio::spawn(async move {
let mut ping_ticker = tokio::time::interval(ping_interval);
ping_ticker.tick().await; // skip first immediate tick
loop {
tokio::select! {
frame = rx.recv() => {
match frame {
Some(frame) => {
let data = frame.encode();
if let Err(e) = sink.send(Message::Binary(data.into())).await {
error!(error = %e, "failed to write frame to WebSocket");
break;
}
}
None => break, // all senders dropped
}
}
_ = ping_ticker.tick() => {
if let Err(e) = sink.send(Message::Ping(vec![])).await {
error!(error = %e, "failed to send WebSocket ping");
break;
}
trace!("sent WebSocket ping");
}
}
}
debug!("writer task exiting");
let _ = sink.close().await;
});
(tx, handle)
}
+444
View File
@@ -0,0 +1,444 @@
use std::future::Future;
use std::io;
use std::net::IpAddr;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use std::time::Duration;
use bytes::Bytes;
use futures_util::Stream;
use http_body_util::combinators::UnsyncBoxBody;
use http_body_util::{BodyExt, 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::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>;
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()
}
#[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>,
}
impl ValidatedResolver {
pub fn new(dns_cache: Arc<DnsCache>) -> Self {
Self { dns_cache }
}
}
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 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, 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>,
}
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);
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(tcp),
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()),
}
})
}
}
pub fn build_upstream_client(config: &Config, dns_cache: Arc<DnsCache>) -> UpstreamClient {
let mut http = HttpConnector::new_with_resolver(ValidatedResolver::new(dns_cache));
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(),
};
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());
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() -> Arc<ClientConfig> {
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 = 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(PlainStream),
Https(TlsStream),
}
impl Connection for MaybeHttpsStream {
fn connected(&self) -> Connected {
match self {
Self::Http(stream) => stream.connected(),
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 hyper::Response;
#[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);
}
}
+51
View File
@@ -0,0 +1,51 @@
# Alembic 配置文件
# 用于数据库版本化迁移
[alembic]
# 迁移脚本存放目录
script_location = alembic
# 模板文件
file_template = %%(year)d%%(month).2d%%(day).2d_%%(hour).2d%%(minute).2d_%%(rev)s_%%(slug)s
# 时区(用于生成迁移文件的时间戳)
timezone = UTC
# 数据库连接 URL(会被 env.py 从环境变量覆盖)
# Docker 环境中会从 DATABASE_URL 环境变量读取
sqlalchemy.url = postgresql://postgres:${DB_PASSWORD}@localhost:5432/aether
# 日志配置
[loggers]
keys = root,sqlalchemy,alembic
[handlers]
keys = console
[formatters]
keys = generic
[logger_root]
level = WARN
handlers = console
qualname =
[logger_sqlalchemy]
level = WARN
handlers =
qualname = sqlalchemy.engine
[logger_alembic]
level = INFO
handlers =
qualname = alembic
[handler_console]
class = StreamHandler
args = (sys.stderr,)
level = NOTSET
formatter = generic
[formatter_generic]
format = %(levelname)-5.5s [%(name)s] %(message)s
datefmt = %H:%M:%S
+121
View File
@@ -0,0 +1,121 @@
"""
Alembic 环境配置
用于数据库迁移的运行时环境设置
"""
import os
import sys
from logging.config import fileConfig
from pathlib import Path
from sqlalchemy import engine_from_config, pool, text
from alembic import context
# 添加项目根目录到 Python 路径
sys.path.insert(0, os.path.dirname(os.path.dirname(__file__)))
# 加载 .env 文件(本地开发时需要)
try:
from dotenv import load_dotenv
env_file = Path(__file__).parent.parent / ".env"
if env_file.exists():
load_dotenv(env_file)
except ImportError:
pass
# 导入所有数据库模型(确保 Alembic 能检测到所有表)
from src.models.database import Base
# Alembic Config 对象
config = context.config
# 从环境变量获取数据库 URL
# 优先使用 DATABASE_URL,否则从 DB_PASSWORD 自动构建(与 docker compose 保持一致)
database_url = os.getenv("DATABASE_URL")
if not database_url:
db_password = os.getenv("DB_PASSWORD", "")
db_host = os.getenv("DB_HOST", "localhost")
db_port = os.getenv("DB_PORT", "5432")
db_name = os.getenv("DB_NAME", "aether")
db_user = os.getenv("DB_USER", "postgres")
database_url = f"postgresql://{db_user}:{db_password}@{db_host}:{db_port}/{db_name}"
config.set_main_option("sqlalchemy.url", database_url)
# 配置日志
if config.config_file_name is not None:
fileConfig(config.config_file_name)
# 目标元数据(包含所有表定义)
target_metadata = Base.metadata
# PostgreSQL 全局迁移锁,避免多进程并发执行 Alembic 导致竞态(重复加列/索引等)
# 使用会话级 advisory lock(pg_advisory_lock),在迁移完成后手动释放。
# ID 由 crc32("aether-alembic-migration") 拼接生成,仅需全局唯一即可。
MIGRATION_ADVISORY_LOCK_ID = 582694137405821
def run_migrations_offline() -> None:
"""
离线模式运行迁移
在离线模式下,不需要连接数据库,
只生成 SQL 脚本
"""
url = config.get_main_option("sqlalchemy.url")
context.configure(
url=url,
target_metadata=target_metadata,
literal_binds=True,
dialect_opts={"paramstyle": "named"},
compare_type=True, # 比较列类型变更
compare_server_default=True, # 比较默认值变更
)
with context.begin_transaction():
context.run_migrations()
def run_migrations_online() -> None:
"""
在线模式运行迁移
在线模式下,直接连接数据库执行迁移
"""
connectable = engine_from_config(
config.get_section(config.config_ini_section, {}),
prefix="sqlalchemy.",
poolclass=pool.NullPool,
)
with connectable.connect() as connection:
try:
# 使用会话级 advisory lock(非事务级),避免干扰 Alembic 的事务管理。
# pg_advisory_lock 在会话结束时自动释放,不受 COMMIT/ROLLBACK 影响。
if connection.dialect.name == "postgresql":
connection.execute(
text("SELECT pg_advisory_lock(:lock_id)"),
{"lock_id": MIGRATION_ADVISORY_LOCK_ID},
)
connection.commit()
context.configure(
connection=connection,
target_metadata=target_metadata,
compare_type=True,
compare_server_default=True,
transaction_per_migration=True, # 每个迁移文件独立事务,完成即提交
)
with context.begin_transaction():
context.run_migrations()
except Exception:
raise
# 根据模式选择运行方式
if context.is_offline_mode():
run_migrations_offline()
else:
run_migrations_online()
+26
View File
@@ -0,0 +1,26 @@
"""${message}
Revision ID: ${up_revision}
Revises: ${down_revision | comma,n}
Create Date: ${create_date}
"""
from alembic import op
import sqlalchemy as sa
${imports if imports else ""}
# revision identifiers, used by Alembic.
revision = ${repr(up_revision)}
down_revision = ${repr(down_revision)}
branch_labels = ${repr(branch_labels)}
depends_on = ${repr(depends_on)}
def upgrade() -> None:
"""应用迁移:升级到新版本"""
${upgrades if upgrades else "pass"}
def downgrade() -> None:
"""回滚迁移:降级到旧版本"""
${downgrades if downgrades else "pass"}
+775
View File
@@ -0,0 +1,775 @@
"""Baseline migration - all tables consolidated
Revision ID: 20251210_baseline
Revises:
Create Date: 2024-12-10
This is the consolidated baseline migration that creates all tables from scratch.
Includes all schema changes up to circuit breaker v2.
"""
from alembic import op
import sqlalchemy as sa
from sqlalchemy.dialects import postgresql
# revision identifiers
revision = "20251210_baseline"
down_revision = None
branch_labels = None
depends_on = None
def upgrade() -> None:
# Create ENUM types (with IF NOT EXISTS for idempotency)
op.execute("DO $$ BEGIN CREATE TYPE userrole AS ENUM ('admin', 'user'); EXCEPTION WHEN duplicate_object THEN NULL; END $$")
op.execute(
"DO $$ BEGIN CREATE TYPE providerbillingtype AS ENUM ('monthly_quota', 'pay_as_you_go', 'free_tier'); EXCEPTION WHEN duplicate_object THEN NULL; END $$"
)
# ==================== users ====================
op.create_table(
"users",
sa.Column("id", sa.String(36), primary_key=True, index=True),
sa.Column("email", sa.String(255), unique=True, index=True, nullable=False),
sa.Column("username", sa.String(100), unique=True, index=True, nullable=False),
sa.Column("password_hash", sa.String(255), nullable=False),
sa.Column(
"role",
postgresql.ENUM("admin", "user", name="userrole", create_type=False),
nullable=False,
server_default="user",
),
sa.Column("allowed_providers", sa.JSON, nullable=True),
sa.Column("allowed_endpoints", sa.JSON, nullable=True),
sa.Column("allowed_models", sa.JSON, nullable=True),
sa.Column("model_capability_settings", sa.JSON, nullable=True),
sa.Column("quota_usd", sa.Float, nullable=True),
sa.Column("used_usd", sa.Float, server_default="0.0"),
sa.Column("total_usd", sa.Float, server_default="0.0"),
sa.Column("is_active", sa.Boolean, server_default="true", nullable=False),
sa.Column("is_deleted", sa.Boolean, server_default="false", nullable=False),
sa.Column(
"created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.Column(
"updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.Column("last_login_at", sa.DateTime(timezone=True), nullable=True),
)
# ==================== providers ====================
op.create_table(
"providers",
sa.Column("id", sa.String(36), primary_key=True, index=True),
sa.Column("name", sa.String(100), unique=True, index=True, nullable=False),
sa.Column("display_name", sa.String(100), nullable=False),
sa.Column("description", sa.Text, nullable=True),
sa.Column("website", sa.String(500), nullable=True),
sa.Column(
"billing_type",
postgresql.ENUM(
"monthly_quota", "pay_as_you_go", "free_tier", name="providerbillingtype", create_type=False
),
nullable=False,
server_default="pay_as_you_go",
),
sa.Column("monthly_quota_usd", sa.Float, nullable=True),
sa.Column("monthly_used_usd", sa.Float, server_default="0.0"),
sa.Column("quota_reset_day", sa.Integer, server_default="30"),
sa.Column("quota_last_reset_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("quota_expires_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("rpm_limit", sa.Integer, nullable=True),
sa.Column("rpm_used", sa.Integer, server_default="0"),
sa.Column("rpm_reset_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("provider_priority", sa.Integer, server_default="100"),
sa.Column("is_active", sa.Boolean, server_default="true", nullable=False),
sa.Column("rate_limit", sa.Integer, nullable=True),
sa.Column("concurrent_limit", sa.Integer, nullable=True),
sa.Column("config", sa.JSON, nullable=True),
sa.Column(
"created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.Column(
"updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
)
# ==================== global_models ====================
op.create_table(
"global_models",
sa.Column("id", sa.String(36), primary_key=True, index=True),
sa.Column("name", sa.String(100), unique=True, index=True, nullable=False),
sa.Column("display_name", sa.String(100), nullable=False),
sa.Column("description", sa.Text, nullable=True),
sa.Column("icon_url", sa.String(500), nullable=True),
sa.Column("official_url", sa.String(500), nullable=True),
sa.Column("default_price_per_request", sa.Float, nullable=True),
sa.Column("default_tiered_pricing", sa.JSON, nullable=False),
sa.Column("default_supports_vision", sa.Boolean, server_default="false", nullable=True),
sa.Column("default_supports_function_calling", sa.Boolean, server_default="false", nullable=True),
sa.Column("default_supports_streaming", sa.Boolean, server_default="true", nullable=True),
sa.Column("default_supports_extended_thinking", sa.Boolean, server_default="false", nullable=True),
sa.Column("default_supports_image_generation", sa.Boolean, server_default="false", nullable=True),
sa.Column("supported_capabilities", sa.JSON, nullable=True),
sa.Column("is_active", sa.Boolean, server_default="true", nullable=False),
sa.Column("usage_count", sa.Integer, server_default="0", nullable=False, index=True),
sa.Column(
"created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.Column(
"updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
)
# ==================== api_keys ====================
op.create_table(
"api_keys",
sa.Column("id", sa.String(36), primary_key=True, index=True),
sa.Column(
"user_id", sa.String(36), sa.ForeignKey("users.id", ondelete="CASCADE"), nullable=False
),
sa.Column("key_hash", sa.String(64), unique=True, index=True, nullable=False),
sa.Column("key_encrypted", sa.Text, nullable=True),
sa.Column("name", sa.String(100), nullable=True),
sa.Column("total_requests", sa.Integer, server_default="0"),
sa.Column("total_cost_usd", sa.Float, server_default="0.0"),
sa.Column("balance_used_usd", sa.Float, server_default="0.0"),
sa.Column("current_balance_usd", sa.Float, nullable=True),
sa.Column("is_standalone", sa.Boolean, server_default="false", nullable=False),
sa.Column("allowed_providers", sa.JSON, nullable=True),
sa.Column("allowed_endpoints", sa.JSON, nullable=True),
sa.Column("allowed_api_formats", sa.JSON, nullable=True),
sa.Column("allowed_models", sa.JSON, nullable=True),
sa.Column("rate_limit", sa.Integer, server_default="100"),
sa.Column("concurrent_limit", sa.Integer, server_default="5", nullable=True),
sa.Column("force_capabilities", sa.JSON, nullable=True),
sa.Column("is_active", sa.Boolean, server_default="true", nullable=False),
sa.Column("last_used_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("expires_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("auto_delete_on_expiry", sa.Boolean, server_default="false", nullable=False),
sa.Column(
"created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.Column(
"updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
)
# ==================== provider_endpoints ====================
op.create_table(
"provider_endpoints",
sa.Column("id", sa.String(36), primary_key=True, index=True),
sa.Column(
"provider_id",
sa.String(36),
sa.ForeignKey("providers.id", ondelete="CASCADE"),
nullable=False,
),
sa.Column("api_format", sa.String(50), nullable=False),
sa.Column("base_url", sa.String(500), nullable=False),
sa.Column("headers", sa.JSON, nullable=True),
sa.Column("timeout", sa.Integer, server_default="300"),
sa.Column("max_retries", sa.Integer, server_default="3"),
sa.Column("max_concurrent", sa.Integer, nullable=True),
sa.Column("rate_limit", sa.Integer, nullable=True),
sa.Column("is_active", sa.Boolean, server_default="true", nullable=False),
sa.Column("custom_path", sa.String(200), nullable=True),
sa.Column("config", sa.JSON, nullable=True),
sa.Column(
"created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.Column(
"updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.UniqueConstraint("provider_id", "api_format", name="uq_provider_api_format"),
)
op.create_index(
"idx_endpoint_format_active", "provider_endpoints", ["api_format", "is_active"]
)
# ==================== models ====================
op.create_table(
"models",
sa.Column("id", sa.String(36), primary_key=True, index=True),
sa.Column(
"provider_id", sa.String(36), sa.ForeignKey("providers.id"), nullable=False
),
sa.Column(
"global_model_id",
sa.String(36),
sa.ForeignKey("global_models.id"),
nullable=False,
index=True,
),
sa.Column("provider_model_name", sa.String(200), nullable=False),
sa.Column("price_per_request", sa.Float, nullable=True),
sa.Column("tiered_pricing", sa.JSON, nullable=True),
sa.Column("supports_vision", sa.Boolean, nullable=True),
sa.Column("supports_function_calling", sa.Boolean, nullable=True),
sa.Column("supports_streaming", sa.Boolean, nullable=True),
sa.Column("supports_extended_thinking", sa.Boolean, nullable=True),
sa.Column("supports_image_generation", sa.Boolean, nullable=True),
sa.Column("is_active", sa.Boolean, server_default="true", nullable=False),
sa.Column("is_available", sa.Boolean, server_default="true"),
sa.Column("config", sa.JSON, nullable=True),
sa.Column(
"created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.Column(
"updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.UniqueConstraint("provider_id", "provider_model_name", name="uq_provider_model"),
)
# ==================== model_mappings ====================
op.create_table(
"model_mappings",
sa.Column("id", sa.String(36), primary_key=True, index=True),
sa.Column("source_model", sa.String(200), nullable=False, index=True),
sa.Column(
"target_global_model_id",
sa.String(36),
sa.ForeignKey("global_models.id", ondelete="CASCADE"),
nullable=False,
index=True,
),
sa.Column(
"provider_id", sa.String(36), sa.ForeignKey("providers.id"), nullable=True, index=True
),
sa.Column("mapping_type", sa.String(20), nullable=False, server_default="alias", index=True),
sa.Column("is_active", sa.Boolean, server_default="true", nullable=False),
sa.Column(
"created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.Column(
"updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.UniqueConstraint("source_model", "provider_id", name="uq_model_mapping_source_provider"),
)
# ==================== provider_api_keys ====================
op.create_table(
"provider_api_keys",
sa.Column("id", sa.String(36), primary_key=True, index=True),
sa.Column(
"endpoint_id",
sa.String(36),
sa.ForeignKey("provider_endpoints.id", ondelete="CASCADE"),
nullable=False,
),
sa.Column("api_key", sa.String(500), nullable=False),
sa.Column("name", sa.String(100), nullable=False),
sa.Column("note", sa.String(500), nullable=True),
sa.Column("rate_multiplier", sa.Float, server_default="1.0", nullable=False),
sa.Column("internal_priority", sa.Integer, server_default="50"),
sa.Column("global_priority", sa.Integer, nullable=True),
sa.Column("max_concurrent", sa.Integer, nullable=True),
sa.Column("rate_limit", sa.Integer, nullable=True),
sa.Column("daily_limit", sa.Integer, nullable=True),
sa.Column("monthly_limit", sa.Integer, nullable=True),
sa.Column("allowed_models", sa.JSON, nullable=True),
sa.Column("capabilities", sa.JSON, nullable=True),
sa.Column("learned_max_concurrent", sa.Integer, nullable=True),
sa.Column("concurrent_429_count", sa.Integer, server_default="0", nullable=False),
sa.Column("rpm_429_count", sa.Integer, server_default="0", nullable=False),
sa.Column("last_429_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("last_429_type", sa.String(50), nullable=True),
sa.Column("last_concurrent_peak", sa.Integer, nullable=True),
sa.Column("adjustment_history", sa.JSON, nullable=True),
# Sliding window fields (replaces high_utilization_start)
sa.Column("utilization_samples", sa.JSON, nullable=True),
sa.Column("last_probe_increase_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("health_score", sa.Float, server_default="1.0"),
sa.Column("consecutive_failures", sa.Integer, server_default="0"),
sa.Column("last_failure_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("cache_ttl_minutes", sa.Integer, server_default="5", nullable=False),
sa.Column("max_probe_interval_minutes", sa.Integer, server_default="32", nullable=False),
sa.Column("circuit_breaker_open", sa.Boolean, server_default="false", nullable=False),
sa.Column("circuit_breaker_open_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("next_probe_at", sa.DateTime(timezone=True), nullable=True),
# Circuit breaker v2 fields
sa.Column("request_results_window", sa.JSON, nullable=True),
sa.Column("half_open_until", sa.DateTime(timezone=True), nullable=True),
sa.Column("half_open_successes", sa.Integer, server_default="0", nullable=True),
sa.Column("half_open_failures", sa.Integer, server_default="0", nullable=True),
sa.Column("request_count", sa.Integer, server_default="0"),
sa.Column("success_count", sa.Integer, server_default="0"),
sa.Column("error_count", sa.Integer, server_default="0"),
sa.Column("total_response_time_ms", sa.Integer, server_default="0"),
sa.Column("last_used_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("last_error_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("last_error_msg", sa.Text, nullable=True),
sa.Column("is_active", sa.Boolean, server_default="true", nullable=False),
sa.Column("expires_at", sa.DateTime(timezone=True), nullable=True),
sa.Column(
"created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.Column(
"updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
)
# ==================== usage ====================
op.create_table(
"usage",
sa.Column("id", sa.String(36), primary_key=True, index=True),
sa.Column(
"user_id",
sa.String(36),
sa.ForeignKey("users.id", ondelete="SET NULL"),
nullable=True,
),
sa.Column(
"api_key_id",
sa.String(36),
sa.ForeignKey("api_keys.id", ondelete="SET NULL"),
nullable=True,
),
sa.Column("request_id", sa.String(100), unique=True, index=True, nullable=False),
sa.Column("provider", sa.String(100), nullable=False),
sa.Column("model", sa.String(100), nullable=False),
sa.Column("target_model", sa.String(100), nullable=True),
sa.Column(
"provider_id",
sa.String(36),
sa.ForeignKey("providers.id", ondelete="SET NULL"),
nullable=True,
),
sa.Column(
"provider_endpoint_id",
sa.String(36),
sa.ForeignKey("provider_endpoints.id", ondelete="SET NULL"),
nullable=True,
),
sa.Column(
"provider_api_key_id",
sa.String(36),
sa.ForeignKey("provider_api_keys.id", ondelete="SET NULL"),
nullable=True,
),
sa.Column("input_tokens", sa.Integer, server_default="0"),
sa.Column("output_tokens", sa.Integer, server_default="0"),
sa.Column("total_tokens", sa.Integer, server_default="0"),
sa.Column("cache_creation_input_tokens", sa.Integer, server_default="0"),
sa.Column("cache_read_input_tokens", sa.Integer, server_default="0"),
sa.Column("input_cost_usd", sa.Float, server_default="0.0"),
sa.Column("output_cost_usd", sa.Float, server_default="0.0"),
sa.Column("cache_cost_usd", sa.Float, server_default="0.0"),
sa.Column("cache_creation_cost_usd", sa.Float, server_default="0.0"),
sa.Column("cache_read_cost_usd", sa.Float, server_default="0.0"),
sa.Column("request_cost_usd", sa.Float, server_default="0.0"),
sa.Column("total_cost_usd", sa.Float, server_default="0.0"),
sa.Column("actual_input_cost_usd", sa.Float, server_default="0.0"),
sa.Column("actual_output_cost_usd", sa.Float, server_default="0.0"),
sa.Column("actual_cache_creation_cost_usd", sa.Float, server_default="0.0"),
sa.Column("actual_cache_read_cost_usd", sa.Float, server_default="0.0"),
sa.Column("actual_request_cost_usd", sa.Float, server_default="0.0"),
sa.Column("actual_total_cost_usd", sa.Float, server_default="0.0"),
sa.Column("rate_multiplier", sa.Float, server_default="1.0"),
sa.Column("input_price_per_1m", sa.Float, nullable=True),
sa.Column("output_price_per_1m", sa.Float, nullable=True),
sa.Column("cache_creation_price_per_1m", sa.Float, nullable=True),
sa.Column("cache_read_price_per_1m", sa.Float, nullable=True),
sa.Column("price_per_request", sa.Float, nullable=True),
sa.Column("request_type", sa.String(50), nullable=True),
sa.Column("api_format", sa.String(50), nullable=True),
sa.Column("is_stream", sa.Boolean, server_default="false"),
sa.Column("status_code", sa.Integer, nullable=True),
sa.Column("error_message", sa.Text, nullable=True),
sa.Column("response_time_ms", sa.Integer, nullable=True),
sa.Column("status", sa.String(20), server_default="completed", nullable=False, index=True),
sa.Column("request_headers", sa.JSON, nullable=True),
sa.Column("request_body", sa.JSON, nullable=True),
sa.Column("provider_request_headers", sa.JSON, nullable=True),
sa.Column("response_headers", sa.JSON, nullable=True),
sa.Column("response_body", sa.JSON, nullable=True),
sa.Column("request_body_compressed", sa.LargeBinary, nullable=True),
sa.Column("response_body_compressed", sa.LargeBinary, nullable=True),
sa.Column("request_metadata", sa.JSON, nullable=True),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
server_default=sa.func.now(),
nullable=False,
index=True,
),
)
# usage 表复合索引(优化常见查询)
op.create_index("idx_usage_user_created", "usage", ["user_id", "created_at"])
op.create_index("idx_usage_apikey_created", "usage", ["api_key_id", "created_at"])
op.create_index("idx_usage_provider_model_created", "usage", ["provider", "model", "created_at"])
# ==================== user_quotas ====================
op.create_table(
"user_quotas",
sa.Column("id", sa.String(36), primary_key=True, index=True),
sa.Column(
"user_id", sa.String(36), sa.ForeignKey("users.id", ondelete="CASCADE"), nullable=False
),
sa.Column("quota_type", sa.String(50), nullable=False),
sa.Column("quota_usd", sa.Float, nullable=False),
sa.Column("period_start", sa.DateTime(timezone=True), nullable=False),
sa.Column("period_end", sa.DateTime(timezone=True), nullable=False),
sa.Column("used_usd", sa.Float, server_default="0.0"),
sa.Column("is_active", sa.Boolean, server_default="true"),
sa.Column(
"created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.Column(
"updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
)
# ==================== system_configs ====================
op.create_table(
"system_configs",
sa.Column("id", sa.String(36), primary_key=True, index=True),
sa.Column("key", sa.String(100), unique=True, nullable=False),
sa.Column("value", sa.JSON, nullable=False),
sa.Column("description", sa.Text, nullable=True),
sa.Column(
"created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.Column(
"updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
)
# ==================== user_preferences ====================
op.create_table(
"user_preferences",
sa.Column("id", sa.String(36), primary_key=True, index=True),
sa.Column(
"user_id",
sa.String(36),
sa.ForeignKey("users.id", ondelete="CASCADE"),
unique=True,
nullable=False,
),
sa.Column("avatar_url", sa.String(500), nullable=True),
sa.Column("bio", sa.Text, nullable=True),
sa.Column(
"default_provider_id", sa.String(36), sa.ForeignKey("providers.id"), nullable=True
),
sa.Column("theme", sa.String(20), server_default="light"),
sa.Column("language", sa.String(10), server_default="zh-CN"),
sa.Column("timezone", sa.String(50), server_default="Asia/Shanghai"),
sa.Column("email_notifications", sa.Boolean, server_default="true"),
sa.Column("usage_alerts", sa.Boolean, server_default="true"),
sa.Column("announcement_notifications", sa.Boolean, server_default="true"),
sa.Column(
"created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.Column(
"updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
)
# ==================== announcements ====================
op.create_table(
"announcements",
sa.Column("id", sa.String(36), primary_key=True, index=True),
sa.Column("title", sa.String(200), nullable=False),
sa.Column("content", sa.Text, nullable=False),
sa.Column("type", sa.String(20), server_default="info"),
sa.Column("priority", sa.Integer, server_default="0"),
sa.Column(
"author_id",
sa.String(36),
sa.ForeignKey("users.id", ondelete="SET NULL"),
nullable=True,
),
sa.Column("is_active", sa.Boolean, server_default="true", index=True),
sa.Column("is_pinned", sa.Boolean, server_default="false"),
sa.Column("start_time", sa.DateTime(timezone=True), nullable=True),
sa.Column("end_time", sa.DateTime(timezone=True), nullable=True),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
server_default=sa.func.now(),
nullable=False,
index=True,
),
sa.Column(
"updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
)
# ==================== announcement_reads ====================
op.create_table(
"announcement_reads",
sa.Column("id", sa.String(36), primary_key=True, index=True),
sa.Column(
"user_id", sa.String(36), sa.ForeignKey("users.id", ondelete="CASCADE"), nullable=False
),
sa.Column(
"announcement_id", sa.String(36), sa.ForeignKey("announcements.id"), nullable=False
),
sa.Column(
"read_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.UniqueConstraint("user_id", "announcement_id", name="uq_user_announcement"),
)
# ==================== audit_logs ====================
op.create_table(
"audit_logs",
sa.Column("id", sa.String(36), primary_key=True, index=True),
sa.Column("event_type", sa.String(50), nullable=False, index=True),
sa.Column(
"user_id",
sa.String(36),
sa.ForeignKey("users.id", ondelete="SET NULL"),
nullable=True,
index=True,
),
sa.Column("api_key_id", sa.String(36), nullable=True),
sa.Column("description", sa.Text, nullable=False),
sa.Column("ip_address", sa.String(45), nullable=True),
sa.Column("user_agent", sa.String(500), nullable=True),
sa.Column("request_id", sa.String(100), nullable=True, index=True),
sa.Column("event_metadata", sa.JSON, nullable=True),
sa.Column("status_code", sa.Integer, nullable=True),
sa.Column("error_message", sa.Text, nullable=True),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
server_default=sa.func.now(),
nullable=False,
index=True,
),
)
# ==================== request_candidates ====================
op.create_table(
"request_candidates",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("request_id", sa.String(100), nullable=False, index=True),
sa.Column(
"user_id", sa.String(36), sa.ForeignKey("users.id", ondelete="CASCADE"), nullable=True
),
sa.Column(
"api_key_id",
sa.String(36),
sa.ForeignKey("api_keys.id", ondelete="CASCADE"),
nullable=True,
),
sa.Column("candidate_index", sa.Integer, nullable=False),
sa.Column("retry_index", sa.Integer, nullable=False, server_default="0"),
sa.Column(
"provider_id",
sa.String(36),
sa.ForeignKey("providers.id", ondelete="CASCADE"),
nullable=True,
),
sa.Column(
"endpoint_id",
sa.String(36),
sa.ForeignKey("provider_endpoints.id", ondelete="CASCADE"),
nullable=True,
),
sa.Column(
"key_id",
sa.String(36),
sa.ForeignKey("provider_api_keys.id", ondelete="CASCADE"),
nullable=True,
),
sa.Column("status", sa.String(20), nullable=False),
sa.Column("skip_reason", sa.Text, nullable=True),
sa.Column("is_cached", sa.Boolean, server_default="false"),
sa.Column("status_code", sa.Integer, nullable=True),
sa.Column("error_type", sa.String(50), nullable=True),
sa.Column("error_message", sa.Text, nullable=True),
sa.Column("latency_ms", sa.Integer, nullable=True),
sa.Column("concurrent_requests", sa.Integer, nullable=True),
sa.Column("extra_data", sa.JSON, nullable=True),
sa.Column("required_capabilities", sa.JSON, nullable=True),
sa.Column(
"created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.Column("started_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("finished_at", sa.DateTime(timezone=True), nullable=True),
sa.UniqueConstraint(
"request_id", "candidate_index", "retry_index", name="uq_request_candidate_with_retry"
),
)
op.create_index("idx_request_candidates_request_id", "request_candidates", ["request_id"])
op.create_index("idx_request_candidates_status", "request_candidates", ["status"])
op.create_index("idx_request_candidates_provider_id", "request_candidates", ["provider_id"])
# ==================== stats_daily ====================
op.create_table(
"stats_daily",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("date", sa.DateTime(timezone=True), nullable=False, unique=True, index=True),
sa.Column("total_requests", sa.Integer, server_default="0", nullable=False),
sa.Column("success_requests", sa.Integer, server_default="0", nullable=False),
sa.Column("error_requests", sa.Integer, server_default="0", nullable=False),
sa.Column("input_tokens", sa.BigInteger, server_default="0", nullable=False),
sa.Column("output_tokens", sa.BigInteger, server_default="0", nullable=False),
sa.Column("cache_creation_tokens", sa.BigInteger, server_default="0", nullable=False),
sa.Column("cache_read_tokens", sa.BigInteger, server_default="0", nullable=False),
sa.Column("total_cost", sa.Float, server_default="0.0", nullable=False),
sa.Column("actual_total_cost", sa.Float, server_default="0.0", nullable=False),
sa.Column("input_cost", sa.Float, server_default="0.0", nullable=False),
sa.Column("output_cost", sa.Float, server_default="0.0", nullable=False),
sa.Column("cache_creation_cost", sa.Float, server_default="0.0", nullable=False),
sa.Column("cache_read_cost", sa.Float, server_default="0.0", nullable=False),
sa.Column("avg_response_time_ms", sa.Float, server_default="0.0", nullable=False),
sa.Column("fallback_count", sa.Integer, server_default="0", nullable=False),
sa.Column("unique_models", sa.Integer, server_default="0", nullable=False),
sa.Column("unique_providers", sa.Integer, server_default="0", nullable=False),
sa.Column(
"created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.Column(
"updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
)
# ==================== stats_summary ====================
op.create_table(
"stats_summary",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("cutoff_date", sa.DateTime(timezone=True), nullable=False),
sa.Column("all_time_requests", sa.Integer, server_default="0", nullable=False),
sa.Column("all_time_success_requests", sa.Integer, server_default="0", nullable=False),
sa.Column("all_time_error_requests", sa.Integer, server_default="0", nullable=False),
sa.Column("all_time_input_tokens", sa.BigInteger, server_default="0", nullable=False),
sa.Column("all_time_output_tokens", sa.BigInteger, server_default="0", nullable=False),
sa.Column(
"all_time_cache_creation_tokens", sa.BigInteger, server_default="0", nullable=False
),
sa.Column("all_time_cache_read_tokens", sa.BigInteger, server_default="0", nullable=False),
sa.Column("all_time_cost", sa.Float, server_default="0.0", nullable=False),
sa.Column("all_time_actual_cost", sa.Float, server_default="0.0", nullable=False),
sa.Column("total_users", sa.Integer, server_default="0", nullable=False),
sa.Column("active_users", sa.Integer, server_default="0", nullable=False),
sa.Column("total_api_keys", sa.Integer, server_default="0", nullable=False),
sa.Column("active_api_keys", sa.Integer, server_default="0", nullable=False),
sa.Column(
"created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.Column(
"updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
)
# ==================== stats_user_daily ====================
op.create_table(
"stats_user_daily",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column(
"user_id", sa.String(36), sa.ForeignKey("users.id", ondelete="CASCADE"), nullable=False
),
sa.Column("date", sa.DateTime(timezone=True), nullable=False, index=True),
sa.Column("total_requests", sa.Integer, server_default="0", nullable=False),
sa.Column("success_requests", sa.Integer, server_default="0", nullable=False),
sa.Column("error_requests", sa.Integer, server_default="0", nullable=False),
sa.Column("input_tokens", sa.BigInteger, server_default="0", nullable=False),
sa.Column("output_tokens", sa.BigInteger, server_default="0", nullable=False),
sa.Column("cache_creation_tokens", sa.BigInteger, server_default="0", nullable=False),
sa.Column("cache_read_tokens", sa.BigInteger, server_default="0", nullable=False),
sa.Column("total_cost", sa.Float, server_default="0.0", nullable=False),
sa.Column(
"created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.Column(
"updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.UniqueConstraint("user_id", "date", name="uq_stats_user_daily"),
)
op.create_index("idx_stats_user_daily_user_date", "stats_user_daily", ["user_id", "date"])
# ==================== api_key_provider_mappings ====================
op.create_table(
"api_key_provider_mappings",
sa.Column("id", sa.String(36), primary_key=True, index=True),
sa.Column(
"api_key_id",
sa.String(36),
sa.ForeignKey("api_keys.id", ondelete="CASCADE"),
nullable=False,
index=True,
),
sa.Column(
"provider_id",
sa.String(36),
sa.ForeignKey("providers.id", ondelete="CASCADE"),
nullable=False,
index=True,
),
sa.Column("priority_adjustment", sa.Integer, server_default="0"),
sa.Column("weight_multiplier", sa.Float, server_default="1.0"),
sa.Column("is_enabled", sa.Boolean, server_default="true", nullable=False),
sa.Column(
"created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.Column(
"updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.UniqueConstraint("api_key_id", "provider_id", name="uq_apikey_provider"),
)
op.create_index(
"idx_apikey_provider_enabled", "api_key_provider_mappings", ["api_key_id", "is_enabled"]
)
# ==================== provider_usage_tracking ====================
op.create_table(
"provider_usage_tracking",
sa.Column("id", sa.String(36), primary_key=True, index=True),
sa.Column(
"provider_id",
sa.String(36),
sa.ForeignKey("providers.id", ondelete="CASCADE"),
nullable=False,
index=True,
),
sa.Column("window_start", sa.DateTime(timezone=True), nullable=False, index=True),
sa.Column("window_end", sa.DateTime(timezone=True), nullable=False),
sa.Column("total_requests", sa.Integer, server_default="0"),
sa.Column("successful_requests", sa.Integer, server_default="0"),
sa.Column("failed_requests", sa.Integer, server_default="0"),
sa.Column("avg_response_time_ms", sa.Float, server_default="0.0"),
sa.Column("total_response_time_ms", sa.Float, server_default="0.0"),
sa.Column("total_cost_usd", sa.Float, server_default="0.0"),
sa.Column(
"created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.Column(
"updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
)
op.create_index(
"idx_provider_window", "provider_usage_tracking", ["provider_id", "window_start"]
)
op.create_index("idx_window_time", "provider_usage_tracking", ["window_start", "window_end"])
def downgrade() -> None:
# Drop tables in reverse order (respecting foreign key dependencies)
op.drop_table("provider_usage_tracking")
op.drop_table("api_key_provider_mappings")
op.drop_table("stats_user_daily")
op.drop_table("stats_summary")
op.drop_table("stats_daily")
op.drop_table("request_candidates")
op.drop_table("audit_logs")
op.drop_table("announcement_reads")
op.drop_table("announcements")
op.drop_table("user_preferences")
op.drop_table("system_configs")
op.drop_table("user_quotas")
op.drop_table("usage")
op.drop_table("provider_api_keys")
op.drop_table("model_mappings")
op.drop_table("models")
op.drop_table("provider_endpoints")
op.drop_table("api_keys")
op.drop_table("global_models")
op.drop_table("providers")
op.drop_table("users")
# Drop ENUM types
op.execute("DROP TYPE IF EXISTS providerbillingtype")
op.execute("DROP TYPE IF EXISTS userrole")
@@ -0,0 +1,315 @@
"""remove_model_mappings_add_aliases
合并迁移:
1. 添加 provider_model_aliases 字段到 models 表
2. 迁移 model_mappings 数据到 provider_model_aliases
3. 删除 model_mappings 表
4. 添加索引优化别名解析性能
Revision ID: e9b3d63f0cbf
Revises: 20251210_baseline
Create Date: 2025-12-14 13:00:22.828183+00:00
"""
import json
from datetime import datetime, timezone
import sqlalchemy as sa
from alembic import op
from sqlalchemy.orm import Session
# revision identifiers, used by Alembic.
revision = 'e9b3d63f0cbf'
down_revision = '20251210_baseline'
branch_labels = None
depends_on = None
def column_exists(bind, table_name: str, column_name: str) -> bool:
"""检查列是否存在"""
result = bind.execute(
sa.text(
"""
SELECT EXISTS (
SELECT 1 FROM information_schema.columns
WHERE table_name = :table_name AND column_name = :column_name
)
"""
),
{"table_name": table_name, "column_name": column_name},
)
return result.scalar()
def table_exists(bind, table_name: str) -> bool:
"""检查表是否存在"""
result = bind.execute(
sa.text(
"""
SELECT EXISTS (
SELECT 1 FROM information_schema.tables
WHERE table_name = :table_name
)
"""
),
{"table_name": table_name},
)
return result.scalar()
def index_exists(bind, index_name: str) -> bool:
"""检查索引是否存在"""
result = bind.execute(
sa.text(
"""
SELECT EXISTS (
SELECT 1 FROM pg_indexes
WHERE indexname = :index_name
)
"""
),
{"index_name": index_name},
)
return result.scalar()
def upgrade() -> None:
"""添加 provider_model_aliases 字段,迁移数据,删除 model_mappings 表"""
bind = op.get_bind()
# 1. 添加 provider_model_aliases 字段(如果不存在)
if not column_exists(bind, "models", "provider_model_aliases"):
op.add_column(
'models',
sa.Column('provider_model_aliases', sa.JSON(), nullable=True)
)
# 2. 迁移 model_mappings 数据(如果表存在)
session = Session(bind=bind)
model_mappings_table = sa.table(
"model_mappings",
sa.column("source_model", sa.String),
sa.column("target_global_model_id", sa.String),
sa.column("provider_id", sa.String),
sa.column("mapping_type", sa.String),
sa.column("is_active", sa.Boolean),
)
models_table = sa.table(
"models",
sa.column("id", sa.String),
sa.column("provider_id", sa.String),
sa.column("global_model_id", sa.String),
sa.column("provider_model_aliases", sa.JSON),
sa.column("updated_at", sa.DateTime(timezone=True)),
)
def normalize_alias_list(value) -> list[dict]:
"""将 DB 返回的 JSON 值规范化为 list[{'name': str, 'priority': int}]"""
if value is None:
return []
if isinstance(value, str):
try:
value = json.loads(value) if value else []
except Exception:
return []
if not isinstance(value, list):
return []
normalized: list[dict] = []
for item in value:
if not isinstance(item, dict):
continue
raw_name = item.get("name")
if not isinstance(raw_name, str):
continue
name = raw_name.strip()
if not name:
continue
raw_priority = item.get("priority", 1)
try:
priority = int(raw_priority)
except Exception:
priority = 1
if priority < 1:
priority = 1
normalized.append({"name": name, "priority": priority})
return normalized
# 查询所有活跃的 provider 级别 alias(只迁移 is_active=True 且 mapping_type='alias' 的)
# 全局别名/映射不迁移(新架构不再支持 source_model -> GlobalModel.name 的解析)
# 仅当 model_mappings 表存在时执行迁移
if table_exists(bind, "model_mappings"):
mappings = session.execute(
sa.select(
model_mappings_table.c.source_model,
model_mappings_table.c.target_global_model_id,
model_mappings_table.c.provider_id,
)
.where(
model_mappings_table.c.is_active.is_(True),
model_mappings_table.c.provider_id.isnot(None),
model_mappings_table.c.mapping_type == "alias",
)
.order_by(model_mappings_table.c.provider_id, model_mappings_table.c.source_model)
).all()
# 按 (provider_id, target_global_model_id) 分组,收集别名
alias_groups: dict = {}
for source_model, target_global_model_id, provider_id in mappings:
if not isinstance(source_model, str):
continue
source_model = source_model.strip()
if not source_model:
continue
if not isinstance(provider_id, str) or not provider_id:
continue
if not isinstance(target_global_model_id, str) or not target_global_model_id:
continue
key = (provider_id, target_global_model_id)
if key not in alias_groups:
alias_groups[key] = []
priority = len(alias_groups[key]) + 1
alias_groups[key].append({"name": source_model, "priority": priority})
# 更新对应的 models 记录
for (provider_id, global_model_id), aliases in alias_groups.items():
model_row = session.execute(
sa.select(models_table.c.id, models_table.c.provider_model_aliases)
.where(
models_table.c.provider_id == provider_id,
models_table.c.global_model_id == global_model_id,
)
.limit(1)
).first()
if model_row:
model_id = model_row[0]
existing_aliases = normalize_alias_list(model_row[1])
existing_names = {a["name"] for a in existing_aliases}
merged_aliases = list(existing_aliases)
for alias in aliases:
name = alias.get("name")
if not isinstance(name, str):
continue
name = name.strip()
if not name or name in existing_names:
continue
merged_aliases.append(
{
"name": name,
"priority": len(merged_aliases) + 1,
}
)
existing_names.add(name)
session.execute(
models_table.update()
.where(models_table.c.id == model_id)
.values(
provider_model_aliases=merged_aliases if merged_aliases else None,
updated_at=datetime.now(timezone.utc),
)
)
session.commit()
# 3. 删除 model_mappings 表
op.drop_table('model_mappings')
# 4. 添加索引优化别名解析性能
# provider_model_name 索引(支持精确匹配,如果不存在)
if not index_exists(bind, "idx_model_provider_model_name"):
op.create_index(
"idx_model_provider_model_name",
"models",
["provider_model_name"],
unique=False,
postgresql_where=sa.text("is_active = true"),
)
# provider_model_aliases GIN 索引(支持 JSONB 查询,仅 PostgreSQL)
if bind.dialect.name == "postgresql":
# 将 json 列转为 jsonb(jsonb 性能更好且支持 GIN 索引)
# 使用 IF NOT EXISTS 风格的检查来避免重复转换
op.execute(
"""
DO $$
BEGIN
IF EXISTS (
SELECT 1 FROM information_schema.columns
WHERE table_name = 'models'
AND column_name = 'provider_model_aliases'
AND data_type = 'json'
) THEN
ALTER TABLE models
ALTER COLUMN provider_model_aliases TYPE jsonb
USING provider_model_aliases::jsonb;
END IF;
END $$;
"""
)
# 创建 GIN 索引
op.execute(
"""
CREATE INDEX IF NOT EXISTS idx_model_provider_model_aliases_gin
ON models USING gin(provider_model_aliases jsonb_path_ops)
WHERE is_active = true
"""
)
def downgrade() -> None:
"""恢复 model_mappings 表,移除 provider_model_aliases 字段和索引"""
bind = op.get_bind()
# 1. 删除索引
op.drop_index("idx_model_provider_model_name", table_name="models")
if bind.dialect.name == "postgresql":
op.execute("DROP INDEX IF EXISTS idx_model_provider_model_aliases_gin")
# 将 jsonb 列还原为 json
op.execute(
"""
ALTER TABLE models
ALTER COLUMN provider_model_aliases TYPE json
USING provider_model_aliases::json
"""
)
# 2. 恢复 model_mappings 表
op.create_table(
'model_mappings',
sa.Column('id', sa.String(36), primary_key=True),
sa.Column('source_model', sa.String(200), nullable=False),
sa.Column(
'target_global_model_id',
sa.String(36),
sa.ForeignKey('global_models.id', ondelete='CASCADE'),
nullable=False,
),
sa.Column('provider_id', sa.String(36), sa.ForeignKey('providers.id'), nullable=True),
sa.Column('mapping_type', sa.String(20), nullable=False, server_default='alias'),
sa.Column('is_active', sa.Boolean(), nullable=False, server_default='true'),
sa.Column('created_at', sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
sa.Column('updated_at', sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
sa.UniqueConstraint('source_model', 'provider_id', name='uq_model_mapping_source_provider'),
)
op.create_index('ix_model_mappings_source_model', 'model_mappings', ['source_model'])
op.create_index('ix_model_mappings_target_global_model_id', 'model_mappings', ['target_global_model_id'])
op.create_index('ix_model_mappings_provider_id', 'model_mappings', ['provider_id'])
op.create_index('ix_model_mappings_mapping_type', 'model_mappings', ['mapping_type'])
# 3. 移除 provider_model_aliases 字段
op.drop_column('models', 'provider_model_aliases')
@@ -0,0 +1,47 @@
"""add first_byte_time_ms to usage table
Revision ID: 180e63a9c83a
Revises: e9b3d63f0cbf
Create Date: 2025-12-15 17:07:44.631032+00:00
"""
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision = '180e63a9c83a'
down_revision = 'e9b3d63f0cbf'
branch_labels = None
depends_on = None
def column_exists(bind, table_name: str, column_name: str) -> bool:
"""检查列是否存在"""
result = bind.execute(
sa.text(
"""
SELECT EXISTS (
SELECT 1 FROM information_schema.columns
WHERE table_name = :table_name AND column_name = :column_name
)
"""
),
{"table_name": table_name, "column_name": column_name},
)
return result.scalar()
def upgrade() -> None:
"""应用迁移:升级到新版本"""
bind = op.get_bind()
# 添加首字时间字段到 usage 表(如果不存在)
if not column_exists(bind, "usage", "first_byte_time_ms"):
op.add_column('usage', sa.Column('first_byte_time_ms', sa.Integer(), nullable=True))
def downgrade() -> None:
"""回滚迁移:降级到旧版本"""
# 删除首字时间字段
op.drop_column('usage', 'first_byte_time_ms')
@@ -0,0 +1,110 @@
"""refactor global_model to use config json field
Revision ID: 1cc6942cf06f
Revises: 180e63a9c83a
Create Date: 2025-12-16 03:11:32.480976+00:00
"""
import sqlalchemy as sa
from alembic import op
from sqlalchemy.dialects import postgresql
# revision identifiers, used by Alembic.
revision = '1cc6942cf06f'
down_revision = '180e63a9c83a'
branch_labels = None
depends_on = None
def column_exists(bind, table_name: str, column_name: str) -> bool:
"""检查列是否存在"""
result = bind.execute(
sa.text(
"""
SELECT EXISTS (
SELECT 1 FROM information_schema.columns
WHERE table_name = :table_name AND column_name = :column_name
)
"""
),
{"table_name": table_name, "column_name": column_name},
)
return result.scalar()
def upgrade() -> None:
"""应用迁移:升级到新版本
1. 添加 config 列
2. 把旧数据迁移到 config
3. 删除旧列
"""
bind = op.get_bind()
# 检查是否已经迁移过(config 列存在且旧列不存在)
has_config = column_exists(bind, "global_models", "config")
has_old_columns = column_exists(bind, "global_models", "default_supports_streaming")
if has_config and not has_old_columns:
# 已完成迁移,跳过
return
# 1. 添加 config 列(使用 JSONB 类型,支持索引和更高效的查询)
if not has_config:
op.add_column('global_models', sa.Column('config', postgresql.JSONB(), nullable=True))
# 2. 迁移数据:把旧字段合并到 config JSON(仅当旧列存在时)
if has_old_columns:
op.execute("""
UPDATE global_models
SET config = jsonb_strip_nulls(jsonb_build_object(
'streaming', COALESCE(default_supports_streaming, true),
'vision', CASE WHEN COALESCE(default_supports_vision, false) THEN true ELSE NULL END,
'function_calling', CASE WHEN COALESCE(default_supports_function_calling, false) THEN true ELSE NULL END,
'extended_thinking', CASE WHEN COALESCE(default_supports_extended_thinking, false) THEN true ELSE NULL END,
'image_generation', CASE WHEN COALESCE(default_supports_image_generation, false) THEN true ELSE NULL END,
'description', description,
'icon_url', icon_url,
'official_url', official_url
))
""")
# 3. 删除旧列
op.drop_column('global_models', 'default_supports_streaming')
op.drop_column('global_models', 'default_supports_vision')
op.drop_column('global_models', 'default_supports_function_calling')
op.drop_column('global_models', 'default_supports_extended_thinking')
op.drop_column('global_models', 'default_supports_image_generation')
op.drop_column('global_models', 'description')
op.drop_column('global_models', 'icon_url')
op.drop_column('global_models', 'official_url')
def downgrade() -> None:
"""回滚迁移:降级到旧版本"""
# 1. 添加旧列
op.add_column('global_models', sa.Column('icon_url', sa.VARCHAR(length=500), nullable=True))
op.add_column('global_models', sa.Column('official_url', sa.VARCHAR(length=500), nullable=True))
op.add_column('global_models', sa.Column('description', sa.TEXT(), nullable=True))
op.add_column('global_models', sa.Column('default_supports_streaming', sa.BOOLEAN(), nullable=True))
op.add_column('global_models', sa.Column('default_supports_vision', sa.BOOLEAN(), nullable=True))
op.add_column('global_models', sa.Column('default_supports_function_calling', sa.BOOLEAN(), nullable=True))
op.add_column('global_models', sa.Column('default_supports_extended_thinking', sa.BOOLEAN(), nullable=True))
op.add_column('global_models', sa.Column('default_supports_image_generation', sa.BOOLEAN(), nullable=True))
# 2. 从 config 恢复数据
op.execute("""
UPDATE global_models
SET
default_supports_streaming = COALESCE((config->>'streaming')::boolean, true),
default_supports_vision = COALESCE((config->>'vision')::boolean, false),
default_supports_function_calling = COALESCE((config->>'function_calling')::boolean, false),
default_supports_extended_thinking = COALESCE((config->>'extended_thinking')::boolean, false),
default_supports_image_generation = COALESCE((config->>'image_generation')::boolean, false),
description = config->>'description',
icon_url = config->>'icon_url',
official_url = config->>'official_url'
""")
# 3. 删除 config 列
op.drop_column('global_models', 'config')
@@ -0,0 +1,57 @@
"""add proxy field to provider_endpoints
Revision ID: f30f9936f6a2
Revises: 1cc6942cf06f
Create Date: 2025-12-18 06:31:58.451112+00:00
"""
from alembic import op
import sqlalchemy as sa
from sqlalchemy.dialects.postgresql import JSONB
from sqlalchemy import inspect
# revision identifiers, used by Alembic.
revision = 'f30f9936f6a2'
down_revision = '1cc6942cf06f'
branch_labels = None
depends_on = None
def column_exists(table_name: str, column_name: str) -> bool:
"""检查列是否存在"""
bind = op.get_bind()
inspector = inspect(bind)
columns = [col['name'] for col in inspector.get_columns(table_name)]
return column_name in columns
def get_column_type(table_name: str, column_name: str) -> str:
"""获取列的类型"""
bind = op.get_bind()
inspector = inspect(bind)
for col in inspector.get_columns(table_name):
if col['name'] == column_name:
return str(col['type']).upper()
return ''
def upgrade() -> None:
"""添加 proxy 字段到 provider_endpoints 表"""
if not column_exists('provider_endpoints', 'proxy'):
# 字段不存在,直接添加 JSONB 类型
op.add_column('provider_endpoints', sa.Column('proxy', JSONB(), nullable=True))
else:
# 字段已存在,检查是否需要转换类型
col_type = get_column_type('provider_endpoints', 'proxy')
if 'JSONB' not in col_type:
# 如果是 JSON 类型,转换为 JSONB
op.execute(
'ALTER TABLE provider_endpoints '
'ALTER COLUMN proxy TYPE JSONB USING proxy::jsonb'
)
def downgrade() -> None:
"""移除 proxy 字段"""
if column_exists('provider_endpoints', 'proxy'):
op.drop_column('provider_endpoints', 'proxy')
@@ -0,0 +1,86 @@
"""add stats_daily_model table and rename provider_model_aliases
Revision ID: a1b2c3d4e5f6
Revises: f30f9936f6a2
Create Date: 2025-12-20 12:00:00.000000+00:00
"""
from alembic import op
import sqlalchemy as sa
from sqlalchemy import inspect
# revision identifiers, used by Alembic.
revision = 'a1b2c3d4e5f6'
down_revision = 'f30f9936f6a2'
branch_labels = None
depends_on = None
def table_exists(table_name: str) -> bool:
"""检查表是否存在"""
bind = op.get_bind()
inspector = inspect(bind)
return table_name in inspector.get_table_names()
def column_exists(table_name: str, column_name: str) -> bool:
"""检查列是否存在"""
bind = op.get_bind()
inspector = inspect(bind)
columns = [col['name'] for col in inspector.get_columns(table_name)]
return column_name in columns
def upgrade() -> None:
"""创建 stats_daily_model 表,重命名 provider_model_aliases 为 provider_model_mappings"""
# 1. 创建 stats_daily_model 表
if not table_exists('stats_daily_model'):
op.create_table(
'stats_daily_model',
sa.Column('id', sa.String(36), primary_key=True),
sa.Column('date', sa.DateTime(timezone=True), nullable=False),
sa.Column('model', sa.String(100), nullable=False),
sa.Column('total_requests', sa.Integer(), nullable=False, default=0),
sa.Column('input_tokens', sa.BigInteger(), nullable=False, default=0),
sa.Column('output_tokens', sa.BigInteger(), nullable=False, default=0),
sa.Column('cache_creation_tokens', sa.BigInteger(), nullable=False, default=0),
sa.Column('cache_read_tokens', sa.BigInteger(), nullable=False, default=0),
sa.Column('total_cost', sa.Float(), nullable=False, default=0.0),
sa.Column('avg_response_time_ms', sa.Float(), nullable=False, default=0.0),
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False,
server_default=sa.func.now()),
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False,
server_default=sa.func.now(), onupdate=sa.func.now()),
sa.UniqueConstraint('date', 'model', name='uq_stats_daily_model'),
)
# 创建索引
op.create_index('idx_stats_daily_model_date', 'stats_daily_model', ['date'])
op.create_index('idx_stats_daily_model_date_model', 'stats_daily_model', ['date', 'model'])
# 2. 重命名 models 表的 provider_model_aliases 为 provider_model_mappings
if column_exists('models', 'provider_model_aliases') and not column_exists('models', 'provider_model_mappings'):
op.alter_column('models', 'provider_model_aliases', new_column_name='provider_model_mappings')
def index_exists(table_name: str, index_name: str) -> bool:
"""检查索引是否存在"""
bind = op.get_bind()
inspector = inspect(bind)
indexes = [idx['name'] for idx in inspector.get_indexes(table_name)]
return index_name in indexes
def downgrade() -> None:
"""删除 stats_daily_model 表,恢复 provider_model_aliases 列名"""
# 恢复列名
if column_exists('models', 'provider_model_mappings') and not column_exists('models', 'provider_model_aliases'):
op.alter_column('models', 'provider_model_mappings', new_column_name='provider_model_aliases')
# 删除表
if table_exists('stats_daily_model'):
if index_exists('stats_daily_model', 'idx_stats_daily_model_date_model'):
op.drop_index('idx_stats_daily_model_date_model', table_name='stats_daily_model')
if index_exists('stats_daily_model', 'idx_stats_daily_model_date'):
op.drop_index('idx_stats_daily_model_date', table_name='stats_daily_model')
op.drop_table('stats_daily_model')
@@ -0,0 +1,65 @@
"""add usage table composite indexes for query optimization
Revision ID: b2c3d4e5f6g7
Revises: a1b2c3d4e5f6
Create Date: 2025-12-20 15:00:00.000000+00:00
"""
from alembic import op
from sqlalchemy import text
# revision identifiers, used by Alembic.
revision = 'b2c3d4e5f6g7'
down_revision = 'a1b2c3d4e5f6'
branch_labels = None
depends_on = None
def upgrade() -> None:
"""为 usage 表添加复合索引以优化常见查询
注意:这些索引已经在 baseline 迁移中创建。
此迁移仅用于从旧版本升级的场景,新安装会跳过。
"""
conn = op.get_bind()
# 检查 usage 表是否存在
result = conn.execute(text(
"SELECT EXISTS (SELECT FROM information_schema.tables WHERE table_name = 'usage')"
))
if not result.scalar():
# 表不存在,跳过
return
# 定义需要创建的索引
indexes = [
("idx_usage_user_created", "ON usage (user_id, created_at)"),
("idx_usage_apikey_created", "ON usage (api_key_id, created_at)"),
("idx_usage_provider_model_created", "ON usage (provider, model, created_at)"),
]
# 分别检查并创建每个索引
for index_name, index_def in indexes:
result = conn.execute(text(
f"SELECT EXISTS (SELECT 1 FROM pg_indexes WHERE indexname = '{index_name}')"
))
if result.scalar():
continue # 索引已存在,跳过
conn.execute(text(f"CREATE INDEX {index_name} {index_def}"))
def downgrade() -> None:
"""删除复合索引"""
conn = op.get_bind()
# 使用 IF EXISTS 避免索引不存在时报错
conn.execute(text(
"DROP INDEX IF EXISTS idx_usage_provider_model_created"
))
conn.execute(text(
"DROP INDEX IF EXISTS idx_usage_apikey_created"
))
conn.execute(text(
"DROP INDEX IF EXISTS idx_usage_user_created"
))
@@ -0,0 +1,161 @@
"""add ldap authentication support
Revision ID: c3d4e5f6g7h8
Revises: b2c3d4e5f6g7
Create Date: 2026-01-01 14:00:00.000000+00:00
"""
from alembic import op
import sqlalchemy as sa
from sqlalchemy import text
# revision identifiers, used by Alembic.
revision = 'c3d4e5f6g7h8'
down_revision = 'b2c3d4e5f6g7'
branch_labels = None
depends_on = None
def _type_exists(conn, type_name: str) -> bool:
"""检查 PostgreSQL 类型是否存在"""
result = conn.execute(
text("SELECT 1 FROM pg_type WHERE typname = :name"),
{"name": type_name}
)
return result.scalar() is not None
def _column_exists(conn, table_name: str, column_name: str) -> bool:
"""检查列是否存在"""
result = conn.execute(
text("""
SELECT 1 FROM information_schema.columns
WHERE table_name = :table AND column_name = :column
"""),
{"table": table_name, "column": column_name}
)
return result.scalar() is not None
def _index_exists(conn, index_name: str) -> bool:
"""检查索引是否存在"""
result = conn.execute(
text("SELECT 1 FROM pg_indexes WHERE indexname = :name"),
{"name": index_name}
)
return result.scalar() is not None
def _table_exists(conn, table_name: str) -> bool:
"""检查表是否存在"""
result = conn.execute(
text("""
SELECT 1 FROM information_schema.tables
WHERE table_name = :name AND table_schema = 'public'
"""),
{"name": table_name}
)
return result.scalar() is not None
def upgrade() -> None:
"""添加 LDAP 认证支持
1. 创建 authsource 枚举类型
2. 在 users 表添加 auth_source 字段和 LDAP 标识字段
3. 创建 ldap_configs 表
"""
conn = op.get_bind()
# 1. 创建 authsource 枚举类型(幂等)
if not _type_exists(conn, 'authsource'):
conn.execute(text("CREATE TYPE authsource AS ENUM ('local', 'ldap')"))
# 2. 在 users 表添加字段(幂等)
if not _column_exists(conn, 'users', 'auth_source'):
op.add_column('users', sa.Column(
'auth_source',
sa.Enum('local', 'ldap', name='authsource', create_type=False),
nullable=False,
server_default='local'
))
if not _column_exists(conn, 'users', 'ldap_dn'):
op.add_column('users', sa.Column('ldap_dn', sa.String(length=512), nullable=True))
if not _column_exists(conn, 'users', 'ldap_username'):
op.add_column('users', sa.Column('ldap_username', sa.String(length=255), nullable=True))
# 创建索引(幂等)
if not _index_exists(conn, 'ix_users_ldap_dn'):
op.create_index('ix_users_ldap_dn', 'users', ['ldap_dn'])
if not _index_exists(conn, 'ix_users_ldap_username'):
op.create_index('ix_users_ldap_username', 'users', ['ldap_username'])
# 3. 创建 ldap_configs 表(幂等)
if not _table_exists(conn, 'ldap_configs'):
op.create_table(
'ldap_configs',
sa.Column('id', sa.Integer(), autoincrement=True, nullable=False),
sa.Column('server_url', sa.String(length=255), nullable=False),
sa.Column('bind_dn', sa.String(length=255), nullable=False),
sa.Column('bind_password_encrypted', sa.Text(), nullable=True),
sa.Column('base_dn', sa.String(length=255), nullable=False),
sa.Column('user_search_filter', sa.String(length=500), nullable=False, server_default='(uid={username})'),
sa.Column('username_attr', sa.String(length=50), nullable=False, server_default='uid'),
sa.Column('email_attr', sa.String(length=50), nullable=False, server_default='mail'),
sa.Column('display_name_attr', sa.String(length=50), nullable=False, server_default='cn'),
sa.Column('is_enabled', sa.Boolean(), nullable=False, server_default='false'),
sa.Column('is_exclusive', sa.Boolean(), nullable=False, server_default='false'),
sa.Column('use_starttls', sa.Boolean(), nullable=False, server_default='false'),
sa.Column('connect_timeout', sa.Integer(), nullable=False, server_default='10'),
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False, server_default=sa.text('now()')),
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False, server_default=sa.text('now()')),
sa.PrimaryKeyConstraint('id')
)
def downgrade() -> None:
"""回滚 LDAP 认证支持
警告:回滚前请确保:
1. 已备份数据库
2. 没有 LDAP 用户需要保留
"""
conn = op.get_bind()
# 检查是否存在 LDAP 用户,防止数据丢失
if _column_exists(conn, 'users', 'auth_source'):
result = conn.execute(text("SELECT COUNT(*) FROM users WHERE auth_source = 'ldap'"))
ldap_user_count = result.scalar()
if ldap_user_count and ldap_user_count > 0:
raise RuntimeError(
f"无法回滚:存在 {ldap_user_count} 个 LDAP 用户。"
f"请先删除或转换这些用户,或使用 --force 参数强制回滚(将丢失数据)。"
)
# 1. 删除 ldap_configs 表(幂等)
if _table_exists(conn, 'ldap_configs'):
op.drop_table('ldap_configs')
# 2. 删除 users 表的 LDAP 相关字段(幂等)
if _index_exists(conn, 'ix_users_ldap_username'):
op.drop_index('ix_users_ldap_username', table_name='users')
if _index_exists(conn, 'ix_users_ldap_dn'):
op.drop_index('ix_users_ldap_dn', table_name='users')
if _column_exists(conn, 'users', 'ldap_username'):
op.drop_column('users', 'ldap_username')
if _column_exists(conn, 'users', 'ldap_dn'):
op.drop_column('users', 'ldap_dn')
if _column_exists(conn, 'users', 'auth_source'):
op.drop_column('users', 'auth_source')
# 3. 删除 authsource 枚举类型(幂等)
# 注意:不使用 CASCADE,因为此时所有依赖应该已被删除
if _type_exists(conn, 'authsource'):
conn.execute(text("DROP TYPE authsource"))
@@ -0,0 +1,131 @@
"""add_management_tokens_table
Revision ID: ad55f1d008b7
Revises: c3d4e5f6g7h8
Create Date: 2026-01-06 15:24:10.660394+00:00
"""
from alembic import op
import sqlalchemy as sa
from sqlalchemy import inspect
# revision identifiers, used by Alembic.
revision = 'ad55f1d008b7'
down_revision = 'c3d4e5f6g7h8'
branch_labels = None
depends_on = None
def table_exists(table_name: str) -> bool:
"""检查表是否存在"""
conn = op.get_bind()
inspector = inspect(conn)
return table_name in inspector.get_table_names()
def index_exists(table_name: str, index_name: str) -> bool:
"""检查索引是否存在"""
conn = op.get_bind()
inspector = inspect(conn)
try:
indexes = inspector.get_indexes(table_name)
return any(idx["name"] == index_name for idx in indexes)
except Exception:
return False
def constraint_exists(table_name: str, constraint_name: str) -> bool:
"""检查约束是否存在"""
conn = op.get_bind()
inspector = inspect(conn)
try:
constraints = inspector.get_unique_constraints(table_name)
if any(c["name"] == constraint_name for c in constraints):
return True
# 也检查 check 约束
check_constraints = inspector.get_check_constraints(table_name)
if any(c["name"] == constraint_name for c in check_constraints):
return True
return False
except Exception:
return False
def upgrade() -> None:
"""应用迁移:创建 management_tokens 表"""
# 幂等性检查
if table_exists("management_tokens"):
# 表已存在,检查是否需要添加约束
if not constraint_exists("management_tokens", "uq_management_tokens_user_name"):
op.create_unique_constraint(
"uq_management_tokens_user_name",
"management_tokens",
["user_id", "name"],
)
# 添加 IP 白名单非空检查约束
if not constraint_exists("management_tokens", "check_allowed_ips_not_empty"):
op.create_check_constraint(
"check_allowed_ips_not_empty",
"management_tokens",
"allowed_ips IS NULL OR allowed_ips::text = 'null' OR json_array_length(allowed_ips) > 0",
)
return
op.create_table('management_tokens',
sa.Column('id', sa.String(length=36), nullable=False),
sa.Column('user_id', sa.String(length=36), nullable=False),
sa.Column('token_hash', sa.String(length=64), nullable=False),
sa.Column('token_prefix', sa.String(length=12), nullable=True),
sa.Column('name', sa.String(length=100), nullable=False),
sa.Column('description', sa.Text(), nullable=True),
sa.Column('allowed_ips', sa.JSON(), nullable=True),
sa.Column('expires_at', sa.DateTime(timezone=True), nullable=True),
sa.Column('last_used_at', sa.DateTime(timezone=True), nullable=True),
sa.Column('last_used_ip', sa.String(length=45), nullable=True),
sa.Column('usage_count', sa.Integer(), server_default='0', nullable=False),
sa.Column('is_active', sa.Boolean(), server_default='true', nullable=False),
sa.Column('created_at', sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
sa.Column('updated_at', sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
sa.ForeignKeyConstraint(['user_id'], ['users.id'], ondelete='CASCADE'),
sa.PrimaryKeyConstraint('id')
)
op.create_index('idx_management_tokens_is_active', 'management_tokens', ['is_active'], unique=False)
op.create_index('idx_management_tokens_user_id', 'management_tokens', ['user_id'], unique=False)
op.create_index(op.f('ix_management_tokens_token_hash'), 'management_tokens', ['token_hash'], unique=True)
# 添加用户名称唯一约束
op.create_unique_constraint(
"uq_management_tokens_user_name",
"management_tokens",
["user_id", "name"],
)
# 添加 IP 白名单非空检查约束
# 注意:JSON 类型的 NULL 可能被序列化为 JSON 'null',需要同时处理
op.create_check_constraint(
"check_allowed_ips_not_empty",
"management_tokens",
"allowed_ips IS NULL OR allowed_ips::text = 'null' OR json_array_length(allowed_ips) > 0",
)
def downgrade() -> None:
"""回滚迁移:删除 management_tokens 表"""
# 幂等性检查
if not table_exists("management_tokens"):
return
# 删除约束
if constraint_exists("management_tokens", "check_allowed_ips_not_empty"):
op.drop_constraint("check_allowed_ips_not_empty", "management_tokens", type_="check")
if constraint_exists("management_tokens", "uq_management_tokens_user_name"):
op.drop_constraint("uq_management_tokens_user_name", "management_tokens", type_="unique")
# 删除索引
if index_exists("management_tokens", "ix_management_tokens_token_hash"):
op.drop_index(op.f('ix_management_tokens_token_hash'), table_name='management_tokens')
if index_exists("management_tokens", "idx_management_tokens_user_id"):
op.drop_index('idx_management_tokens_user_id', table_name='management_tokens')
if index_exists("management_tokens", "idx_management_tokens_is_active"):
op.drop_index('idx_management_tokens_is_active', table_name='management_tokens')
# 删除表
op.drop_table('management_tokens')
@@ -0,0 +1,73 @@
"""cleanup ambiguous database fields
Revision ID: 02a45b66b7c4
Revises: ad55f1d008b7
Create Date: 2026-01-07 11:20:12.684426+00:00
变更内容:
1. users 表:重命名 allowed_endpoints 为 allowed_api_formats(修正历史命名错误)
2. api_keys 表:删除 allowed_endpoints 字段(未使用的功能)
3. providers 表:删除 rate_limit 字段(与 rpm_limit 功能重复,且未使用)
4. usage 表:重命名 provider 为 provider_name(避免与 provider_id 外键混淆)
"""
from alembic import op
import sqlalchemy as sa
from sqlalchemy import inspect
# revision identifiers, used by Alembic.
revision = '02a45b66b7c4'
down_revision = 'ad55f1d008b7'
branch_labels = None
depends_on = None
def _column_exists(table_name: str, column_name: str) -> bool:
"""检查列是否存在"""
bind = op.get_bind()
inspector = inspect(bind)
columns = [col['name'] for col in inspector.get_columns(table_name)]
return column_name in columns
def upgrade() -> None:
"""
1. users.allowed_endpoints -> allowed_api_formats(重命名)
2. api_keys.allowed_endpoints 删除
3. providers.rate_limit 删除(与 rpm_limit 重复)
4. usage.provider -> provider_name(重命名)
"""
# 1. users 表:重命名 allowed_endpoints 为 allowed_api_formats
if _column_exists('users', 'allowed_endpoints'):
op.alter_column('users', 'allowed_endpoints', new_column_name='allowed_api_formats')
# 2. api_keys 表:删除 allowed_endpoints 字段
if _column_exists('api_keys', 'allowed_endpoints'):
op.drop_column('api_keys', 'allowed_endpoints')
# 3. providers 表:删除 rate_limit 字段(与 rpm_limit 功能重复)
if _column_exists('providers', 'rate_limit'):
op.drop_column('providers', 'rate_limit')
# 4. usage 表:重命名 provider 为 provider_name
if _column_exists('usage', 'provider'):
op.alter_column('usage', 'provider', new_column_name='provider_name')
def downgrade() -> None:
"""回滚:恢复原字段"""
# 4. usage 表:将 provider_name 改回 provider
if _column_exists('usage', 'provider_name'):
op.alter_column('usage', 'provider_name', new_column_name='provider')
# 3. providers 表:恢复 rate_limit 字段
if not _column_exists('providers', 'rate_limit'):
op.add_column('providers', sa.Column('rate_limit', sa.Integer(), nullable=True))
# 2. api_keys 表:恢复 allowed_endpoints 字段
if not _column_exists('api_keys', 'allowed_endpoints'):
op.add_column('api_keys', sa.Column('allowed_endpoints', sa.JSON(), nullable=True))
# 1. users 表:将 allowed_api_formats 改回 allowed_endpoints
if _column_exists('users', 'allowed_api_formats'):
op.alter_column('users', 'allowed_api_formats', new_column_name='allowed_endpoints')
@@ -0,0 +1,604 @@
"""consolidated schema updates
Revision ID: m4n5o6p7q8r9
Revises: 02a45b66b7c4
Create Date: 2026-01-10 20:00:00.000000
This migration consolidates all schema changes from 2026-01-08 to 2026-01-10:
1. provider_api_keys: Key 直接关联 Provider (provider_id, api_formats)
2. provider_api_keys: 添加 rate_multipliers JSON 字段(按格式费率)
3. models: global_model_id 改为可空(支持独立 ProviderModel)
4. providers: 添加 timeout, max_retries, proxy(从 endpoint 迁移)
5. providers: display_name 重命名为 name,删除原 name
6. provider_api_keys: max_concurrent -> rpm_limit(并发改 RPM)
7. provider_api_keys: 健康度改为按格式存储(health_by_format, circuit_breaker_by_format)
8. provider_endpoints: 删除废弃的 rate_limit 列
9. usage: 添加 client_response_headers 字段
10. provider_api_keys: 删除 endpoint_id(Key 不再与 Endpoint 绑定)
11. provider_endpoints: 删除废弃的 max_concurrent 列
12. providers: 删除废弃的 rpm_limit, rpm_used, rpm_reset_at 列
"""
import logging
import sqlalchemy as sa
from sqlalchemy.dialects import postgresql
from sqlalchemy.exc import ProgrammingError
from alembic import op
# 配置日志
alembic_logger = logging.getLogger("alembic.runtime.migration")
revision = "m4n5o6p7q8r9"
down_revision = "02a45b66b7c4"
branch_labels = None
depends_on = None
def _column_exists(table_name: str, column_name: str) -> bool:
"""Check if a column exists in the table (bypasses inspector cache)"""
bind = op.get_bind()
result = bind.execute(
sa.text(
"SELECT 1 FROM information_schema.columns "
"WHERE table_name = :table AND column_name = :col"
),
{"table": table_name, "col": column_name},
)
return result.scalar() is not None
def _constraint_exists(table_name: str, constraint_name: str) -> bool:
"""Check if a constraint exists (bypasses inspector cache)"""
bind = op.get_bind()
result = bind.execute(
sa.text(
"SELECT 1 FROM information_schema.table_constraints "
"WHERE table_name = :table AND constraint_name = :name"
),
{"table": table_name, "name": constraint_name},
)
return result.scalar() is not None
def _index_exists(table_name: str, index_name: str) -> bool:
"""Check if an index exists (bypasses inspector cache)"""
bind = op.get_bind()
result = bind.execute(
sa.text("SELECT 1 FROM pg_indexes WHERE indexname = :name"),
{"name": index_name},
)
return result.scalar() is not None
def upgrade() -> None:
"""Apply all consolidated schema changes"""
bind = op.get_bind()
# ========== 1. provider_api_keys: 添加 provider_id 和 api_formats ==========
if not _column_exists("provider_api_keys", "provider_id"):
conn = op.get_bind()
conn.execute(sa.text("SAVEPOINT sp_add_provider_id"))
try:
op.add_column(
"provider_api_keys", sa.Column("provider_id", sa.String(36), nullable=True)
)
conn.execute(sa.text("RELEASE SAVEPOINT sp_add_provider_id"))
except ProgrammingError as exc:
if getattr(getattr(exc, "orig", None), "pgcode", None) == "42701":
conn.execute(sa.text("ROLLBACK TO SAVEPOINT sp_add_provider_id"))
alembic_logger.warning("provider_api_keys.provider_id already exists; skipping add")
else:
conn.execute(sa.text("ROLLBACK TO SAVEPOINT sp_add_provider_id"))
raise
# 数据迁移:从 endpoint 获取 provider_id(如果 endpoint_id 仍存在)
if _column_exists("provider_api_keys", "endpoint_id"):
op.execute("""
UPDATE provider_api_keys k
SET provider_id = e.provider_id
FROM provider_endpoints e
WHERE k.endpoint_id = e.id AND k.provider_id IS NULL
""")
# 检查无法关联的孤儿 Key
result = bind.execute(
sa.text("SELECT COUNT(*) FROM provider_api_keys WHERE provider_id IS NULL")
)
orphan_count = result.scalar() or 0
if orphan_count > 0:
# 使用 logger 记录更明显的告警
alembic_logger.warning("=" * 60)
alembic_logger.warning(
f"[MIGRATION WARNING] 发现 {orphan_count} 个无法关联 Provider 的孤儿 Key"
)
alembic_logger.warning("=" * 60)
alembic_logger.info("正在备份孤儿 Key 到 _orphan_api_keys_backup 表...")
# 先备份孤儿数据到临时表,避免数据丢失
op.execute("""
CREATE TABLE IF NOT EXISTS _orphan_api_keys_backup AS
SELECT *, NOW() as backup_at
FROM provider_api_keys
WHERE provider_id IS NULL
""")
# 记录备份的 Key ID
orphan_ids = bind.execute(
sa.text("SELECT id, name FROM provider_api_keys WHERE provider_id IS NULL")
).fetchall()
alembic_logger.info("备份的孤儿 Key 列表:")
for key_id, key_name in orphan_ids:
alembic_logger.info(f" - Key: {key_name} (ID: {key_id})")
# 删除孤儿数据
op.execute("DELETE FROM provider_api_keys WHERE provider_id IS NULL")
alembic_logger.info(f"已备份并删除 {orphan_count} 个孤儿 Key")
# 提供恢复指南
alembic_logger.warning("-" * 60)
alembic_logger.warning("[恢复指南] 如需恢复孤儿 Key:")
alembic_logger.warning(" 1. 查询备份表: SELECT * FROM _orphan_api_keys_backup;")
alembic_logger.warning(" 2. 确定正确的 provider_id")
alembic_logger.warning(" 3. 执行恢复:")
alembic_logger.warning(" INSERT INTO provider_api_keys (...)")
alembic_logger.warning(" SELECT ... FROM _orphan_api_keys_backup WHERE ...;")
alembic_logger.warning("-" * 60)
# 设置 NOT NULL 并创建外键
op.alter_column("provider_api_keys", "provider_id", nullable=False)
if not _constraint_exists("provider_api_keys", "fk_provider_api_keys_provider"):
op.create_foreign_key(
"fk_provider_api_keys_provider",
"provider_api_keys",
"providers",
["provider_id"],
["id"],
ondelete="CASCADE",
)
if not _index_exists("provider_api_keys", "idx_provider_api_keys_provider_id"):
op.create_index("idx_provider_api_keys_provider_id", "provider_api_keys", ["provider_id"])
if not _column_exists("provider_api_keys", "api_formats"):
op.add_column("provider_api_keys", sa.Column("api_formats", sa.JSON(), nullable=True))
# 数据迁移:从 endpoint 获取 api_format
op.execute("""
UPDATE provider_api_keys k
SET api_formats = json_build_array(e.api_format)
FROM provider_endpoints e
WHERE k.endpoint_id = e.id AND k.api_formats IS NULL
""")
op.alter_column("provider_api_keys", "api_formats", nullable=False, server_default="[]")
# 修改 endpoint_id 为可空,外键改为 SET NULL
if _constraint_exists("provider_api_keys", "provider_api_keys_endpoint_id_fkey"):
op.drop_constraint(
"provider_api_keys_endpoint_id_fkey", "provider_api_keys", type_="foreignkey"
)
op.alter_column("provider_api_keys", "endpoint_id", nullable=True)
# 不再重建外键,因为后面会删除这个字段
# ========== 2. provider_api_keys: 添加 rate_multipliers ==========
if not _column_exists("provider_api_keys", "rate_multipliers"):
op.add_column(
"provider_api_keys",
sa.Column("rate_multipliers", postgresql.JSON(astext_type=sa.Text()), nullable=True),
)
# 数据迁移:将 rate_multiplier 按 api_formats 转换
op.execute("""
UPDATE provider_api_keys
SET rate_multipliers = (
SELECT jsonb_object_agg(elem, rate_multiplier)
FROM jsonb_array_elements_text(api_formats::jsonb) AS elem
)
WHERE api_formats IS NOT NULL
AND api_formats::text != '[]'
AND api_formats::text != 'null'
AND rate_multipliers IS NULL
""")
# ========== 3. models: global_model_id 改为可空 ==========
op.alter_column("models", "global_model_id", existing_type=sa.String(36), nullable=True)
# ========== 4. providers: 添加 timeout, max_retries, proxy ==========
if not _column_exists("providers", "timeout"):
op.add_column(
"providers",
sa.Column("timeout", sa.Integer(), nullable=True, comment="请求超时(秒)"),
)
if not _column_exists("providers", "max_retries"):
op.add_column(
"providers",
sa.Column("max_retries", sa.Integer(), nullable=True, comment="最大重试次数"),
)
if not _column_exists("providers", "proxy"):
op.add_column(
"providers",
sa.Column("proxy", postgresql.JSONB(), nullable=True, comment="代理配置"),
)
# 从端点迁移数据到 provider(动态构建 SQL,仅引用存在的列)
ep_has_timeout = _column_exists("provider_endpoints", "timeout")
ep_has_max_retries = _column_exists("provider_endpoints", "max_retries")
ep_has_proxy = _column_exists("provider_endpoints", "proxy")
set_clauses = []
if _column_exists("providers", "timeout"):
if ep_has_timeout:
set_clauses.append("""
timeout = COALESCE(
p.timeout,
(SELECT MAX(e.timeout) FROM provider_endpoints e WHERE e.provider_id = p.id AND e.timeout IS NOT NULL),
300
)""")
else:
set_clauses.append("timeout = COALESCE(p.timeout, 300)")
if _column_exists("providers", "max_retries"):
if ep_has_max_retries:
set_clauses.append("""
max_retries = COALESCE(
p.max_retries,
(SELECT MAX(e.max_retries) FROM provider_endpoints e WHERE e.provider_id = p.id AND e.max_retries IS NOT NULL),
2
)""")
else:
set_clauses.append("max_retries = COALESCE(p.max_retries, 2)")
if _column_exists("providers", "proxy") and ep_has_proxy:
set_clauses.append("""
proxy = COALESCE(
p.proxy,
(SELECT e.proxy FROM provider_endpoints e WHERE e.provider_id = p.id AND e.proxy IS NOT NULL ORDER BY e.created_at LIMIT 1)
)""")
if set_clauses:
where_parts = []
if _column_exists("providers", "timeout"):
where_parts.append("p.timeout IS NULL")
if _column_exists("providers", "max_retries"):
where_parts.append("p.max_retries IS NULL")
where_clause = " OR ".join(where_parts) if where_parts else "TRUE"
sql = "UPDATE providers p SET " + ", ".join(set_clauses) + " WHERE " + where_clause
op.execute(sql)
# ========== 5. providers: display_name -> name ==========
# 注意:这里假设 display_name 已经被重命名为 name
# 如果 display_name 仍然存在,则需要执行重命名
if _column_exists("providers", "display_name"):
# 删除旧的 name 索引
if _index_exists("providers", "ix_providers_name"):
op.drop_index("ix_providers_name", table_name="providers")
# 如果存在旧的 name 列,先删除
if _column_exists("providers", "name"):
op.drop_column("providers", "name")
# 重命名 display_name 为 name
op.alter_column("providers", "display_name", new_column_name="name")
# 创建新索引
op.create_index("ix_providers_name", "providers", ["name"], unique=True)
# ========== 6. provider_api_keys: max_concurrent -> rpm_limit ==========
if _column_exists("provider_api_keys", "max_concurrent"):
op.alter_column("provider_api_keys", "max_concurrent", new_column_name="rpm_limit")
if _column_exists("provider_api_keys", "learned_max_concurrent"):
op.alter_column(
"provider_api_keys", "learned_max_concurrent", new_column_name="learned_rpm_limit"
)
if _column_exists("provider_api_keys", "last_concurrent_peak"):
op.alter_column(
"provider_api_keys", "last_concurrent_peak", new_column_name="last_rpm_peak"
)
# 删除废弃字段
for col in ["rate_limit", "daily_limit", "monthly_limit"]:
if _column_exists("provider_api_keys", col):
op.drop_column("provider_api_keys", col)
# ========== 7. provider_api_keys: 健康度改为按格式存储 ==========
if not _column_exists("provider_api_keys", "health_by_format"):
op.add_column(
"provider_api_keys",
sa.Column(
"health_by_format",
postgresql.JSONB(astext_type=sa.Text()),
nullable=True,
comment="按API格式存储的健康度数据",
),
)
if not _column_exists("provider_api_keys", "circuit_breaker_by_format"):
op.add_column(
"provider_api_keys",
sa.Column(
"circuit_breaker_by_format",
postgresql.JSONB(astext_type=sa.Text()),
nullable=True,
comment="按API格式存储的熔断器状态",
),
)
# 数据迁移:如果存在旧字段,迁移数据到新结构
if _column_exists("provider_api_keys", "health_score"):
op.execute("""
UPDATE provider_api_keys
SET health_by_format = (
SELECT jsonb_object_agg(
elem,
jsonb_build_object(
'health_score', COALESCE(health_score, 1.0),
'consecutive_failures', COALESCE(consecutive_failures, 0),
'last_failure_at', last_failure_at,
'request_results_window', COALESCE(request_results_window::jsonb, '[]'::jsonb)
)
)
FROM jsonb_array_elements_text(api_formats::jsonb) AS elem
)
WHERE api_formats IS NOT NULL
AND api_formats::text != '[]'
AND health_by_format IS NULL
""")
# Circuit Breaker 迁移策略:
# 不复制旧的 circuit_breaker_open 状态到所有 format,而是全部重置为 closed
# 原因:旧的单一 circuit breaker 状态可能因某一个 format 失败而打开,
# 如果复制到所有 format,会导致其他正常工作的 format 被错误标记为不可用
if _column_exists("provider_api_keys", "circuit_breaker_open"):
op.execute("""
UPDATE provider_api_keys
SET circuit_breaker_by_format = (
SELECT jsonb_object_agg(
elem,
jsonb_build_object(
'open', false,
'open_at', NULL,
'next_probe_at', NULL,
'half_open_until', NULL,
'half_open_successes', 0,
'half_open_failures', 0
)
)
FROM jsonb_array_elements_text(api_formats::jsonb) AS elem
)
WHERE api_formats IS NOT NULL
AND api_formats::text != '[]'
AND circuit_breaker_by_format IS NULL
""")
# 设置默认空对象
op.execute("""
UPDATE provider_api_keys
SET health_by_format = '{}'::jsonb
WHERE health_by_format IS NULL
""")
op.execute("""
UPDATE provider_api_keys
SET circuit_breaker_by_format = '{}'::jsonb
WHERE circuit_breaker_by_format IS NULL
""")
# 创建 GIN 索引
if not _index_exists("provider_api_keys", "ix_provider_api_keys_health_by_format"):
op.create_index(
"ix_provider_api_keys_health_by_format",
"provider_api_keys",
["health_by_format"],
postgresql_using="gin",
)
if not _index_exists("provider_api_keys", "ix_provider_api_keys_circuit_breaker_by_format"):
op.create_index(
"ix_provider_api_keys_circuit_breaker_by_format",
"provider_api_keys",
["circuit_breaker_by_format"],
postgresql_using="gin",
)
# 删除旧字段
old_health_columns = [
"health_score",
"consecutive_failures",
"last_failure_at",
"request_results_window",
"circuit_breaker_open",
"circuit_breaker_open_at",
"next_probe_at",
"half_open_until",
"half_open_successes",
"half_open_failures",
]
for col in old_health_columns:
if _column_exists("provider_api_keys", col):
op.drop_column("provider_api_keys", col)
# ========== 8. provider_endpoints: 删除废弃的 rate_limit 列 ==========
if _column_exists("provider_endpoints", "rate_limit"):
op.drop_column("provider_endpoints", "rate_limit")
# ========== 9. usage: 添加 client_response_headers ==========
if not _column_exists("usage", "client_response_headers"):
op.add_column(
"usage",
sa.Column("client_response_headers", sa.JSON(), nullable=True),
)
# ========== 10. provider_api_keys: 删除 endpoint_id ==========
# Key 不再与 Endpoint 绑定,通过 provider_id + api_formats 关联
if _column_exists("provider_api_keys", "endpoint_id"):
# 查找 endpoint_id 上的外键并删除(用 savepoint 保护,避免事务中止)
conn = op.get_bind()
fk_rows = conn.execute(
sa.text(
"SELECT con.conname FROM pg_constraint con "
"JOIN pg_attribute att ON att.attnum = ANY(con.conkey) "
" AND att.attrelid = con.conrelid "
"WHERE con.conrelid = 'provider_api_keys'::regclass "
" AND con.contype = 'f' AND att.attname = 'endpoint_id'"
)
).fetchall()
for (fk_name,) in fk_rows:
conn.execute(sa.text(f"SAVEPOINT sp_drop_fk_{fk_name}"))
try:
op.drop_constraint(fk_name, "provider_api_keys", type_="foreignkey")
conn.execute(sa.text(f"RELEASE SAVEPOINT sp_drop_fk_{fk_name}"))
except Exception:
conn.execute(sa.text(f"ROLLBACK TO SAVEPOINT sp_drop_fk_{fk_name}"))
op.drop_column("provider_api_keys", "endpoint_id")
# ========== 11. provider_endpoints: 删除废弃的 max_concurrent 列 ==========
if _column_exists("provider_endpoints", "max_concurrent"):
op.drop_column("provider_endpoints", "max_concurrent")
# ========== 12. providers: 删除废弃的 RPM 相关字段 ==========
if _column_exists("providers", "rpm_limit"):
op.drop_column("providers", "rpm_limit")
if _column_exists("providers", "rpm_used"):
op.drop_column("providers", "rpm_used")
if _column_exists("providers", "rpm_reset_at"):
op.drop_column("providers", "rpm_reset_at")
alembic_logger.info("[OK] Consolidated migration completed successfully")
def downgrade() -> None:
"""
Downgrade is complex due to data migrations.
For safety, this only removes new columns without restoring old structure.
Manual intervention may be required for full rollback.
"""
bind = op.get_bind()
# 12. 恢复 providers RPM 相关字段
if not _column_exists("providers", "rpm_limit"):
op.add_column("providers", sa.Column("rpm_limit", sa.Integer(), nullable=True))
if not _column_exists("providers", "rpm_used"):
op.add_column(
"providers",
sa.Column("rpm_used", sa.Integer(), server_default="0", nullable=True),
)
if not _column_exists("providers", "rpm_reset_at"):
op.add_column(
"providers",
sa.Column("rpm_reset_at", sa.DateTime(timezone=True), nullable=True),
)
# 11. 恢复 provider_endpoints.max_concurrent
if not _column_exists("provider_endpoints", "max_concurrent"):
op.add_column(
"provider_endpoints", sa.Column("max_concurrent", sa.Integer(), nullable=True)
)
# 10. 恢复 endpoint_id
if not _column_exists("provider_api_keys", "endpoint_id"):
op.add_column("provider_api_keys", sa.Column("endpoint_id", sa.String(36), nullable=True))
# 9. 删除 client_response_headers
if _column_exists("usage", "client_response_headers"):
op.drop_column("usage", "client_response_headers")
# 8. 恢复 provider_endpoints.rate_limit(如果需要)
if not _column_exists("provider_endpoints", "rate_limit"):
op.add_column("provider_endpoints", sa.Column("rate_limit", sa.Integer(), nullable=True))
# 7. 删除健康度 JSON 字段
bind.execute(sa.text("DROP INDEX IF EXISTS ix_provider_api_keys_health_by_format"))
bind.execute(sa.text("DROP INDEX IF EXISTS ix_provider_api_keys_circuit_breaker_by_format"))
if _column_exists("provider_api_keys", "health_by_format"):
op.drop_column("provider_api_keys", "health_by_format")
if _column_exists("provider_api_keys", "circuit_breaker_by_format"):
op.drop_column("provider_api_keys", "circuit_breaker_by_format")
# 6. rpm_limit -> max_concurrent(简化版:仅重命名)
if _column_exists("provider_api_keys", "rpm_limit"):
op.alter_column("provider_api_keys", "rpm_limit", new_column_name="max_concurrent")
if _column_exists("provider_api_keys", "learned_rpm_limit"):
op.alter_column(
"provider_api_keys", "learned_rpm_limit", new_column_name="learned_max_concurrent"
)
if _column_exists("provider_api_keys", "last_rpm_peak"):
op.alter_column(
"provider_api_keys", "last_rpm_peak", new_column_name="last_concurrent_peak"
)
# 恢复已删除的字段
if not _column_exists("provider_api_keys", "rate_limit"):
op.add_column("provider_api_keys", sa.Column("rate_limit", sa.Integer(), nullable=True))
if not _column_exists("provider_api_keys", "daily_limit"):
op.add_column("provider_api_keys", sa.Column("daily_limit", sa.Integer(), nullable=True))
if not _column_exists("provider_api_keys", "monthly_limit"):
op.add_column("provider_api_keys", sa.Column("monthly_limit", sa.Integer(), nullable=True))
# 5. name -> display_name (需要先删除索引)
if _column_exists("providers", "name") and not _column_exists("providers", "display_name"):
if _index_exists("providers", "ix_providers_name"):
op.drop_index("ix_providers_name", table_name="providers")
op.alter_column("providers", "name", new_column_name="display_name")
if not _column_exists("providers", "name"):
op.add_column("providers", sa.Column("name", sa.String(100), nullable=True))
op.execute("""
UPDATE providers
SET name = LOWER(REPLACE(REPLACE(display_name, ' ', '_'), '-', '_'))
""")
op.alter_column("providers", "name", nullable=False)
if not _index_exists("providers", "ix_providers_name"):
op.create_index("ix_providers_name", "providers", ["name"], unique=True)
# 4. 删除 providers 的 timeout, max_retries, proxy
if _column_exists("providers", "proxy"):
op.drop_column("providers", "proxy")
if _column_exists("providers", "max_retries"):
op.drop_column("providers", "max_retries")
if _column_exists("providers", "timeout"):
op.drop_column("providers", "timeout")
# 3. models: global_model_id 改回 NOT NULL
result = bind.execute(sa.text("SELECT COUNT(*) FROM models WHERE global_model_id IS NULL"))
orphan_model_count = result.scalar() or 0
if orphan_model_count > 0:
alembic_logger.warning(
f"[WARN] 发现 {orphan_model_count} 个无 global_model_id 的独立模型,将被删除"
)
op.execute("DELETE FROM models WHERE global_model_id IS NULL")
alembic_logger.info(f"已删除 {orphan_model_count} 个独立模型")
op.alter_column("models", "global_model_id", nullable=False)
# 2. 删除 rate_multipliers
if _column_exists("provider_api_keys", "rate_multipliers"):
op.drop_column("provider_api_keys", "rate_multipliers")
# 1. 删除 provider_id 和 api_formats
if _index_exists("provider_api_keys", "idx_provider_api_keys_provider_id"):
op.drop_index("idx_provider_api_keys_provider_id", table_name="provider_api_keys")
if _constraint_exists("provider_api_keys", "fk_provider_api_keys_provider"):
op.drop_constraint("fk_provider_api_keys_provider", "provider_api_keys", type_="foreignkey")
if _column_exists("provider_api_keys", "api_formats"):
op.drop_column("provider_api_keys", "api_formats")
if _column_exists("provider_api_keys", "provider_id"):
op.drop_column("provider_api_keys", "provider_id")
# 恢复 endpoint_id 外键(简化版:仅创建外键,不强制 NOT NULL)
if _column_exists("provider_api_keys", "endpoint_id"):
if not _constraint_exists("provider_api_keys", "provider_api_keys_endpoint_id_fkey"):
op.create_foreign_key(
"provider_api_keys_endpoint_id_fkey",
"provider_api_keys",
"provider_endpoints",
["endpoint_id"],
["id"],
ondelete="SET NULL",
)
alembic_logger.info("[OK] Downgrade completed (simplified version)")
@@ -0,0 +1,95 @@
"""add auto_fetch_models and locked_models to provider_api_keys
Revision ID: e4ebe3233b40
Revises: m4n5o6p7q8r9
Create Date: 2026-01-13 17:59:53.119479+00:00
为 provider_api_keys 表添加自动获取模型相关字段:
1. auto_fetch_models: 是否启用自动获取模型
2. last_models_fetch_at: 最后获取时间
3. last_models_fetch_error: 最后获取错误信息
4. locked_models: 被锁定的模型列表(刷新时不会被删除)
注意: downgrade 操作会永久删除 auto_fetch_models 配置和 locked_models 数据
"""
from alembic import op
import sqlalchemy as sa
from sqlalchemy import inspect
def _index_exists(index_name: str) -> bool:
"""Check if an index exists"""
bind = op.get_bind()
inspector = inspect(bind)
indexes = inspector.get_indexes("provider_api_keys")
return any(idx["name"] == index_name for idx in indexes)
# revision identifiers, used by Alembic.
revision = 'e4ebe3233b40'
down_revision = 'm4n5o6p7q8r9'
branch_labels = None
depends_on = None
def _column_exists(table_name: str, column_name: str) -> bool:
"""Check if a column exists in the table"""
bind = op.get_bind()
inspector = inspect(bind)
columns = [col["name"] for col in inspector.get_columns(table_name)]
return column_name in columns
def upgrade() -> None:
"""添加自动获取模型相关字段"""
if not _column_exists("provider_api_keys", "auto_fetch_models"):
op.add_column(
"provider_api_keys",
sa.Column("auto_fetch_models", sa.Boolean(), nullable=False, server_default="false"),
)
if not _column_exists("provider_api_keys", "last_models_fetch_at"):
op.add_column(
"provider_api_keys",
sa.Column("last_models_fetch_at", sa.DateTime(timezone=True), nullable=True),
)
if not _column_exists("provider_api_keys", "last_models_fetch_error"):
op.add_column(
"provider_api_keys",
sa.Column("last_models_fetch_error", sa.Text(), nullable=True),
)
if not _column_exists("provider_api_keys", "locked_models"):
op.add_column(
"provider_api_keys",
sa.Column("locked_models", sa.JSON(), nullable=True),
)
# 添加复合索引以优化调度器查询
if not _index_exists("ix_provider_api_keys_auto_fetch_active"):
op.create_index(
"ix_provider_api_keys_auto_fetch_active",
"provider_api_keys",
["auto_fetch_models", "is_active"],
postgresql_where=sa.text("auto_fetch_models = true AND is_active = true"),
)
def downgrade() -> None:
"""移除自动获取模型相关字段"""
# 先删除索引
if _index_exists("ix_provider_api_keys_auto_fetch_active"):
op.drop_index("ix_provider_api_keys_auto_fetch_active", table_name="provider_api_keys")
if _column_exists("provider_api_keys", "locked_models"):
op.drop_column("provider_api_keys", "locked_models")
if _column_exists("provider_api_keys", "last_models_fetch_error"):
op.drop_column("provider_api_keys", "last_models_fetch_error")
if _column_exists("provider_api_keys", "last_models_fetch_at"):
op.drop_column("provider_api_keys", "last_models_fetch_at")
if _column_exists("provider_api_keys", "auto_fetch_models"):
op.drop_column("provider_api_keys", "auto_fetch_models")
@@ -0,0 +1,104 @@
"""add header_rules to provider_endpoints and is_locked to api_keys
Revision ID: 6d579000e511
Revises: e4ebe3233b40
Create Date: 2026-01-15 23:00:00.000000+00:00
变更:
1. provider_endpoints 表: 添加 header_rules 字段,迁移 headers 数据
2. api_keys 表: 添加 is_locked 字段(管理员锁定标志)
"""
from alembic import op
import sqlalchemy as sa
from sqlalchemy.dialects.postgresql import JSON
# revision identifiers, used by Alembic.
revision = '6d579000e511'
down_revision = 'e4ebe3233b40'
branch_labels = None
depends_on = None
def _column_exists(connection, table: str, column: str) -> bool:
"""检查列是否存在"""
result = connection.execute(
sa.text("""
SELECT 1 FROM information_schema.columns
WHERE table_name = :table AND column_name = :column
"""),
{"table": table, "column": column}
)
return result.fetchone() is not None
def upgrade() -> None:
"""添加 header_rules 字段并迁移现有 headers 数据;添加 is_locked 字段"""
connection = op.get_bind()
# ========== provider_endpoints.header_rules ==========
# 1. 添加 header_rules 列(幂等)
if not _column_exists(connection, 'provider_endpoints', 'header_rules'):
op.add_column('provider_endpoints', sa.Column('header_rules', JSON, nullable=True))
# 2. 批量迁移:headers -> header_rules
# 使用纯 SQL 将 {"k1":"v1", "k2":"v2"} 转换为 [{"action":"set","key":"k1","value":"v1"}, ...]
if _column_exists(connection, 'provider_endpoints', 'headers'):
connection.execute(
sa.text("""
UPDATE provider_endpoints
SET header_rules = (
SELECT jsonb_agg(
jsonb_build_object('action', 'set', 'key', key, 'value', value)
)
FROM jsonb_each_text(headers::jsonb)
)
WHERE headers IS NOT NULL
AND headers::text != '{}'
AND jsonb_typeof(headers::jsonb) = 'object'
AND header_rules IS NULL
""")
)
# 3. 删除旧列
op.drop_column('provider_endpoints', 'headers')
# ========== api_keys.is_locked ==========
if not _column_exists(connection, 'api_keys', 'is_locked'):
op.add_column(
'api_keys',
sa.Column('is_locked', sa.Boolean(), nullable=False, server_default='false')
)
def downgrade() -> None:
"""移除 header_rules 字段,恢复 headers 字段;移除 is_locked 字段"""
connection = op.get_bind()
# ========== api_keys.is_locked ==========
if _column_exists(connection, 'api_keys', 'is_locked'):
op.drop_column('api_keys', 'is_locked')
# ========== provider_endpoints.header_rules ==========
# 1. 添加 headers 列(幂等)
if not _column_exists(connection, 'provider_endpoints', 'headers'):
op.add_column('provider_endpoints', sa.Column('headers', JSON, nullable=True))
# 2. 批量迁移:header_rules -> headers(仅提取 set 操作)
if _column_exists(connection, 'provider_endpoints', 'header_rules'):
connection.execute(
sa.text("""
UPDATE provider_endpoints
SET headers = (
SELECT jsonb_object_agg(rule->>'key', rule->>'value')
FROM jsonb_array_elements(header_rules::jsonb) AS rule
WHERE rule->>'action' = 'set'
AND rule->>'key' IS NOT NULL
)
WHERE header_rules IS NOT NULL
AND jsonb_typeof(header_rules::jsonb) = 'array'
AND jsonb_array_length(header_rules::jsonb) > 0
""")
)
# 3. 删除 header_rules 列
op.drop_column('provider_endpoints', 'header_rules')
@@ -0,0 +1,127 @@
"""add global_priority_by_format and remove deprecated fields
Revision ID: ddd59cdf0349
Revises: 6d579000e511
Create Date: 2026-01-16 12:00:00.000000+00:00
变更:
1. provider_api_keys 表: 添加 global_priority_by_format 字段(按 API 格式的全局优先级)
2. 迁移现有 global_priority 数据到新字段
3. 删除已废弃的 global_priority 字段
4. 删除已废弃的 rate_multiplier 字段(已被 rate_multipliers 替代)
5. 删除已废弃的 providers.timeout 字段(由环境变量控制)
6. 删除已废弃的 provider_endpoints.timeout 字段(由环境变量控制)
"""
from alembic import op
import sqlalchemy as sa
from sqlalchemy.dialects.postgresql import JSON
# revision identifiers, used by Alembic.
revision = 'ddd59cdf0349'
down_revision = '6d579000e511'
branch_labels = None
depends_on = None
def _column_exists(connection, table: str, column: str) -> bool:
"""检查列是否存在"""
result = connection.execute(
sa.text("""
SELECT 1 FROM information_schema.columns
WHERE table_name = :table AND column_name = :column
"""),
{"table": table, "column": column}
)
return result.fetchone() is not None
def upgrade():
connection = op.get_bind()
# 1. 添加 global_priority_by_format 字段
if not _column_exists(connection, 'provider_api_keys', 'global_priority_by_format'):
op.add_column(
'provider_api_keys',
sa.Column('global_priority_by_format', JSON, nullable=True)
)
# 2. 迁移现有 global_priority 数据到新字段
# 对于有 global_priority 的 Key,将其值应用到所有支持的 api_formats
if _column_exists(connection, 'provider_api_keys', 'global_priority'):
# 将 JSON 数组转换为 text[] 后使用 unnest
connection.execute(sa.text("""
UPDATE provider_api_keys
SET global_priority_by_format = (
SELECT jsonb_object_agg(format, global_priority)
FROM jsonb_array_elements_text(api_formats::jsonb) AS format
)
WHERE global_priority IS NOT NULL
AND api_formats IS NOT NULL
AND jsonb_array_length(api_formats::jsonb) > 0
AND global_priority_by_format IS NULL
"""))
# 3. 删除 global_priority 字段
op.drop_column('provider_api_keys', 'global_priority')
# 4. 删除 rate_multiplier 字段(已被 rate_multipliers 替代)
if _column_exists(connection, 'provider_api_keys', 'rate_multiplier'):
op.drop_column('provider_api_keys', 'rate_multiplier')
# 5. 删除 providers.timeout 字段(由环境变量控制)
if _column_exists(connection, 'providers', 'timeout'):
op.drop_column('providers', 'timeout')
# 6. 删除 provider_endpoints.timeout 字段(由环境变量控制)
if _column_exists(connection, 'provider_endpoints', 'timeout'):
op.drop_column('provider_endpoints', 'timeout')
def downgrade():
connection = op.get_bind()
# 1. 恢复 rate_multiplier 字段
if not _column_exists(connection, 'provider_api_keys', 'rate_multiplier'):
op.add_column(
'provider_api_keys',
sa.Column('rate_multiplier', sa.Float, nullable=False, server_default='1.0')
)
# 2. 恢复 global_priority 字段并迁移数据
if not _column_exists(connection, 'provider_api_keys', 'global_priority'):
op.add_column(
'provider_api_keys',
sa.Column('global_priority', sa.Integer, nullable=True)
)
# 从 global_priority_by_format 迁移数据(取第一个格式的优先级值)
if _column_exists(connection, 'provider_api_keys', 'global_priority_by_format'):
connection.execute(sa.text("""
UPDATE provider_api_keys
SET global_priority = (
SELECT (value::text)::integer
FROM jsonb_each(global_priority_by_format::jsonb)
LIMIT 1
)
WHERE global_priority_by_format IS NOT NULL
AND jsonb_typeof(global_priority_by_format::jsonb) = 'object'
AND global_priority IS NULL
"""))
# 3. 删除 global_priority_by_format 字段
if _column_exists(connection, 'provider_api_keys', 'global_priority_by_format'):
op.drop_column('provider_api_keys', 'global_priority_by_format')
# 4. 恢复 providers.timeout 字段
if not _column_exists(connection, 'providers', 'timeout'):
op.add_column(
'providers',
sa.Column('timeout', sa.Integer, nullable=True, server_default='300')
)
# 5. 恢复 provider_endpoints.timeout 字段
if not _column_exists(connection, 'provider_endpoints', 'timeout'):
op.add_column(
'provider_endpoints',
sa.Column('timeout', sa.Integer, nullable=True, server_default='300')
)
@@ -0,0 +1,223 @@
"""make users email/password nullable add email_verified and oauth tables
Revision ID: 33e347f97c0c
Revises: ddd59cdf0349
Create Date: 2026-01-18 11:18:15.940559+00:00
"""
from alembic import op
import sqlalchemy as sa
from sqlalchemy import inspect
# revision identifiers, used by Alembic.
revision = "33e347f97c0c"
down_revision = "ddd59cdf0349"
branch_labels = None
depends_on = None
def column_exists(table_name: str, column_name: str) -> bool:
"""检查列是否存在"""
bind = op.get_bind()
inspector = inspect(bind)
columns = [col["name"] for col in inspector.get_columns(table_name)]
return column_name in columns
def table_exists(table_name: str) -> bool:
"""检查表是否存在"""
bind = op.get_bind()
inspector = inspect(bind)
return table_name in inspector.get_table_names()
def column_is_nullable(table_name: str, column_name: str) -> bool:
"""检查列是否允许 NULL"""
bind = op.get_bind()
inspector = inspect(bind)
for col in inspector.get_columns(table_name):
if col["name"] == column_name:
return col["nullable"]
return False
def enum_value_exists(enum_name: str, value: str) -> bool:
"""检查 PostgreSQL ENUM 是否包含指定值"""
bind = op.get_bind()
if bind.dialect.name != "postgresql":
return True # 非 PostgreSQL 跳过检查
result = bind.execute(
sa.text(
"SELECT 1 FROM pg_enum WHERE enumlabel = :value "
"AND enumtypid = (SELECT oid FROM pg_type WHERE typname = :enum_name)"
),
{"value": value, "enum_name": enum_name},
).first()
return result is not None
def upgrade() -> None:
"""应用迁移:升级到新版本"""
bind = op.get_bind()
# ========== Part 1: users 表修改 ==========
# 1) 新增 email_verified
if not column_exists("users", "email_verified"):
op.add_column("users", sa.Column("email_verified", sa.Boolean(), nullable=True))
# 历史数据回填:已有邮箱的用户默认视为已验证
op.execute(sa.text("UPDATE users SET email_verified = true WHERE email IS NOT NULL"))
op.execute(sa.text("UPDATE users SET email_verified = false WHERE email IS NULL"))
# 收紧约束
op.alter_column("users", "email_verified", existing_type=sa.Boolean(), nullable=False)
# 2) email 放宽为可空
if not column_is_nullable("users", "email"):
op.alter_column(
"users",
"email",
existing_type=sa.String(length=255),
nullable=True,
)
# 3) password_hash 放宽为可空
if not column_is_nullable("users", "password_hash"):
op.alter_column(
"users",
"password_hash",
existing_type=sa.String(length=255),
nullable=True,
)
# ========== Part 2: OAuth 相关 ==========
# 4) 扩展 authsource enum
if bind.dialect.name == "postgresql" and not enum_value_exists("authsource", "oauth"):
ctx = op.get_context()
with ctx.autocommit_block():
op.execute("ALTER TYPE authsource ADD VALUE IF NOT EXISTS 'oauth'")
# 5) OAuth provider 配置表
if not table_exists("oauth_providers"):
op.create_table(
"oauth_providers",
sa.Column("provider_type", sa.String(length=50), primary_key=True),
sa.Column("display_name", sa.String(length=100), nullable=False),
sa.Column("client_id", sa.String(length=255), nullable=False),
sa.Column("client_secret_encrypted", sa.Text(), nullable=True),
sa.Column("authorization_url_override", sa.String(length=500), nullable=True),
sa.Column("token_url_override", sa.String(length=500), nullable=True),
sa.Column("userinfo_url_override", sa.String(length=500), nullable=True),
sa.Column("scopes", sa.JSON(), nullable=True),
sa.Column("redirect_uri", sa.String(length=500), nullable=False),
sa.Column("frontend_callback_url", sa.String(length=500), nullable=False),
sa.Column("attribute_mapping", sa.JSON(), nullable=True),
sa.Column("extra_config", sa.JSON(), nullable=True),
sa.Column(
"is_enabled", sa.Boolean(), nullable=False, server_default=sa.text("false")
),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("CURRENT_TIMESTAMP"),
),
sa.Column(
"updated_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("CURRENT_TIMESTAMP"),
),
)
# 6) 用户 OAuth 绑定关系表
if not table_exists("user_oauth_links"):
op.create_table(
"user_oauth_links",
sa.Column("id", sa.String(length=36), primary_key=True),
sa.Column(
"user_id",
sa.String(length=36),
sa.ForeignKey("users.id", ondelete="CASCADE"),
nullable=False,
),
sa.Column(
"provider_type",
sa.String(length=50),
sa.ForeignKey("oauth_providers.provider_type", ondelete="CASCADE"),
nullable=False,
),
sa.Column("provider_user_id", sa.String(length=255), nullable=False),
sa.Column("provider_username", sa.String(length=255), nullable=True),
sa.Column("provider_email", sa.String(length=255), nullable=True),
sa.Column("extra_data", sa.JSON(), nullable=True),
sa.Column(
"linked_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("CURRENT_TIMESTAMP"),
),
sa.Column("last_login_at", sa.DateTime(timezone=True), nullable=True),
sa.UniqueConstraint(
"provider_type", "provider_user_id", name="uq_oauth_provider_user"
),
sa.UniqueConstraint("user_id", "provider_type", name="uq_user_oauth_provider"),
)
op.create_index("ix_user_oauth_links_user_id", "user_oauth_links", ["user_id"])
op.create_index(
"ix_user_oauth_links_provider_type", "user_oauth_links", ["provider_type"]
)
def downgrade() -> None:
"""回滚迁移:降级到旧版本"""
bind = op.get_bind()
# ========== Part 2: OAuth 相关(先删除,因为有外键依赖) ==========
if table_exists("user_oauth_links"):
op.drop_index("ix_user_oauth_links_provider_type", table_name="user_oauth_links")
op.drop_index("ix_user_oauth_links_user_id", table_name="user_oauth_links")
op.drop_table("user_oauth_links")
if table_exists("oauth_providers"):
op.drop_table("oauth_providers")
# 注意:Postgres 不支持从 ENUM 删除值,authsource 不回退
# ========== Part 1: users 表修改 ==========
# 降级前检查:避免把包含 NULL 的列强制改回 NOT NULL
has_null_email = bind.execute(
sa.text("SELECT 1 FROM users WHERE email IS NULL LIMIT 1")
).first()
if has_null_email:
raise RuntimeError("Cannot downgrade: users.email contains NULL values")
has_null_password = bind.execute(
sa.text("SELECT 1 FROM users WHERE password_hash IS NULL LIMIT 1")
).first()
if has_null_password:
raise RuntimeError("Cannot downgrade: users.password_hash contains NULL values")
# 恢复 NOT NULL 约束
if column_is_nullable("users", "email"):
op.alter_column(
"users",
"email",
existing_type=sa.String(length=255),
nullable=False,
)
if column_is_nullable("users", "password_hash"):
op.alter_column(
"users",
"password_hash",
existing_type=sa.String(length=255),
nullable=False,
)
if column_exists("users", "email_verified"):
op.drop_column("users", "email_verified")
@@ -0,0 +1,65 @@
"""add_stats_daily_provider_table
Revision ID: c868729753ad
Revises: 33e347f97c0c
Create Date: 2026-01-19 05:19:49.634662+00:00
"""
from alembic import op
import sqlalchemy as sa
from sqlalchemy import inspect
# revision identifiers, used by Alembic.
revision = 'c868729753ad'
down_revision = '33e347f97c0c'
branch_labels = None
depends_on = None
def table_exists(table_name: str) -> bool:
"""检查表是否存在"""
bind = op.get_bind()
inspector = inspect(bind)
return table_name in inspector.get_table_names()
def index_exists(table_name: str, index_name: str) -> bool:
"""检查索引是否存在"""
bind = op.get_bind()
inspector = inspect(bind)
indexes = [idx['name'] for idx in inspector.get_indexes(table_name)]
return index_name in indexes
def upgrade() -> None:
"""应用迁移:升级到新版本"""
if not table_exists('stats_daily_provider'):
op.create_table(
'stats_daily_provider',
sa.Column('id', sa.String(length=36), nullable=False),
sa.Column('date', sa.DateTime(timezone=True), nullable=False),
sa.Column('provider_name', sa.String(length=100), nullable=False),
sa.Column('total_requests', sa.Integer(), nullable=False),
sa.Column('input_tokens', sa.BigInteger(), nullable=False),
sa.Column('output_tokens', sa.BigInteger(), nullable=False),
sa.Column('cache_creation_tokens', sa.BigInteger(), nullable=False),
sa.Column('cache_read_tokens', sa.BigInteger(), nullable=False),
sa.Column('total_cost', sa.Float(), nullable=False),
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
sa.PrimaryKeyConstraint('id'),
sa.UniqueConstraint('date', 'provider_name', name='uq_stats_daily_provider')
)
op.create_index('idx_stats_daily_provider_date', 'stats_daily_provider', ['date'], unique=False)
op.create_index('idx_stats_daily_provider_date_provider', 'stats_daily_provider', ['date', 'provider_name'], unique=False)
def downgrade() -> None:
"""回滚迁移:降级到旧版本"""
if table_exists('stats_daily_provider'):
if index_exists('stats_daily_provider', 'idx_stats_daily_provider_date_provider'):
op.drop_index('idx_stats_daily_provider_date_provider', table_name='stats_daily_provider')
if index_exists('stats_daily_provider', 'idx_stats_daily_provider_date'):
op.drop_index('idx_stats_daily_provider_date', table_name='stats_daily_provider')
op.drop_table('stats_daily_provider')
@@ -0,0 +1,51 @@
"""add_format_acceptance_config_to_provider_endpoints
Revision ID: 4b4c7b0df1a2
Revises: c868729753ad
Create Date: 2026-01-21 18:45:00+00:00
"""
from alembic import op
import sqlalchemy as sa
from sqlalchemy import inspect
# revision identifiers, used by Alembic.
revision = "4b4c7b0df1a2"
down_revision = "c868729753ad"
branch_labels = None
depends_on = None
def table_exists(table_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
return table_name in inspector.get_table_names()
def column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
columns = [col["name"] for col in inspector.get_columns(table_name)]
return column_name in columns
def upgrade() -> None:
if not table_exists("provider_endpoints"):
return
if column_exists("provider_endpoints", "format_acceptance_config"):
return
op.add_column(
"provider_endpoints",
sa.Column("format_acceptance_config", sa.JSON(), nullable=True),
)
def downgrade() -> None:
if not table_exists("provider_endpoints"):
return
if not column_exists("provider_endpoints", "format_acceptance_config"):
return
op.drop_column("provider_endpoints", "format_acceptance_config")
@@ -0,0 +1,114 @@
"""add_format_conversion_tracking_and_model_filter_patterns_and_provider_timeout
Revision ID: f7c8d9e0a1b2
Revises: 4b4c7b0df1a2
Create Date: 2026-01-27 10:00:00+00:00
Changes:
1. usage 表: 添加 endpoint_api_format 和 has_format_conversion 字段
2. provider_api_keys 表: 添加 model_include_patterns 和 model_exclude_patterns 字段
- 支持通配符规则自动过滤从上游获取的模型列表
- 包含规则和排除规则(支持 * 和 ? 通配符)
3. providers 表: 添加 stream_first_byte_timeout 和 request_timeout 字段
- 允许每个提供商单独配置超时时间
"""
from alembic import op
import sqlalchemy as sa
from sqlalchemy import inspect
# revision identifiers, used by Alembic.
revision = "f7c8d9e0a1b2"
down_revision = "4b4c7b0df1a2"
branch_labels = None
depends_on = None
def table_exists(table_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
return table_name in inspector.get_table_names()
def column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
columns = [col["name"] for col in inspector.get_columns(table_name)]
return column_name in columns
def upgrade() -> None:
# === usage 表: 格式转换追踪 ===
if table_exists("usage"):
# 添加 endpoint_api_format 字段(端点原生 API 格式)
if not column_exists("usage", "endpoint_api_format"):
op.add_column(
"usage",
sa.Column("endpoint_api_format", sa.String(50), nullable=True),
)
# 添加 has_format_conversion 字段(是否发生了格式转换)
if not column_exists("usage", "has_format_conversion"):
op.add_column(
"usage",
sa.Column("has_format_conversion", sa.Boolean(), nullable=True, server_default="false"),
)
# === provider_api_keys 表: 模型过滤规则 ===
if table_exists("provider_api_keys"):
# 添加 model_include_patterns 字段(包含规则,支持 * 和 ? 通配符)
if not column_exists("provider_api_keys", "model_include_patterns"):
op.add_column(
"provider_api_keys",
sa.Column("model_include_patterns", sa.JSON(), nullable=True),
)
# 添加 model_exclude_patterns 字段(排除规则,支持 * 和 ? 通配符)
if not column_exists("provider_api_keys", "model_exclude_patterns"):
op.add_column(
"provider_api_keys",
sa.Column("model_exclude_patterns", sa.JSON(), nullable=True),
)
# === providers 表: 超时配置 ===
if table_exists("providers"):
# 添加 stream_first_byte_timeout 字段(流式请求首字节超时)
if not column_exists("providers", "stream_first_byte_timeout"):
op.add_column(
"providers",
sa.Column("stream_first_byte_timeout", sa.Float(), nullable=True),
)
# 添加 request_timeout 字段(非流式请求整体超时)
if not column_exists("providers", "request_timeout"):
op.add_column(
"providers",
sa.Column("request_timeout", sa.Float(), nullable=True),
)
def downgrade() -> None:
# === providers 表: 移除超时配置 ===
if table_exists("providers"):
if column_exists("providers", "request_timeout"):
op.drop_column("providers", "request_timeout")
if column_exists("providers", "stream_first_byte_timeout"):
op.drop_column("providers", "stream_first_byte_timeout")
# === provider_api_keys 表: 移除模型过滤规则 ===
if table_exists("provider_api_keys"):
if column_exists("provider_api_keys", "model_exclude_patterns"):
op.drop_column("provider_api_keys", "model_exclude_patterns")
if column_exists("provider_api_keys", "model_include_patterns"):
op.drop_column("provider_api_keys", "model_include_patterns")
# === usage 表: 移除格式转换追踪 ===
if table_exists("usage"):
if column_exists("usage", "has_format_conversion"):
op.drop_column("usage", "has_format_conversion")
if column_exists("usage", "endpoint_api_format"):
op.drop_column("usage", "endpoint_api_format")
@@ -0,0 +1,58 @@
"""add_keep_priority_on_conversion_to_providers
Revision ID: 364680d1bc99
Revises: f7c8d9e0a1b2
Create Date: 2026-01-28 12:00:00+00:00
Changes:
1. providers 表: 添加 keep_priority_on_conversion 字段
- 格式转换时是否保持提供商原优先级
- 默认 False:需要格式转换时,候选会被降级到不需要转换的候选之后
- 设为 True:即使需要格式转换,也保持原优先级排名
"""
from alembic import op
import sqlalchemy as sa
from sqlalchemy import inspect
# revision identifiers, used by Alembic.
revision = "364680d1bc99"
down_revision = "f7c8d9e0a1b2"
branch_labels = None
depends_on = None
def table_exists(table_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
return table_name in inspector.get_table_names()
def column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
columns = [col["name"] for col in inspector.get_columns(table_name)]
return column_name in columns
def upgrade() -> None:
# === providers 表: 添加格式转换优先级保持配置 ===
if table_exists("providers"):
if not column_exists("providers", "keep_priority_on_conversion"):
op.add_column(
"providers",
sa.Column(
"keep_priority_on_conversion",
sa.Boolean(),
nullable=False,
server_default="false",
),
)
def downgrade() -> None:
# === providers 表: 移除格式转换优先级保持配置 ===
if table_exists("providers"):
if column_exists("providers", "keep_priority_on_conversion"):
op.drop_column("providers", "keep_priority_on_conversion")
@@ -0,0 +1,51 @@
"""Add auth_type and auth_config fields to provider_api_keys table
Revision ID: 7f6f8065f517
Revises: 364680d1bc99
Create Date: 2026-01-30 10:00:00.000000
"""
from typing import Sequence, Union
import sqlalchemy as sa
from alembic import op
from sqlalchemy import inspect
# revision identifiers, used by Alembic.
revision: str = "7f6f8065f517"
down_revision: Union[str, None] = "364680d1bc99"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def column_exists(table_name: str, column_name: str) -> bool:
"""检查列是否已存在"""
bind = op.get_bind()
inspector = inspect(bind)
columns = [col["name"] for col in inspector.get_columns(table_name)]
return column_name in columns
def upgrade() -> None:
# 添加 auth_type 字段,默认值为 "api_key"
if not column_exists("provider_api_keys", "auth_type"):
op.add_column(
"provider_api_keys",
sa.Column("auth_type", sa.String(20), nullable=False, server_default="api_key"),
)
# 添加 auth_config 字段(Text,存储加密后的认证配置)
if not column_exists("provider_api_keys", "auth_config"):
op.add_column(
"provider_api_keys",
sa.Column("auth_config", sa.Text, nullable=True),
)
def downgrade() -> None:
if column_exists("provider_api_keys", "auth_config"):
op.drop_column("provider_api_keys", "auth_config")
if column_exists("provider_api_keys", "auth_type"):
op.drop_column("provider_api_keys", "auth_type")
@@ -0,0 +1,112 @@
"""Add video_tasks table
Revision ID: b6f1a2c5d8e9
Revises: 7f6f8065f517
Create Date: 2026-01-30 18:00:00.000000
"""
from typing import Sequence, Union
import sqlalchemy as sa
from sqlalchemy import inspect
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "b6f1a2c5d8e9"
down_revision: Union[str, None] = "7f6f8065f517"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def table_exists(table_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
return table_name in inspector.get_table_names()
def upgrade() -> None:
if table_exists("video_tasks"):
return
op.create_table(
"video_tasks",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("external_task_id", sa.String(200), nullable=True, index=False),
sa.Column("user_id", sa.String(36), sa.ForeignKey("users.id"), nullable=False),
sa.Column("api_key_id", sa.String(36), sa.ForeignKey("api_keys.id"), nullable=True),
sa.Column("provider_id", sa.String(36), sa.ForeignKey("providers.id"), nullable=True),
sa.Column(
"endpoint_id", sa.String(36), sa.ForeignKey("provider_endpoints.id"), nullable=True
),
sa.Column("key_id", sa.String(36), sa.ForeignKey("provider_api_keys.id"), nullable=True),
sa.Column("client_api_format", sa.String(50), nullable=False),
sa.Column("provider_api_format", sa.String(50), nullable=False),
sa.Column("format_converted", sa.Boolean(), server_default=sa.false()),
sa.Column("model", sa.String(100), nullable=False),
sa.Column("prompt", sa.Text(), nullable=False),
sa.Column("original_request_body", sa.JSON(), nullable=True),
sa.Column("converted_request_body", sa.JSON(), nullable=True),
sa.Column("duration_seconds", sa.Integer(), server_default=sa.text("4")),
sa.Column("resolution", sa.String(20), server_default=sa.text("'720p'")),
sa.Column("aspect_ratio", sa.String(10), server_default=sa.text("'16:9'")),
sa.Column("size", sa.String(20), nullable=True),
sa.Column("status", sa.String(20), server_default=sa.text("'pending'")),
sa.Column("progress_percent", sa.Integer(), server_default=sa.text("0")),
sa.Column("progress_message", sa.String(500), nullable=True),
sa.Column("video_url", sa.String(2000), nullable=True),
sa.Column("video_urls", sa.JSON(), nullable=True),
sa.Column("thumbnail_url", sa.String(2000), nullable=True),
sa.Column("video_size_bytes", sa.BigInteger(), nullable=True),
sa.Column("video_expires_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("stored_video_path", sa.String(500), nullable=True),
sa.Column("storage_provider", sa.String(50), nullable=True),
sa.Column("error_code", sa.String(50), nullable=True),
sa.Column("error_message", sa.Text(), nullable=True),
sa.Column("retry_count", sa.Integer(), server_default=sa.text("0")),
sa.Column("max_retries", sa.Integer(), server_default=sa.text("3")),
sa.Column("poll_interval_seconds", sa.Integer(), server_default=sa.text("10")),
sa.Column("next_poll_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("poll_count", sa.Integer(), server_default=sa.text("0")),
sa.Column("max_poll_count", sa.Integer(), server_default=sa.text("360")),
sa.Column(
"remixed_from_task_id",
sa.String(36),
sa.ForeignKey("video_tasks.id", ondelete="SET NULL"),
nullable=True,
),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
server_default=sa.text("CURRENT_TIMESTAMP"),
),
sa.Column("submitted_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("completed_at", sa.DateTime(timezone=True), nullable=True),
sa.Column(
"updated_at",
sa.DateTime(timezone=True),
server_default=sa.text("CURRENT_TIMESTAMP"),
),
)
op.create_index("idx_video_tasks_user_status", "video_tasks", ["user_id", "status"])
op.create_index("idx_video_tasks_next_poll", "video_tasks", ["next_poll_at"])
op.create_index("idx_video_tasks_external_id", "video_tasks", ["external_task_id"])
# 唯一约束:同一用户不能有重复的 external_task_id
op.create_unique_constraint(
"uq_video_tasks_user_external_id",
"video_tasks",
["user_id", "external_task_id"],
)
def downgrade() -> None:
if not table_exists("video_tasks"):
return
op.drop_constraint("uq_video_tasks_user_external_id", "video_tasks", type_="unique")
op.drop_index("idx_video_tasks_external_id", table_name="video_tasks")
op.drop_index("idx_video_tasks_next_poll", table_name="video_tasks")
op.drop_index("idx_video_tasks_user_status", table_name="video_tasks")
op.drop_table("video_tasks")
@@ -0,0 +1,180 @@
"""Add billing system tables and video_tasks.request_metadata
Revision ID: c8d2e4f6a1b3
Revises: b6f1a2c5d8e9
Create Date: 2026-01-31 12:00:00.000000
"""
from __future__ import annotations
from typing import Sequence, Union
import sqlalchemy as sa
from sqlalchemy import inspect
from sqlalchemy.dialects.postgresql import JSONB
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "c8d2e4f6a1b3"
down_revision: Union[str, None] = "b6f1a2c5d8e9"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def table_exists(table_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
return table_name in inspector.get_table_names()
def column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
columns = [col["name"] for col in inspector.get_columns(table_name)]
return column_name in columns
def index_exists(table_name: str, index_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
try:
indexes = inspector.get_indexes(table_name)
except Exception:
return False
return any(idx.get("name") == index_name for idx in indexes)
def upgrade() -> None:
# ==================== video_tasks.request_metadata ====================
if not column_exists("video_tasks", "request_metadata"):
op.add_column(
"video_tasks",
sa.Column("request_metadata", sa.JSON(), nullable=True),
)
# ==================== billing_rules ====================
if not table_exists("billing_rules"):
op.create_table(
"billing_rules",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column(
"global_model_id",
sa.String(36),
sa.ForeignKey("global_models.id", ondelete="CASCADE"),
nullable=True,
),
sa.Column(
"model_id",
sa.String(36),
sa.ForeignKey("models.id", ondelete="CASCADE"),
nullable=True,
),
sa.Column("name", sa.String(100), nullable=False),
sa.Column("task_type", sa.String(20), nullable=False, server_default="chat"),
sa.Column("expression", sa.Text(), nullable=False),
sa.Column("variables", JSONB, nullable=False, server_default=sa.text("'{}'::jsonb")),
sa.Column(
"dimension_mappings", JSONB, nullable=False, server_default=sa.text("'{}'::jsonb")
),
sa.Column("is_enabled", sa.Boolean(), nullable=False, server_default=sa.text("true")),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("now()"),
),
sa.Column(
"updated_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("now()"),
),
sa.CheckConstraint(
"(global_model_id IS NOT NULL AND model_id IS NULL) OR "
"(global_model_id IS NULL AND model_id IS NOT NULL)",
name="chk_billing_rules_model_ref",
),
)
# Partial unique indexes for enabled rules
if table_exists("billing_rules"):
if not index_exists("billing_rules", "uq_billing_rules_global_model_task"):
op.create_index(
"uq_billing_rules_global_model_task",
"billing_rules",
["global_model_id", "task_type"],
unique=True,
postgresql_where=sa.text("is_enabled = TRUE AND global_model_id IS NOT NULL"),
)
if not index_exists("billing_rules", "uq_billing_rules_model_task"):
op.create_index(
"uq_billing_rules_model_task",
"billing_rules",
["model_id", "task_type"],
unique=True,
postgresql_where=sa.text("is_enabled = TRUE AND model_id IS NOT NULL"),
)
# ==================== dimension_collectors ====================
if not table_exists("dimension_collectors"):
op.create_table(
"dimension_collectors",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("api_format", sa.String(50), nullable=False),
sa.Column("task_type", sa.String(20), nullable=False),
sa.Column("dimension_name", sa.String(100), nullable=False),
sa.Column("source_type", sa.String(20), nullable=False),
sa.Column("source_path", sa.String(200), nullable=True),
sa.Column("value_type", sa.String(20), nullable=False, server_default="float"),
sa.Column("transform_expression", sa.Text(), nullable=True),
sa.Column("default_value", sa.String(100), nullable=True),
sa.Column("priority", sa.Integer(), nullable=False, server_default="0"),
sa.Column("is_enabled", sa.Boolean(), nullable=False, server_default=sa.text("true")),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("now()"),
),
sa.Column(
"updated_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("now()"),
),
sa.CheckConstraint(
"(source_type = 'computed' AND source_path IS NULL AND transform_expression IS NOT NULL) OR "
"(source_type != 'computed' AND source_path IS NOT NULL)",
name="chk_dimension_collectors_source_config",
),
)
if table_exists("dimension_collectors"):
if not index_exists("dimension_collectors", "uq_dimension_collectors_enabled"):
op.create_index(
"uq_dimension_collectors_enabled",
"dimension_collectors",
["api_format", "task_type", "dimension_name", "priority"],
unique=True,
postgresql_where=sa.text("is_enabled = TRUE"),
)
def downgrade() -> None:
# Drop in reverse order
if table_exists("dimension_collectors"):
if index_exists("dimension_collectors", "uq_dimension_collectors_enabled"):
op.drop_index("uq_dimension_collectors_enabled", table_name="dimension_collectors")
op.drop_table("dimension_collectors")
if table_exists("billing_rules"):
if index_exists("billing_rules", "uq_billing_rules_model_task"):
op.drop_index("uq_billing_rules_model_task", table_name="billing_rules")
if index_exists("billing_rules", "uq_billing_rules_global_model_task"):
op.drop_index("uq_billing_rules_global_model_task", table_name="billing_rules")
op.drop_table("billing_rules")
if column_exists("video_tasks", "request_metadata"):
op.drop_column("video_tasks", "request_metadata")
@@ -0,0 +1,462 @@
"""Add api_family/endpoint_kind and migrate api_format to endpoint signature keys
Revision ID: cf40e6a5c5b1
Revises: c8d2e4f6a1b3
Create Date: 2026-01-31 15:30:00.000000
"""
from __future__ import annotations
import json
from datetime import datetime, timezone
from typing import Sequence, Union
from uuid import uuid4
import sqlalchemy as sa
from sqlalchemy import inspect, text
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "cf40e6a5c5b1"
down_revision: Union[str, None] = "c8d2e4f6a1b3"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def _json_loads(val):
if val is None:
return None
if isinstance(val, (dict, list)):
return val
if isinstance(val, str):
try:
return json.loads(val)
except Exception:
return None
return None
def _json_dumps(val):
"""将 dict/list 转为 JSON 字符串,None 保持 None"""
if val is None:
return None
if isinstance(val, str):
return val
return json.dumps(val)
def _normalize_signature(value: str | None) -> str | None:
"""
Normalize legacy api_format / signature-ish strings to canonical signature key.
- canonical: `<family>:<kind>` (lowercase)
- legacy examples: "OPENAI", "OPENAI_CLI", "GEMINI_VIDEO"
"""
if value is None:
return None
raw = str(value).strip()
if not raw:
return None
if ":" in raw:
fam, kind = raw.split(":", 1)
fam = fam.strip().lower()
kind = kind.strip().lower()
if not fam or not kind:
return None
return f"{fam}:{kind}"
upper = raw.upper()
if upper.startswith("CLAUDE"):
fam = "claude"
elif upper.startswith("OPENAI"):
fam = "openai"
elif upper.startswith("GEMINI"):
fam = "gemini"
else:
return None
kind = "chat"
if upper.endswith("_CLI"):
kind = "cli"
elif upper.endswith("_VIDEO"):
kind = "video"
return f"{fam}:{kind}"
def _normalize_signature_list(values) -> list[str] | None:
if values is None:
return None
if isinstance(values, str):
values = _json_loads(values)
if not isinstance(values, list):
return None
out: list[str] = []
seen: set[str] = set()
for v in values:
sig = _normalize_signature(str(v) if v is not None else None)
if not sig:
continue
if sig in seen:
continue
seen.add(sig)
out.append(sig)
return out
def _normalize_signature_dict(values) -> dict | None:
if values is None:
return None
if isinstance(values, str):
values = _json_loads(values)
if not isinstance(values, dict):
return None
out: dict = {}
for k, v in values.items():
sig = _normalize_signature(str(k) if k is not None else None)
if not sig:
continue
out[sig] = v
return out
def _add_video_variants(formats: list[str]) -> list[str]:
"""
保持原有格式,不自动补齐 video 变体。
"""
return formats
def table_exists(table_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
return table_name in inspector.get_table_names()
def column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
columns = [col["name"] for col in inspector.get_columns(table_name)]
return column_name in columns
def index_exists(table_name: str, index_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
try:
indexes = inspector.get_indexes(table_name)
except Exception:
return False
return any(idx.get("name") == index_name for idx in indexes)
def _migrate_format_acceptance_config(cfg) -> dict | None:
cfg_obj = _json_loads(cfg)
if not isinstance(cfg_obj, dict):
return cfg_obj if cfg_obj is None else None
for key in ("accept_formats", "reject_formats"):
raw = cfg_obj.get(key)
if not isinstance(raw, list):
continue
normalized = _normalize_signature_list(raw) or []
cfg_obj[key] = normalized
return cfg_obj
def migrate_provider_endpoints(connection) -> None:
"""
- 将 provider_endpoints.api_format 统一迁移为 signature key(小写)
- 填充/校准 api_family / endpoint_kind
- 迁移 format_acceptance_config 中的 accept/reject formats
"""
rows = connection.execute(text("""
SELECT
id,
api_format,
api_family,
endpoint_kind,
format_acceptance_config
FROM provider_endpoints
""")).fetchall()
for row in rows:
sig = _normalize_signature(row.api_format)
if not sig:
continue
fam, kind = sig.split(":", 1)
cfg = _migrate_format_acceptance_config(row.format_acceptance_config)
connection.execute(
text("""
UPDATE provider_endpoints
SET
api_format = :api_format,
api_family = :api_family,
endpoint_kind = :endpoint_kind,
format_acceptance_config = CAST(:format_acceptance_config AS json)
WHERE id = :id
"""),
{
"id": row.id,
"api_format": sig,
"api_family": fam,
"endpoint_kind": kind,
"format_acceptance_config": _json_dumps(cfg),
},
)
def create_video_endpoints(connection) -> None:
"""
不再自动创建 video endpoint,保持原有配置。
"""
pass
def migrate_provider_api_keys(connection) -> None:
"""
迁移 provider_api_keys:
- api_formats -> signature keys(并补齐 video 变体)
- dict 字段 key -> signature keys(rate_multipliers/global_priority/health/circuit_breaker)
- rate_multipliers/global_priority_by_format 复制 chat -> video(如 openai:chat -> openai:video)
"""
rows = connection.execute(text("""
SELECT
id,
api_formats,
rate_multipliers,
global_priority_by_format,
health_by_format,
circuit_breaker_by_format
FROM provider_api_keys
""")).fetchall()
for row in rows:
api_formats = _normalize_signature_list(row.api_formats)
if api_formats is not None:
api_formats = _add_video_variants(api_formats)
rate_multipliers = _normalize_signature_dict(row.rate_multipliers)
global_priority_by_format = _normalize_signature_dict(row.global_priority_by_format)
health_by_format = _normalize_signature_dict(row.health_by_format)
circuit_breaker_by_format = _normalize_signature_dict(row.circuit_breaker_by_format)
connection.execute(
text("""
UPDATE provider_api_keys
SET
api_formats = CAST(:api_formats AS json),
rate_multipliers = CAST(:rate_multipliers AS json),
global_priority_by_format = CAST(:global_priority_by_format AS json),
health_by_format = CAST(:health_by_format AS json),
circuit_breaker_by_format = CAST(:circuit_breaker_by_format AS json)
WHERE id = :id
"""),
{
"id": row.id,
"api_formats": _json_dumps(api_formats),
"rate_multipliers": _json_dumps(rate_multipliers),
"global_priority_by_format": _json_dumps(global_priority_by_format),
"health_by_format": _json_dumps(health_by_format),
"circuit_breaker_by_format": _json_dumps(circuit_breaker_by_format),
},
)
def migrate_allowed_api_formats(connection, *, table_name: str) -> None:
"""迁移 users/api_keys.allowed_api_formats 为 signature keys(并补齐 video 变体)。"""
if not table_exists(table_name):
return
rows = connection.execute(text(f"""
SELECT id, allowed_api_formats
FROM {table_name}
""")).fetchall()
for row in rows:
allowed = _normalize_signature_list(row.allowed_api_formats)
if allowed is None:
continue
allowed = _add_video_variants(allowed)
connection.execute(
text(f"""
UPDATE {table_name}
SET allowed_api_formats = CAST(:allowed_api_formats AS json)
WHERE id = :id
"""),
{"id": row.id, "allowed_api_formats": _json_dumps(allowed)},
)
def migrate_video_tasks(connection) -> None:
"""
video_tasks.*_api_format 迁移为 signature keys。
注意:video_tasks 表天然是 video 任务,因此将 openai/gemini 的 kind 强制归一为 video,
以兼容历史上复用 chat 格式存储的旧记录。
"""
if not table_exists("video_tasks"):
return
rows = connection.execute(text("""
SELECT id, client_api_format, provider_api_format
FROM video_tasks
""")).fetchall()
for row in rows:
client_sig = _normalize_signature(row.client_api_format) or ""
provider_sig = _normalize_signature(row.provider_api_format) or ""
def _force_video(sig: str) -> str:
if not sig or ":" not in sig:
return sig
fam, _kind = sig.split(":", 1)
fam = fam.strip().lower()
if fam in ("openai", "gemini"):
return f"{fam}:video"
return sig
client_sig = _force_video(client_sig)
provider_sig = _force_video(provider_sig)
if not client_sig or not provider_sig:
continue
connection.execute(
text("""
UPDATE video_tasks
SET client_api_format = :client_api_format,
provider_api_format = :provider_api_format
WHERE id = :id
"""),
{
"id": row.id,
"client_api_format": client_sig,
"provider_api_format": provider_sig,
},
)
def migrate_model_provider_mappings(connection) -> None:
"""迁移 models.provider_model_mappings[*].api_formats 为 signature keys。"""
if not table_exists("models"):
return
rows = connection.execute(text("""
SELECT id, provider_model_mappings
FROM models
WHERE provider_model_mappings IS NOT NULL
""")).fetchall()
for row in rows:
mappings = _json_loads(row.provider_model_mappings)
if not isinstance(mappings, list):
continue
changed = False
new_mappings: list = []
for item in mappings:
if not isinstance(item, dict):
new_mappings.append(item)
continue
raw_formats = item.get("api_formats")
if isinstance(raw_formats, list):
normalized = _normalize_signature_list(raw_formats) or []
# 内容比较(而非引用比较),避免已迁移数据被无意义地重复 UPDATE
if set(normalized) != set(raw_formats):
changed = True
item = dict(item)
item["api_formats"] = normalized
new_mappings.append(item)
if not changed:
continue
connection.execute(
text("""
UPDATE models
SET provider_model_mappings = CAST(:provider_model_mappings AS json)
WHERE id = :id
"""),
{"id": row.id, "provider_model_mappings": _json_dumps(new_mappings)},
)
def migrate_dimension_collectors(connection) -> None:
"""迁移 dimension_collectors.api_format 为 signature keys(如果存在历史数据)。"""
if not table_exists("dimension_collectors"):
return
rows = connection.execute(text("""
SELECT id, api_format
FROM dimension_collectors
WHERE api_format IS NOT NULL
""")).fetchall()
for row in rows:
sig = _normalize_signature(row.api_format)
if not sig:
continue
connection.execute(
text("""
UPDATE dimension_collectors
SET api_format = :api_format
WHERE id = :id
"""),
{"id": row.id, "api_format": sig},
)
def upgrade() -> None:
if not table_exists("provider_endpoints"):
return
# ==================== provider_endpoints.api_family / endpoint_kind ====================
if not column_exists("provider_endpoints", "api_family"):
op.add_column("provider_endpoints", sa.Column("api_family", sa.String(50), nullable=True))
if not column_exists("provider_endpoints", "endpoint_kind"):
op.add_column(
"provider_endpoints", sa.Column("endpoint_kind", sa.String(50), nullable=True)
)
# ==================== idx_provider_family_kind ====================
if not index_exists("provider_endpoints", "idx_provider_family_kind"):
op.create_index(
"idx_provider_family_kind",
"provider_endpoints",
["provider_id", "api_family", "endpoint_kind"],
)
# ==================== data migrations (idempotent) ====================
conn = op.get_bind()
migrate_provider_endpoints(conn)
create_video_endpoints(conn)
if table_exists("provider_api_keys"):
migrate_provider_api_keys(conn)
migrate_allowed_api_formats(conn, table_name="users")
migrate_allowed_api_formats(conn, table_name="api_keys")
migrate_video_tasks(conn)
migrate_model_provider_mappings(conn)
migrate_dimension_collectors(conn)
def downgrade() -> None:
# Drop index/columns only; data changes are intentionally kept (safe rollback strategy).
if table_exists("provider_endpoints"):
if index_exists("provider_endpoints", "idx_provider_family_kind"):
op.drop_index("idx_provider_family_kind", table_name="provider_endpoints")
if column_exists("provider_endpoints", "endpoint_kind"):
op.drop_column("provider_endpoints", "endpoint_kind")
if column_exists("provider_endpoints", "api_family"):
op.drop_column("provider_endpoints", "api_family")
@@ -0,0 +1,329 @@
"""Add usage billing, video_tasks fields, gemini_file_mappings, provider format conversion, and indexes
Revision ID: a2f1b3c4d5e6
Revises: cf40e6a5c5b1
Create Date: 2026-02-01 12:00:00+00:00
Changes:
1. usage 表:
- 添加 billing_status (pending/settled/void),用于表示结算状态
- 添加 finalized_at,用于记录结算完成时间
- 添加 (provider_name, created_at) 和 (model, created_at) 索引
2. video_tasks 表:
- 添加 request_id(全局唯一),用于与 Usage/RequestCandidate 建立稳定关联
- 添加 short_id (Gemini-style short ID)
3. gemini_file_mappings 表:
- 创建新表用于文件映射
- 添加 source_hash 字段用于关联相同源文件
4. providers 表:
- 添加 enable_format_conversion 开关字段
5. request_candidates 表:
- 添加 created_at 索引
"""
from __future__ import annotations
import secrets
import string
from typing import Sequence, Union
import sqlalchemy as sa
from sqlalchemy import inspect, text
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "a2f1b3c4d5e6"
down_revision: Union[str, None] = "cf40e6a5c5b1"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def table_exists(table_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
inspector.clear_cache()
return table_name in inspector.get_table_names()
def column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
# Clear cached schema info to get fresh data
inspector.clear_cache()
columns = [col["name"] for col in inspector.get_columns(table_name)]
return column_name in columns
def index_exists(table_name: str, index_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
inspector.clear_cache()
indexes = inspector.get_indexes(table_name)
return any(idx.get("name") == index_name for idx in indexes)
def unique_constraint_exists(table_name: str, constraint_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
inspector.clear_cache()
constraints = inspector.get_unique_constraints(table_name)
return any(c.get("name") == constraint_name for c in constraints)
def generate_short_id(length: int = 12) -> str:
"""Generate a Gemini-style short ID (lowercase letters + digits)"""
alphabet = string.ascii_lowercase + string.digits
return "".join(secrets.choice(alphabet) for _ in range(length))
def upgrade() -> None:
bind = op.get_bind()
dialect = bind.dialect.name
# =========================================================================
# 1. usage 表: billing_status + finalized_at + 索引
# =========================================================================
if table_exists("usage"):
if not column_exists("usage", "billing_status"):
op.add_column(
"usage",
sa.Column(
"billing_status",
sa.String(20),
nullable=False,
server_default="settled",
),
)
if not column_exists("usage", "finalized_at"):
op.add_column(
"usage",
sa.Column("finalized_at", sa.DateTime(timezone=True), nullable=True),
)
if not index_exists("usage", "idx_usage_billing_status"):
op.create_index("idx_usage_billing_status", "usage", ["billing_status"])
# (provider_name, created_at) — provider list / dashboard queries
if (
column_exists("usage", "provider_name")
and column_exists("usage", "created_at")
and not index_exists("usage", "idx_usage_provider_created")
):
op.create_index("idx_usage_provider_created", "usage", ["provider_name", "created_at"])
# (model, created_at) — model analytics / recent requests queries
if (
column_exists("usage", "model")
and column_exists("usage", "created_at")
and not index_exists("usage", "idx_usage_model_created")
):
op.create_index("idx_usage_model_created", "usage", ["model", "created_at"])
# =========================================================================
# 2. video_tasks 表: request_id + short_id
# =========================================================================
if table_exists("video_tasks"):
# --- request_id ---
if not column_exists("video_tasks", "request_id"):
op.add_column(
"video_tasks",
sa.Column("request_id", sa.String(100), nullable=True),
)
# 回填 request_id
if dialect == "postgresql":
op.execute("""
UPDATE video_tasks
SET request_id = COALESCE(request_metadata->>'request_id', id)
WHERE request_id IS NULL
""")
elif dialect == "sqlite":
op.execute("""
UPDATE video_tasks
SET request_id = COALESCE(json_extract(request_metadata, '$.request_id'), id)
WHERE request_id IS NULL
""")
else:
op.execute("""
UPDATE video_tasks
SET request_id = id
WHERE request_id IS NULL
""")
if dialect == "postgresql":
op.alter_column("video_tasks", "request_id", nullable=False)
if not index_exists("video_tasks", "idx_video_tasks_request_id"):
op.create_index("idx_video_tasks_request_id", "video_tasks", ["request_id"])
if not unique_constraint_exists("video_tasks", "uq_video_tasks_request_id"):
op.create_unique_constraint(
"uq_video_tasks_request_id",
"video_tasks",
["request_id"],
)
# --- short_id ---
if not column_exists("video_tasks", "short_id"):
op.add_column(
"video_tasks",
sa.Column("short_id", sa.String(16), nullable=True),
)
# Populate existing rows with unique short_ids
result = bind.execute(text("SELECT id FROM video_tasks WHERE short_id IS NULL"))
for row in result:
short_id = generate_short_id()
bind.execute(
text("UPDATE video_tasks SET short_id = :short_id WHERE id = :id"),
{"short_id": short_id, "id": row[0]},
)
op.alter_column("video_tasks", "short_id", nullable=False)
op.create_index("ix_video_tasks_short_id", "video_tasks", ["short_id"], unique=True)
# =========================================================================
# 3. gemini_file_mappings 表
# =========================================================================
if not table_exists("gemini_file_mappings"):
op.create_table(
"gemini_file_mappings",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("file_name", sa.String(255), nullable=False, unique=True),
sa.Column(
"key_id",
sa.String(36),
sa.ForeignKey("provider_api_keys.id", ondelete="CASCADE"),
nullable=False,
),
sa.Column(
"user_id",
sa.String(36),
sa.ForeignKey("users.id", ondelete="CASCADE"),
nullable=True,
),
sa.Column("display_name", sa.String(255), nullable=True),
sa.Column("mime_type", sa.String(100), nullable=True),
sa.Column("source_hash", sa.String(64), nullable=True),
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
sa.Column("expires_at", sa.DateTime(timezone=True), nullable=False),
)
op.create_index("ix_gemini_file_mappings_id", "gemini_file_mappings", ["id"])
op.create_index(
"ix_gemini_file_mappings_file_name", "gemini_file_mappings", ["file_name"], unique=True
)
op.create_index("ix_gemini_file_mappings_key_id", "gemini_file_mappings", ["key_id"])
op.create_index("ix_gemini_file_mappings_user_id", "gemini_file_mappings", ["user_id"])
op.create_index("idx_gemini_file_mappings_expires", "gemini_file_mappings", ["expires_at"])
op.create_index(
"idx_gemini_file_mappings_source_hash", "gemini_file_mappings", ["source_hash"]
)
else:
# 表已存在,只添加 source_hash
if not column_exists("gemini_file_mappings", "source_hash"):
op.add_column(
"gemini_file_mappings",
sa.Column("source_hash", sa.String(64), nullable=True),
)
op.create_index(
"idx_gemini_file_mappings_source_hash",
"gemini_file_mappings",
["source_hash"],
)
# =========================================================================
# 4. providers 表: enable_format_conversion
# =========================================================================
if table_exists("providers") and not column_exists("providers", "enable_format_conversion"):
op.add_column(
"providers",
sa.Column(
"enable_format_conversion",
sa.Boolean(),
nullable=False,
server_default=sa.text("false"),
),
)
# =========================================================================
# 5. request_candidates 表: created_at 索引
# =========================================================================
if table_exists("request_candidates"):
if not index_exists("request_candidates", "idx_request_candidates_created_at"):
op.create_index(
"idx_request_candidates_created_at",
"request_candidates",
["created_at"],
unique=False,
)
def downgrade() -> None:
bind = op.get_bind()
dialect = bind.dialect.name
# =========================================================================
# 5. request_candidates 表回滚
# =========================================================================
if table_exists("request_candidates"):
if index_exists("request_candidates", "idx_request_candidates_created_at"):
op.drop_index("idx_request_candidates_created_at", table_name="request_candidates")
# =========================================================================
# 4. providers 表回滚
# =========================================================================
if table_exists("providers") and column_exists("providers", "enable_format_conversion"):
op.drop_column("providers", "enable_format_conversion")
# =========================================================================
# 3. gemini_file_mappings 表回滚
# =========================================================================
if table_exists("gemini_file_mappings"):
op.drop_index("idx_gemini_file_mappings_source_hash", table_name="gemini_file_mappings")
op.drop_index("idx_gemini_file_mappings_expires", table_name="gemini_file_mappings")
op.drop_index("ix_gemini_file_mappings_user_id", table_name="gemini_file_mappings")
op.drop_index("ix_gemini_file_mappings_key_id", table_name="gemini_file_mappings")
op.drop_index("ix_gemini_file_mappings_file_name", table_name="gemini_file_mappings")
op.drop_index("ix_gemini_file_mappings_id", table_name="gemini_file_mappings")
op.drop_table("gemini_file_mappings")
# =========================================================================
# 2. video_tasks 表回滚
# =========================================================================
if table_exists("video_tasks"):
# short_id
if column_exists("video_tasks", "short_id"):
if index_exists("video_tasks", "ix_video_tasks_short_id"):
op.drop_index("ix_video_tasks_short_id", table_name="video_tasks")
op.drop_column("video_tasks", "short_id")
# request_id
if column_exists("video_tasks", "request_id"):
if dialect == "postgresql":
if unique_constraint_exists("video_tasks", "uq_video_tasks_request_id"):
op.drop_constraint("uq_video_tasks_request_id", "video_tasks", type_="unique")
if index_exists("video_tasks", "idx_video_tasks_request_id"):
op.drop_index("idx_video_tasks_request_id", table_name="video_tasks")
op.drop_column("video_tasks", "request_id")
# =========================================================================
# 1. usage 表回滚
# =========================================================================
if table_exists("usage"):
if index_exists("usage", "idx_usage_model_created"):
op.drop_index("idx_usage_model_created", table_name="usage")
if index_exists("usage", "idx_usage_provider_created"):
op.drop_index("idx_usage_provider_created", table_name="usage")
if index_exists("usage", "idx_usage_billing_status"):
op.drop_index("idx_usage_billing_status", table_name="usage")
if column_exists("usage", "finalized_at"):
op.drop_column("usage", "finalized_at")
if column_exists("usage", "billing_status"):
op.drop_column("usage", "billing_status")
@@ -0,0 +1,60 @@
"""Add video_duration_seconds to video_tasks and body_rules to provider_endpoints
Revision ID: b3c4d5e6f7a8
Revises: a2f1b3c4d5e6
Create Date: 2026-02-03 15:00:00.000000
"""
from __future__ import annotations
from typing import Sequence, Union
import sqlalchemy as sa
from sqlalchemy import inspect
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "b3c4d5e6f7a8"
down_revision: Union[str, None] = "a2f1b3c4d5e6"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def _column_exists(table_name: str, column_name: str) -> bool:
"""Check if a column exists in a table."""
bind = op.get_bind()
inspector = inspect(bind)
columns = [col["name"] for col in inspector.get_columns(table_name)]
return column_name in columns
def upgrade() -> None:
# 1. Add video_duration_seconds to video_tasks
if not _column_exists("video_tasks", "video_duration_seconds"):
op.add_column(
"video_tasks",
sa.Column("video_duration_seconds", sa.Float(), nullable=True),
)
# 2. Add body_rules to provider_endpoints
# 请求体规则支持三种操作:
# - set: 设置/覆盖字段 {"action": "set", "path": "metadata", "value": {"custom": "val"}}
# - drop: 删除字段 {"action": "drop", "path": "unwanted_field"}
# - rename: 重命名字段 {"action": "rename", "from": "old_key", "to": "new_key"}
if not _column_exists("provider_endpoints", "body_rules"):
op.add_column(
"provider_endpoints",
sa.Column("body_rules", sa.JSON(), nullable=True),
)
def downgrade() -> None:
# Remove body_rules from provider_endpoints
if _column_exists("provider_endpoints", "body_rules"):
op.drop_column("provider_endpoints", "body_rules")
# Remove video_duration_seconds from video_tasks
if _column_exists("video_tasks", "video_duration_seconds"):
op.drop_column("video_tasks", "video_duration_seconds")
@@ -0,0 +1,347 @@
"""add_stats_hourly_and_daily_complete_flag
Revision ID: c4e8f9a1b2c3
Revises: b3c4d5e6f7a8
Create Date: 2026-02-04 12:00:00.000000
"""
from __future__ import annotations
from typing import Sequence, Union
import sqlalchemy as sa
from sqlalchemy import inspect
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "c4e8f9a1b2c3"
down_revision: Union[str, None] = "b3c4d5e6f7a8"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def _table_exists(table_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
return table_name in inspector.get_table_names()
def _index_exists(table_name: str, index_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
indexes = [idx["name"] for idx in inspector.get_indexes(table_name)]
return index_name in indexes
def _column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
# Use information_schema for more reliable detection (inspector can have caching issues)
result = bind.execute(
sa.text(
"SELECT EXISTS ("
"SELECT 1 FROM information_schema.columns "
"WHERE table_name = :table AND column_name = :column"
")"
),
{"table": table_name, "column": column_name},
)
return bool(result.scalar())
def upgrade() -> None:
if _table_exists("stats_daily"):
if not _column_exists("stats_daily", "is_complete"):
op.add_column(
"stats_daily",
sa.Column("is_complete", sa.Boolean(), nullable=False, server_default=sa.false()),
)
op.execute("UPDATE stats_daily SET is_complete = true")
if not _column_exists("stats_daily", "aggregated_at"):
op.add_column(
"stats_daily",
sa.Column("aggregated_at", sa.DateTime(timezone=True), nullable=True),
)
if not _table_exists("stats_hourly"):
op.create_table(
"stats_hourly",
sa.Column("id", sa.String(length=36), nullable=False),
sa.Column("hour_utc", sa.DateTime(timezone=True), nullable=False),
sa.Column("total_requests", sa.Integer(), nullable=False),
sa.Column("success_requests", sa.Integer(), nullable=False),
sa.Column("error_requests", sa.Integer(), nullable=False),
sa.Column("input_tokens", sa.BigInteger(), nullable=False),
sa.Column("output_tokens", sa.BigInteger(), nullable=False),
sa.Column("cache_creation_tokens", sa.BigInteger(), nullable=False),
sa.Column("cache_read_tokens", sa.BigInteger(), nullable=False),
sa.Column("total_cost", sa.Float(), nullable=False),
sa.Column("actual_total_cost", sa.Float(), nullable=False),
sa.Column("avg_response_time_ms", sa.Float(), nullable=False),
sa.Column("is_complete", sa.Boolean(), nullable=False),
sa.Column("aggregated_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False),
sa.PrimaryKeyConstraint("id"),
sa.UniqueConstraint("hour_utc", name="uq_stats_hourly_hour"),
)
op.create_index("idx_stats_hourly_hour", "stats_hourly", ["hour_utc"], unique=False)
if not _table_exists("stats_hourly_user"):
op.create_table(
"stats_hourly_user",
sa.Column("id", sa.String(length=36), nullable=False),
sa.Column("hour_utc", sa.DateTime(timezone=True), nullable=False),
sa.Column("user_id", sa.String(length=36), nullable=False),
sa.Column("total_requests", sa.Integer(), nullable=False),
sa.Column("success_requests", sa.Integer(), nullable=False),
sa.Column("error_requests", sa.Integer(), nullable=False),
sa.Column("input_tokens", sa.BigInteger(), nullable=False),
sa.Column("output_tokens", sa.BigInteger(), nullable=False),
sa.Column("total_cost", sa.Float(), nullable=False),
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False),
sa.PrimaryKeyConstraint("id"),
sa.UniqueConstraint("hour_utc", "user_id", name="uq_stats_hourly_user"),
)
op.create_index(
"idx_stats_hourly_user_hour", "stats_hourly_user", ["hour_utc"], unique=False
)
op.create_index(
"idx_stats_hourly_user_user_hour",
"stats_hourly_user",
["user_id", "hour_utc"],
unique=False,
)
if not _table_exists("stats_hourly_model"):
op.create_table(
"stats_hourly_model",
sa.Column("id", sa.String(length=36), nullable=False),
sa.Column("hour_utc", sa.DateTime(timezone=True), nullable=False),
sa.Column("model", sa.String(length=100), nullable=False),
sa.Column("total_requests", sa.Integer(), nullable=False),
sa.Column("input_tokens", sa.BigInteger(), nullable=False),
sa.Column("output_tokens", sa.BigInteger(), nullable=False),
sa.Column("total_cost", sa.Float(), nullable=False),
sa.Column("avg_response_time_ms", sa.Float(), nullable=False),
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False),
sa.PrimaryKeyConstraint("id"),
sa.UniqueConstraint("hour_utc", "model", name="uq_stats_hourly_model"),
)
op.create_index(
"idx_stats_hourly_model_hour", "stats_hourly_model", ["hour_utc"], unique=False
)
op.create_index(
"idx_stats_hourly_model_model_hour",
"stats_hourly_model",
["model", "hour_utc"],
unique=False,
)
if not _table_exists("stats_hourly_provider"):
op.create_table(
"stats_hourly_provider",
sa.Column("id", sa.String(length=36), nullable=False),
sa.Column("hour_utc", sa.DateTime(timezone=True), nullable=False),
sa.Column("provider_name", sa.String(length=100), nullable=False),
sa.Column("total_requests", sa.Integer(), nullable=False),
sa.Column("input_tokens", sa.BigInteger(), nullable=False),
sa.Column("output_tokens", sa.BigInteger(), nullable=False),
sa.Column("total_cost", sa.Float(), nullable=False),
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False),
sa.PrimaryKeyConstraint("id"),
sa.UniqueConstraint("hour_utc", "provider_name", name="uq_stats_hourly_provider"),
)
op.create_index(
"idx_stats_hourly_provider_hour",
"stats_hourly_provider",
["hour_utc"],
unique=False,
)
if not _table_exists("stats_daily_api_key"):
op.create_table(
"stats_daily_api_key",
sa.Column("id", sa.String(length=36), primary_key=True),
sa.Column(
"api_key_id",
sa.String(length=36),
sa.ForeignKey("api_keys.id", ondelete="CASCADE"),
nullable=False,
),
sa.Column("date", sa.DateTime(timezone=True), nullable=False),
sa.Column("total_requests", sa.Integer(), nullable=False, server_default="0"),
sa.Column("success_requests", sa.Integer(), nullable=False, server_default="0"),
sa.Column("error_requests", sa.Integer(), nullable=False, server_default="0"),
sa.Column("input_tokens", sa.BigInteger(), nullable=False, server_default="0"),
sa.Column("output_tokens", sa.BigInteger(), nullable=False, server_default="0"),
sa.Column("cache_creation_tokens", sa.BigInteger(), nullable=False, server_default="0"),
sa.Column("cache_read_tokens", sa.BigInteger(), nullable=False, server_default="0"),
sa.Column("total_cost", sa.Float(), nullable=False, server_default="0"),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("CURRENT_TIMESTAMP"),
),
sa.Column(
"updated_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("CURRENT_TIMESTAMP"),
),
sa.UniqueConstraint("api_key_id", "date", name="uq_stats_daily_api_key"),
)
if _table_exists("stats_daily_api_key"):
if not _index_exists("stats_daily_api_key", "idx_stats_daily_api_key_date"):
op.create_index("idx_stats_daily_api_key_date", "stats_daily_api_key", ["date"])
if not _index_exists("stats_daily_api_key", "idx_stats_daily_api_key_key_date"):
op.create_index(
"idx_stats_daily_api_key_key_date",
"stats_daily_api_key",
["api_key_id", "date"],
)
if not _index_exists("stats_daily_api_key", "idx_stats_daily_api_key_date_requests"):
op.create_index(
"idx_stats_daily_api_key_date_requests",
"stats_daily_api_key",
["date", "total_requests"],
)
if not _index_exists("stats_daily_api_key", "idx_stats_daily_api_key_date_cost"):
op.create_index(
"idx_stats_daily_api_key_date_cost",
"stats_daily_api_key",
["date", "total_cost"],
)
if _table_exists("usage"):
if not _column_exists("usage", "error_category"):
op.add_column(
"usage",
sa.Column("error_category", sa.String(length=50), nullable=True),
)
op.create_index("idx_usage_error_category", "usage", ["error_category"], unique=False)
if _table_exists("stats_daily"):
for name in (
"p50_response_time_ms",
"p90_response_time_ms",
"p99_response_time_ms",
"p50_first_byte_time_ms",
"p90_first_byte_time_ms",
"p99_first_byte_time_ms",
):
if not _column_exists("stats_daily", name):
op.add_column("stats_daily", sa.Column(name, sa.Integer(), nullable=True))
if not _table_exists("stats_daily_error"):
op.create_table(
"stats_daily_error",
sa.Column("id", sa.String(length=36), primary_key=True),
sa.Column("date", sa.DateTime(timezone=True), nullable=False),
sa.Column("error_category", sa.String(length=50), nullable=False),
sa.Column("provider_name", sa.String(length=100), nullable=True),
sa.Column("model", sa.String(length=100), nullable=True),
sa.Column("count", sa.Integer(), nullable=False, server_default="0"),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("CURRENT_TIMESTAMP"),
),
sa.Column(
"updated_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("CURRENT_TIMESTAMP"),
),
sa.UniqueConstraint(
"date",
"error_category",
"provider_name",
"model",
name="uq_stats_daily_error",
),
)
if _table_exists("stats_daily_error"):
if not _index_exists("stats_daily_error", "idx_stats_daily_error_date"):
op.create_index("idx_stats_daily_error_date", "stats_daily_error", ["date"])
if not _index_exists("stats_daily_error", "idx_stats_daily_error_category"):
op.create_index(
"idx_stats_daily_error_category",
"stats_daily_error",
["date", "error_category"],
)
def downgrade() -> None:
if _table_exists("stats_daily_error"):
if _index_exists("stats_daily_error", "idx_stats_daily_error_category"):
op.drop_index("idx_stats_daily_error_category", table_name="stats_daily_error")
if _index_exists("stats_daily_error", "idx_stats_daily_error_date"):
op.drop_index("idx_stats_daily_error_date", table_name="stats_daily_error")
op.drop_table("stats_daily_error")
if _table_exists("stats_daily"):
for name in (
"p50_response_time_ms",
"p90_response_time_ms",
"p99_response_time_ms",
"p50_first_byte_time_ms",
"p90_first_byte_time_ms",
"p99_first_byte_time_ms",
):
if _column_exists("stats_daily", name):
op.drop_column("stats_daily", name)
if _table_exists("usage") and _column_exists("usage", "error_category"):
if _index_exists("usage", "idx_usage_error_category"):
op.drop_index("idx_usage_error_category", table_name="usage")
op.drop_column("usage", "error_category")
if _table_exists("stats_daily_api_key"):
if _index_exists("stats_daily_api_key", "idx_stats_daily_api_key_date_cost"):
op.drop_index("idx_stats_daily_api_key_date_cost", table_name="stats_daily_api_key")
if _index_exists("stats_daily_api_key", "idx_stats_daily_api_key_date_requests"):
op.drop_index("idx_stats_daily_api_key_date_requests", table_name="stats_daily_api_key")
if _index_exists("stats_daily_api_key", "idx_stats_daily_api_key_key_date"):
op.drop_index("idx_stats_daily_api_key_key_date", table_name="stats_daily_api_key")
if _index_exists("stats_daily_api_key", "idx_stats_daily_api_key_date"):
op.drop_index("idx_stats_daily_api_key_date", table_name="stats_daily_api_key")
op.drop_table("stats_daily_api_key")
if _table_exists("stats_hourly_provider"):
if _index_exists("stats_hourly_provider", "idx_stats_hourly_provider_hour"):
op.drop_index("idx_stats_hourly_provider_hour", table_name="stats_hourly_provider")
op.drop_table("stats_hourly_provider")
if _table_exists("stats_hourly_model"):
if _index_exists("stats_hourly_model", "idx_stats_hourly_model_model_hour"):
op.drop_index("idx_stats_hourly_model_model_hour", table_name="stats_hourly_model")
if _index_exists("stats_hourly_model", "idx_stats_hourly_model_hour"):
op.drop_index("idx_stats_hourly_model_hour", table_name="stats_hourly_model")
op.drop_table("stats_hourly_model")
if _table_exists("stats_hourly_user"):
if _index_exists("stats_hourly_user", "idx_stats_hourly_user_user_hour"):
op.drop_index("idx_stats_hourly_user_user_hour", table_name="stats_hourly_user")
if _index_exists("stats_hourly_user", "idx_stats_hourly_user_hour"):
op.drop_index("idx_stats_hourly_user_hour", table_name="stats_hourly_user")
op.drop_table("stats_hourly_user")
if _table_exists("stats_hourly"):
if _index_exists("stats_hourly", "idx_stats_hourly_hour"):
op.drop_index("idx_stats_hourly_hour", table_name="stats_hourly")
op.drop_table("stats_hourly")
if _table_exists("stats_daily"):
if _column_exists("stats_daily", "aggregated_at"):
op.drop_column("stats_daily", "aggregated_at")
if _column_exists("stats_daily", "is_complete"):
op.drop_column("stats_daily", "is_complete")
@@ -0,0 +1,205 @@
"""Add provider_type, upstream_metadata, oauth_invalid fields and expand string columns to TEXT
- Add providers.provider_type (String(20), server_default="custom")
- Add provider_api_keys.upstream_metadata (JSON, nullable)
- Add provider_api_keys.oauth_invalid_at (DateTime, nullable) - OAuth Token 失效时间
- Add provider_api_keys.oauth_invalid_reason (String(255), nullable) - OAuth Token 失效原因
- Expand multiple VARCHAR columns to TEXT for long values (OAuth tokens, LDAP DN, URLs, etc.)
Revision ID: b5c6d7e8f9a0
Revises: c4e8f9a1b2c3
Create Date: 2026-02-04 15:00:00.000000
"""
from typing import Sequence, Union
import sqlalchemy as sa
from sqlalchemy import inspect
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "b5c6d7e8f9a0"
down_revision: Union[str, None] = "c4e8f9a1b2c3"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
# 需要扩展为 TEXT 的列(表名, 列名, 原始类型长度)
COLUMNS_TO_EXPAND = [
("provider_api_keys", "api_key", 500), # OAuth tokens can be very long
(
"provider_api_keys",
"auth_config",
None,
), # 确保 auth_config 是 TEXT 类型(可能从 JSON 迁移过来)
("ldap_configs", "bind_dn", 255), # LDAP DN can be deeply nested
("ldap_configs", "base_dn", 255), # LDAP DN can be deeply nested
("ldap_configs", "user_search_filter", 500), # Complex LDAP filters
("oauth_providers", "client_id", 255), # Some OAuth providers use JWT client_id
]
def column_exists(table_name: str, column_name: str) -> bool:
"""检查列是否已存在"""
bind = op.get_bind()
inspector = inspect(bind)
columns = [col["name"] for col in inspector.get_columns(table_name)]
return column_name in columns
def table_exists(table_name: str) -> bool:
"""检查表是否存在"""
bind = op.get_bind()
inspector = inspect(bind)
return table_name in inspector.get_table_names()
def is_sqlite() -> bool:
"""检查是否为 SQLite 数据库"""
bind = op.get_bind()
return bind.dialect.name == "sqlite"
def get_column_type(table_name: str, column_name: str) -> str | None:
"""获取列的数据类型"""
bind = op.get_bind()
inspector = inspect(bind)
for col in inspector.get_columns(table_name):
if col["name"] == column_name:
return str(col["type"]).upper()
return None
def expand_column_to_text(table_name: str, column_name: str, original_length: int | None) -> None:
"""将 VARCHAR 列扩展为 TEXT(兼容 SQLite)"""
if not table_exists(table_name):
return
if not column_exists(table_name, column_name):
return
# 检查当前列类型,如果已经是 TEXT 则跳过
col_type = get_column_type(table_name, column_name)
if col_type and "TEXT" in col_type:
return
# 如果是 JSON 类型(可能是历史遗留),先将 JSON 数据转为文本表示再变更类型
is_json_col = col_type and "JSON" in col_type
if is_json_col and not is_sqlite():
# PostgreSQL: 先用 CAST 把 JSON 值转为 TEXT,保留数据
op.execute(
sa.text(
f"ALTER TABLE {table_name} ALTER COLUMN {column_name} "
f"TYPE TEXT USING {column_name}::TEXT"
)
)
return
if is_sqlite():
# SQLite 不支持直接 ALTER COLUMN,需要用 batch 模式
# batch 模式会自动处理 JSON->TEXT 的数据迁移
with op.batch_alter_table(table_name) as batch_op:
batch_op.alter_column(
column_name,
type_=sa.Text(),
existing_type=sa.String(original_length) if original_length else sa.Text(),
)
else:
op.alter_column(
table_name,
column_name,
type_=sa.Text(),
existing_type=sa.String(original_length) if original_length else sa.Text(),
existing_nullable=True,
)
def shrink_column_to_varchar(
table_name: str, column_name: str, target_length: int, nullable: bool = False
) -> None:
"""将 TEXT 列缩小为 VARCHAR(兼容 SQLite)
WARNING: 如果数据超过 target_length 会失败
"""
if not table_exists(table_name):
return
if not column_exists(table_name, column_name):
return
if is_sqlite():
with op.batch_alter_table(table_name) as batch_op:
batch_op.alter_column(
column_name,
type_=sa.String(target_length),
existing_type=sa.Text(),
existing_nullable=nullable,
)
else:
op.alter_column(
table_name,
column_name,
type_=sa.String(target_length),
existing_type=sa.Text(),
existing_nullable=nullable,
)
def upgrade() -> None:
# Add providers.provider_type
if not column_exists("providers", "provider_type"):
op.add_column(
"providers",
sa.Column("provider_type", sa.String(20), nullable=False, server_default="custom"),
)
# Add provider_api_keys.upstream_metadata
if not column_exists("provider_api_keys", "upstream_metadata"):
op.add_column(
"provider_api_keys",
sa.Column("upstream_metadata", sa.JSON(), nullable=True),
)
# Add provider_api_keys.oauth_invalid_at
if not column_exists("provider_api_keys", "oauth_invalid_at"):
op.add_column(
"provider_api_keys",
sa.Column("oauth_invalid_at", sa.DateTime(timezone=True), nullable=True),
)
# Add provider_api_keys.oauth_invalid_reason
if not column_exists("provider_api_keys", "oauth_invalid_reason"):
op.add_column(
"provider_api_keys",
sa.Column("oauth_invalid_reason", sa.String(255), nullable=True),
)
# Expand VARCHAR columns to TEXT
for table_name, column_name, original_length in COLUMNS_TO_EXPAND:
expand_column_to_text(table_name, column_name, original_length)
def downgrade() -> None:
# Shrink TEXT columns back to VARCHAR
# WARNING: Downgrade may fail if any values exceed original length
for table_name, column_name, original_length in reversed(COLUMNS_TO_EXPAND):
# 跳过没有原始长度的列(如 auth_config,由其他迁移创建)
if original_length is None:
continue
shrink_column_to_varchar(table_name, column_name, original_length)
# Drop provider_api_keys.oauth_invalid_reason
if column_exists("provider_api_keys", "oauth_invalid_reason"):
op.drop_column("provider_api_keys", "oauth_invalid_reason")
# Drop provider_api_keys.oauth_invalid_at
if column_exists("provider_api_keys", "oauth_invalid_at"):
op.drop_column("provider_api_keys", "oauth_invalid_at")
# Drop provider_api_keys.upstream_metadata
if column_exists("provider_api_keys", "upstream_metadata"):
op.drop_column("provider_api_keys", "upstream_metadata")
# Drop providers.provider_type
if column_exists("providers", "provider_type"):
op.drop_column("providers", "provider_type")
@@ -0,0 +1,254 @@
"""Antigravity endpoint signature to gemini:chat & add proxy_nodes table (with manual fields)
Revision ID: e1b2c3d4f5a6
Revises: b5c6d7e8f9a0
Create Date: 2026-02-06 23:45:00.000000
"""
from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from sqlalchemy import inspect, text
from sqlalchemy.dialects import postgresql
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "e1b2c3d4f5a6"
down_revision: str | None = "b5c6d7e8f9a0"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def table_exists(table_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
return table_name in inspector.get_table_names()
def upgrade() -> None:
conn = op.get_bind()
# =========================================================================
# Part 1: Antigravity endpoint signature migration (gemini:cli -> gemini:chat)
# =========================================================================
# --- provider_endpoints ---
# Update only when there is no conflicting gemini:chat endpoint for the same provider
# (provider_endpoints has a unique constraint on (provider_id, api_format)).
conn.execute(text("""
UPDATE provider_endpoints pe
SET
api_format = 'gemini:chat',
api_family = 'gemini',
endpoint_kind = 'chat'
WHERE pe.api_format = 'gemini:cli'
AND pe.provider_id IN (
SELECT id FROM providers WHERE lower(provider_type) = 'antigravity'
)
AND NOT EXISTS (
SELECT 1 FROM provider_endpoints pe2
WHERE pe2.provider_id = pe.provider_id
AND pe2.api_format = 'gemini:chat'
)
"""))
# Best-effort normalization for already-existing Antigravity gemini:chat endpoints.
conn.execute(text("""
UPDATE provider_endpoints pe
SET
api_family = 'gemini',
endpoint_kind = 'chat'
WHERE pe.api_format = 'gemini:chat'
AND pe.provider_id IN (
SELECT id FROM providers WHERE lower(provider_type) = 'antigravity'
)
"""))
# --- provider_api_keys.api_formats (JSON array) ---
# Replace "gemini:cli" with "gemini:chat" in the JSON array for Antigravity keys.
# Uses text-level replace on the serialized JSON -- safe because the value is a
# simple string with no special characters that could cause ambiguous replacements.
conn.execute(text("""
UPDATE provider_api_keys pak
SET api_formats = replace(pak.api_formats::text, '"gemini:cli"', '"gemini:chat"')::json
WHERE pak.provider_id IN (
SELECT id FROM providers WHERE lower(provider_type) = 'antigravity'
)
AND pak.api_formats IS NOT NULL
AND pak.api_formats::text LIKE '%"gemini:cli"%'
"""))
# =========================================================================
# Part 2: Create proxy_nodes table with manual proxy fields (idempotent)
# =========================================================================
# Create ENUM type (idempotent)
op.execute(
"DO $$ BEGIN "
"CREATE TYPE proxynodestatus AS ENUM ('online', 'unhealthy', 'offline'); "
"EXCEPTION WHEN duplicate_object THEN NULL; "
"END $$"
)
if table_exists("proxy_nodes"):
# Table already exists — ensure manual proxy columns are present
inspector = inspect(conn)
existing_columns = {c["name"] for c in inspector.get_columns("proxy_nodes")}
# ip 列扩容:手动节点的 ip 存储 "socks5://hostname" 形式,45 字符可能不够
ip_col = next((c for c in inspector.get_columns("proxy_nodes") if c["name"] == "ip"), None)
if ip_col and hasattr(ip_col["type"], "length") and (ip_col["type"].length or 0) < 512:
op.alter_column("proxy_nodes", "ip", type_=sa.String(512), existing_nullable=False)
manual_columns = [
("is_manual", sa.Boolean(), False, sa.text("false"), "是否为手动添加的代理节点"),
("proxy_url", sa.String(500), True, None, "手动节点的完整代理 URL"),
("proxy_username", sa.String(255), True, None, "手动节点的代理用户名"),
("proxy_password", sa.String(500), True, None, "手动节点的代理密码"),
]
for col_name, col_type, nullable, default, comment in manual_columns:
if col_name not in existing_columns:
op.add_column(
"proxy_nodes",
sa.Column(
col_name,
col_type, # type: ignore[arg-type]
nullable=nullable,
server_default=default,
comment=comment,
),
)
return
op.create_table(
"proxy_nodes",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("name", sa.String(100), nullable=False),
sa.Column("ip", sa.String(512), nullable=False),
sa.Column("port", sa.Integer(), nullable=False),
sa.Column("region", sa.String(100), nullable=True),
sa.Column(
"status",
postgresql.ENUM(
"online",
"unhealthy",
"offline",
name="proxynodestatus",
create_type=False,
),
nullable=False,
server_default=sa.text("'online'"),
),
sa.Column(
"registered_by",
sa.String(36),
sa.ForeignKey("users.id", ondelete="SET NULL"),
nullable=True,
),
sa.Column("last_heartbeat_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("heartbeat_interval", sa.Integer(), nullable=False, server_default=sa.text("30")),
sa.Column("active_connections", sa.Integer(), nullable=False, server_default=sa.text("0")),
sa.Column("total_requests", sa.BigInteger(), nullable=False, server_default=sa.text("0")),
sa.Column("avg_latency_ms", sa.Float(), nullable=True),
# --- Manual proxy node fields ---
sa.Column(
"is_manual",
sa.Boolean(),
nullable=False,
server_default=sa.text("false"),
comment="是否为手动添加的代理节点",
),
sa.Column(
"proxy_url",
sa.String(500),
nullable=True,
comment="手动节点的完整代理 URL",
),
sa.Column(
"proxy_username",
sa.String(255),
nullable=True,
comment="手动节点的代理用户名",
),
sa.Column(
"proxy_password",
sa.String(500),
nullable=True,
comment="手动节点的代理密码",
),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
server_default=sa.text("CURRENT_TIMESTAMP"),
nullable=False,
),
sa.Column(
"updated_at",
sa.DateTime(timezone=True),
server_default=sa.text("CURRENT_TIMESTAMP"),
nullable=False,
),
sa.UniqueConstraint("ip", "port", name="uq_proxy_node_ip_port"),
)
def downgrade() -> None:
conn = op.get_bind()
# =========================================================================
# Part 2 rollback: Drop proxy_nodes table (and manual columns if present)
# =========================================================================
if table_exists("proxy_nodes"):
op.drop_table("proxy_nodes")
# Best-effort: drop type (only used by proxy_nodes)
op.execute("DROP TYPE IF EXISTS proxynodestatus")
# =========================================================================
# Part 1 rollback: Revert Antigravity endpoint signature (gemini:chat -> gemini:cli)
# =========================================================================
# --- provider_endpoints ---
conn.execute(text("""
UPDATE provider_endpoints pe
SET
api_format = 'gemini:cli',
api_family = 'gemini',
endpoint_kind = 'cli'
WHERE pe.api_format = 'gemini:chat'
AND pe.provider_id IN (
SELECT id FROM providers WHERE lower(provider_type) = 'antigravity'
)
AND NOT EXISTS (
SELECT 1 FROM provider_endpoints pe2
WHERE pe2.provider_id = pe.provider_id
AND pe2.api_format = 'gemini:cli'
)
"""))
# Best-effort normalization for already-existing Antigravity gemini:cli endpoints.
conn.execute(text("""
UPDATE provider_endpoints pe
SET
api_family = 'gemini',
endpoint_kind = 'cli'
WHERE pe.api_format = 'gemini:cli'
AND pe.provider_id IN (
SELECT id FROM providers WHERE lower(provider_type) = 'antigravity'
)
"""))
# --- provider_api_keys.api_formats (JSON array) ---
conn.execute(text("""
UPDATE provider_api_keys pak
SET api_formats = replace(pak.api_formats::text, '"gemini:chat"', '"gemini:cli"')::json
WHERE pak.provider_id IN (
SELECT id FROM providers WHERE lower(provider_type) = 'antigravity'
)
AND pak.api_formats IS NOT NULL
AND pak.api_formats::text LIKE '%"gemini:chat"%'
"""))
@@ -0,0 +1,61 @@
"""Add remote_config and config_version to proxy_nodes
Revision ID: 3aff3ffc4a0e
Revises: e1b2c3d4f5a6
Create Date: 2026-02-07 15:00:00.000000
"""
from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from sqlalchemy import inspect
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "3aff3ffc4a0e"
down_revision: str | None = "e1b2c3d4f5a6"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
columns = [c["name"] for c in inspector.get_columns(table_name)]
return column_name in columns
def upgrade() -> None:
if not column_exists("proxy_nodes", "remote_config"):
op.add_column(
"proxy_nodes",
sa.Column(
"remote_config",
sa.JSON(),
nullable=True,
comment="管理端下发的远程配置 (allowed_ports, log_level, heartbeat_interval, timestamp_tolerance)",
),
)
if not column_exists("proxy_nodes", "config_version"):
op.add_column(
"proxy_nodes",
sa.Column(
"config_version",
sa.Integer(),
nullable=False,
server_default="0",
comment="远程配置版本号,每次更新 +1",
),
)
def downgrade() -> None:
if column_exists("proxy_nodes", "config_version"):
op.drop_column("proxy_nodes", "config_version")
if column_exists("proxy_nodes", "remote_config"):
op.drop_column("proxy_nodes", "remote_config")
@@ -0,0 +1,61 @@
"""Add tls_enabled and tls_cert_fingerprint to proxy_nodes
Revision ID: 4b5c6d7e8f9a
Revises: 3aff3ffc4a0e
Create Date: 2026-02-07 18:00:00.000000
"""
from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from sqlalchemy import inspect
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "4b5c6d7e8f9a"
down_revision: str | None = "3aff3ffc4a0e"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
columns = [c["name"] for c in inspector.get_columns(table_name)]
return column_name in columns
def upgrade() -> None:
if not column_exists("proxy_nodes", "tls_enabled"):
op.add_column(
"proxy_nodes",
sa.Column(
"tls_enabled",
sa.Boolean(),
nullable=False,
server_default="false",
comment="是否启用 TLS 加密",
),
)
if not column_exists("proxy_nodes", "tls_cert_fingerprint"):
op.add_column(
"proxy_nodes",
sa.Column(
"tls_cert_fingerprint",
sa.String(128),
nullable=True,
comment="TLS 证书 SHA-256 指纹(hex)",
),
)
def downgrade() -> None:
if column_exists("proxy_nodes", "tls_cert_fingerprint"):
op.drop_column("proxy_nodes", "tls_cert_fingerprint")
if column_exists("proxy_nodes", "tls_enabled"):
op.drop_column("proxy_nodes", "tls_enabled")
@@ -0,0 +1,60 @@
"""Add hardware_info and estimated_max_concurrency to proxy_nodes
Revision ID: 5c6d7e8f9a0b
Revises: 4b5c6d7e8f9a
Create Date: 2026-02-08 12:00:00.000000
"""
from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from sqlalchemy import inspect
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "5c6d7e8f9a0b"
down_revision: str | None = "4b5c6d7e8f9a"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
columns = [c["name"] for c in inspector.get_columns(table_name)]
return column_name in columns
def upgrade() -> None:
if not column_exists("proxy_nodes", "hardware_info"):
op.add_column(
"proxy_nodes",
sa.Column(
"hardware_info",
sa.JSON(),
nullable=True,
comment="硬件信息 (cpu_cores, total_memory_mb, os_info, fd_limit, ...)",
),
)
if not column_exists("proxy_nodes", "estimated_max_concurrency"):
op.add_column(
"proxy_nodes",
sa.Column(
"estimated_max_concurrency",
sa.Integer(),
nullable=True,
comment="基于硬件估算的最大并发连接数",
),
)
def downgrade() -> None:
if column_exists("proxy_nodes", "estimated_max_concurrency"):
op.drop_column("proxy_nodes", "estimated_max_concurrency")
if column_exists("proxy_nodes", "hardware_info"):
op.drop_column("proxy_nodes", "hardware_info")
@@ -0,0 +1,47 @@
"""Add proxy column to provider_api_keys for per-key proxy configuration
Revision ID: 6d7e8f9a0b1c
Revises: 5c6d7e8f9a0b
Create Date: 2026-02-08 15:00:00.000000
"""
from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from sqlalchemy import inspect
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "6d7e8f9a0b1c"
down_revision: str | None = "5c6d7e8f9a0b"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
columns = [c["name"] for c in inspector.get_columns(table_name)]
return column_name in columns
def upgrade() -> None:
if not column_exists("provider_api_keys", "proxy"):
op.add_column(
"provider_api_keys",
sa.Column(
"proxy",
sa.JSON(),
nullable=True,
comment="Key 级别代理配置(覆盖 Provider 级别代理),如 {node_id, enabled}",
),
)
def downgrade() -> None:
if column_exists("provider_api_keys", "proxy"):
op.drop_column("provider_api_keys", "proxy")
@@ -0,0 +1,46 @@
"""Add provider_request_body and client_response_body columns to usage table
Revision ID: 7e8f9a0b1c2d
Revises: 6d7e8f9a0b1c
Create Date: 2026-02-20 18:00:00.000000
"""
from __future__ import annotations
from collections.abc import Sequence
from sqlalchemy import text
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "7e8f9a0b1c2d"
down_revision: str | None = "6d7e8f9a0b1c"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
conn = op.get_bind()
# Use PostgreSQL native IF NOT EXISTS to avoid duplicate-column races
# when migrations are triggered concurrently (e.g. startup + manual run).
conn.execute(text("ALTER TABLE usage ADD COLUMN IF NOT EXISTS provider_request_body JSON"))
conn.execute(
text("ALTER TABLE usage ADD COLUMN IF NOT EXISTS provider_request_body_compressed BYTEA")
)
conn.execute(text("ALTER TABLE usage ADD COLUMN IF NOT EXISTS client_response_body JSON"))
conn.execute(
text("ALTER TABLE usage ADD COLUMN IF NOT EXISTS client_response_body_compressed BYTEA")
)
def downgrade() -> None:
conn = op.get_bind()
for col in (
"client_response_body_compressed",
"client_response_body",
"provider_request_body_compressed",
"provider_request_body",
):
conn.execute(text(f"ALTER TABLE usage DROP COLUMN IF EXISTS {col}"))
@@ -0,0 +1,91 @@
"""Add api_family and endpoint_kind columns to usage table
Revision ID: 8f9a0b1c2d3e
Revises: 7e8f9a0b1c2d
Create Date: 2026-02-21 15:00:00.000000
"""
from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from sqlalchemy import inspect, text
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "8f9a0b1c2d3e"
down_revision: str | None = "7e8f9a0b1c2d"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
# Usage 表新增列
NEW_COLUMNS = [
("api_family", sa.String(50)),
("endpoint_kind", sa.String(50)),
("provider_api_family", sa.String(50)),
("provider_endpoint_kind", sa.String(50)),
]
# 新增索引
NEW_INDEXES = [
("idx_usage_api_family", "usage", ["api_family"]),
("idx_usage_endpoint_kind", "usage", ["endpoint_kind"]),
("idx_usage_family_kind", "usage", ["api_family", "endpoint_kind"]),
]
def upgrade() -> None:
conn = op.get_bind()
# 使用 PostgreSQL 原生 IF NOT EXISTS,比 inspect 更可靠(避免同一事务内缓存问题)
col_definitions = {
"api_family": "VARCHAR(50)",
"endpoint_kind": "VARCHAR(50)",
"provider_api_family": "VARCHAR(50)",
"provider_endpoint_kind": "VARCHAR(50)",
}
for col_name, col_type_sql in col_definitions.items():
conn.execute(text(f"ALTER TABLE usage ADD COLUMN IF NOT EXISTS {col_name} {col_type_sql}"))
# 数据迁移:从 api_format 解析 api_family + endpoint_kind
conn.execute(text("""
UPDATE usage SET
api_family = lower(split_part(api_format, ':', 1)),
endpoint_kind = lower(split_part(api_format, ':', 2))
WHERE api_format IS NOT NULL
AND api_format LIKE '%%:%%'
AND api_family IS NULL
"""))
conn.execute(text("""
UPDATE usage SET
provider_api_family = lower(split_part(endpoint_api_format, ':', 1)),
provider_endpoint_kind = lower(split_part(endpoint_api_format, ':', 2))
WHERE endpoint_api_format IS NOT NULL
AND endpoint_api_format LIKE '%%:%%'
AND provider_api_family IS NULL
"""))
# 创建索引
inspector = inspect(conn)
existing_indexes = {idx["name"] for idx in inspector.get_indexes("usage")}
for idx_name, table, columns in NEW_INDEXES:
if idx_name not in existing_indexes:
op.create_index(idx_name, table, columns)
def downgrade() -> None:
conn = op.get_bind()
inspector = inspect(conn)
existing_indexes = {idx["name"] for idx in inspector.get_indexes("usage")}
for idx_name, _, _ in reversed(NEW_INDEXES):
if idx_name in existing_indexes:
op.drop_index(idx_name, table_name="usage")
existing_columns = {col["name"] for col in inspector.get_columns("usage")}
for col_name, _ in reversed(NEW_COLUMNS):
if col_name in existing_columns:
op.drop_column("usage", col_name)
@@ -0,0 +1,106 @@
"""Add tunnel mode fields and remove IP forwarding fields
Revision ID: 9a0b1c2d3e4f
Revises: 8f9a0b1c2d3e
Create Date: 2026-02-24 17:00:00.000000
"""
from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from sqlalchemy import inspect
from alembic import op
revision: str = "9a0b1c2d3e4f"
down_revision: str | None = "8f9a0b1c2d3e"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
columns = [c["name"] for c in inspector.get_columns(table_name)]
return column_name in columns
def upgrade() -> None:
# 添加 tunnel 模式字段
if not column_exists("proxy_nodes", "tunnel_mode"):
op.add_column(
"proxy_nodes",
sa.Column(
"tunnel_mode",
sa.Boolean(),
nullable=False,
server_default=sa.text("false"),
comment="是否使用 WebSocket 隧道模式",
),
)
if not column_exists("proxy_nodes", "tunnel_connected"):
op.add_column(
"proxy_nodes",
sa.Column(
"tunnel_connected",
sa.Boolean(),
nullable=False,
server_default=sa.text("false"),
comment="隧道是否已连接",
),
)
if not column_exists("proxy_nodes", "tunnel_connected_at"):
op.add_column(
"proxy_nodes",
sa.Column(
"tunnel_connected_at",
sa.DateTime(timezone=True),
nullable=True,
comment="隧道最近一次建立时间",
),
)
# tunnel 模式节点不需要 port,将其置零
op.execute("UPDATE proxy_nodes SET port = 0 WHERE tunnel_mode = true")
# 移除旧的 IP 转发字段
if column_exists("proxy_nodes", "tls_enabled"):
op.drop_column("proxy_nodes", "tls_enabled")
if column_exists("proxy_nodes", "tls_cert_fingerprint"):
op.drop_column("proxy_nodes", "tls_cert_fingerprint")
def downgrade() -> None:
# 恢复 IP 转发字段
if not column_exists("proxy_nodes", "tls_cert_fingerprint"):
op.add_column(
"proxy_nodes",
sa.Column(
"tls_cert_fingerprint",
sa.String(128),
nullable=True,
comment="TLS 证书 SHA-256 指纹(hex)",
),
)
if not column_exists("proxy_nodes", "tls_enabled"):
op.add_column(
"proxy_nodes",
sa.Column(
"tls_enabled",
sa.Boolean(),
nullable=False,
server_default=sa.text("false"),
comment="是否启用 TLS 加密",
),
)
# 移除 tunnel 模式字段
if column_exists("proxy_nodes", "tunnel_connected_at"):
op.drop_column("proxy_nodes", "tunnel_connected_at")
if column_exists("proxy_nodes", "tunnel_connected"):
op.drop_column("proxy_nodes", "tunnel_connected")
if column_exists("proxy_nodes", "tunnel_mode"):
op.drop_column("proxy_nodes", "tunnel_mode")

Some files were not shown because too many files have changed in this diff Show More