Compare 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
2499 changed files with 38540 additions and 555968 deletions
+16 -4
View File
@@ -1,12 +1,24 @@
# Build artifacts
build/
# 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/.vite/
# frontend/dist/ - 注释掉,因为我们需要预构建的dist文件
# Development
.git/
@@ -48,4 +60,4 @@ Dockerfile.*
# Deployment
deploy/
scripts/
scripts/
+82 -30
View File
@@ -1,35 +1,17 @@
# ==================== 必须配置(启动前) ====================
# 以下配置项必须在项目启动前设置
# 应用端口(默认 8084)
APP_PORT=8084
# 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密钥(使用 python generate_keys.py 生成)
# 用于用户登录 token 签名,更换后所有用户需重新登录
@@ -39,21 +21,91 @@ JWT_SECRET_KEY=change-this-to-a-secure-random-string
# 注意:更换此密钥后需要在管理面板重新配置所有 Provider API Key
ENCRYPTION_KEY=change-this-to-another-secure-random-string
# 启动自举管理员(仅在当前库里还没有活动管理员时生效)
# 支付回调共享密钥(公开 /api/payment/callback/* 入口必须携带 x-payment-callback-token)
# 建议使用 32+ 位随机字符串
PAYMENT_CALLBACK_SECRET=change-this-to-a-secure-callback-secret
# 管理员账号(仅首次初始化时使用, 创建完成后可在系统内修改密码)
ADMIN_EMAIL=[email protected]
ADMIN_USERNAME=admin
ADMIN_PASSWORD=admin123456
# ==================== 可选配置(有默认值) ====================
# 以下配置项有合理的默认值,可按需调整
# 支付回调共享密钥(公开 /api/payment/callback/* 入口必须携带 x-payment-callback-token)
# 建议使用 32+ 位随机字符串
# PAYMENT_CALLBACK_SECRET=change-this-to-a-secure-callback-secret
# 应用端口(默认 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 连接池配置(默认适合单实例/小型部署;高并发可按需调大)
# AETHER_GATEWAY_DATA_POSTGRES_MIN_CONNECTIONS=1
# AETHER_GATEWAY_DATA_POSTGRES_MAX_CONNECTIONS=20
# AETHER_GATEWAY_DATA_POSTGRES_IDLE_TIMEOUT_MS=30000
# Gunicorn Worker 数量(默认 2)
# Tunnel 请求统一经 Hub 转发,可安全使用多 worker。
# 非 Docker 运行时若使用 ProxyNode tunnel,请确保 aether-hub 可达(默认 ws://127.0.0.1:8085)。
# GUNICORN_WORKERS=2
# 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
+81 -39
View File
@@ -7,6 +7,12 @@ on:
permissions:
contents: write
packages: write
env:
REGISTRY: ghcr.io
GHCR_IMAGE: fawney19/aether-proxy
DOCKERHUB_IMAGE: fawney19/aether-proxy
jobs:
build:
@@ -24,21 +30,13 @@ jobs:
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
os: macos-latest
use_cross: false
- name: macos-arm64
target: aarch64-apple-darwin
os: macos-15
os: macos-latest
use_cross: false
- name: windows-amd64
target: x86_64-pc-windows-msvc
@@ -53,13 +51,10 @@ jobs:
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-proxy -> target
workspaces: aether-proxy -> target
key: ${{ matrix.target }}
- name: Install cross
@@ -67,7 +62,7 @@ jobs:
uses: taiki-e/install-action@cross
- name: Build
working-directory: apps/aether-proxy
working-directory: aether-proxy
shell: bash
run: |
if [ "${{ matrix.use_cross }}" = "true" ]; then
@@ -80,16 +75,16 @@ jobs:
if: runner.os != 'Windows'
shell: bash
run: |
cd target/${{ matrix.target }}/release
cd aether-proxy/target/${{ matrix.target }}/release
chmod +x aether-proxy
tar czf ../../../aether-proxy-${{ matrix.name }}.tar.gz aether-proxy
tar czf ../../../../aether-proxy-${{ matrix.name }}.tar.gz aether-proxy
- name: Package (Windows)
if: runner.os == 'Windows'
shell: bash
run: |
cd target/${{ matrix.target }}/release
7z a ../../../aether-proxy-${{ matrix.name }}.zip aether-proxy.exe
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
@@ -125,6 +120,68 @@ jobs:
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
@@ -132,7 +189,7 @@ jobs:
steps:
- uses: actions/checkout@v5
with:
ref: aether-rust-pioneer
ref: master
- name: Update README download links
env:
@@ -140,20 +197,11 @@ jobs:
run: |
VERSION="${TAG#proxy-v}"
BASE="https://github.com/fawney19/Aether/releases/download/${TAG}"
if [ -d apps/aether-proxy ]; then
PROXY_DIR="apps/aether-proxy"
else
PROXY_DIR="aether-proxy"
fi
cd "$PROXY_DIR"
cd aether-proxy
TABLE="| Platform | Download |\n|----------|----------|\n"
TABLE+="| Linux x86_64 (GNU) | [aether-proxy-linux-amd64.tar.gz](${BASE}/aether-proxy-linux-amd64.tar.gz) |\n"
TABLE+="| Linux ARM64 (GNU) | [aether-proxy-linux-arm64.tar.gz](${BASE}/aether-proxy-linux-arm64.tar.gz) |\n"
TABLE+="| Linux x86_64 (musl) | [aether-proxy-linux-musl-amd64.tar.gz](${BASE}/aether-proxy-linux-musl-amd64.tar.gz) |\n"
TABLE+="| Linux ARM64 (musl) | [aether-proxy-linux-musl-arm64.tar.gz](${BASE}/aether-proxy-linux-musl-arm64.tar.gz) |\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) |"
@@ -169,13 +217,7 @@ jobs:
- name: Commit and push
run: |
if [ -d apps/aether-proxy ]; then
PROXY_DIR="apps/aether-proxy"
else
PROXY_DIR="aether-proxy"
fi
cd "$PROXY_DIR"
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
+270 -94
View File
@@ -4,114 +4,212 @@ on:
push:
tags: ['v*']
workflow_dispatch:
permissions:
contents: read
packages: write
inputs:
build_base:
description: 'Rebuild base image'
required: false
default: false
type: boolean
env:
REGISTRY: ghcr.io
GHCR_IMAGE: fawney19/aether
DOCKERHUB_IMAGE: fawney19/aether
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:
frontend:
name: Build frontend
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: Setup Node.js
uses: actions/setup-node@v4
- name: Log in to Container Registry
uses: docker/login-action@v3
with:
node-version: 22
cache: npm
cache-dependency-path: frontend/package-lock.json
registry: ${{ env.REGISTRY }}
username: ${{ github.actor }}
password: ${{ secrets.GITHUB_TOKEN }}
- name: Install & build
working-directory: frontend
- name: Check if base image needs rebuild
id: check
run: |
npm ci
npm run build
if [ "${{ github.event.inputs.build_base }}" == "true" ]; then
echo "base_changed=true" >> $GITHUB_OUTPUT
exit 0
fi
- name: Upload frontend artifact
uses: actions/upload-artifact@v5
with:
name: frontend-dist
path: frontend/dist/
if-no-files-found: error
retention-days: 1
# 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"
build:
name: Build ${{ matrix.name }}
# 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
strategy:
fail-fast: true
matrix:
include:
- name: linux-amd64
target: x86_64-unknown-linux-musl
arch: amd64
- name: linux-arm64
target: aarch64-unknown-linux-musl
arch: arm64
permissions:
contents: read
packages: write
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
uses: taiki-e/install-action@cross
- name: Build
env:
CARGO_TERM_COLOR: always
run: cross build --release --locked -p aether-gateway --target ${{ matrix.target }}
- name: Upload binary artifact
uses: actions/upload-artifact@v5
with:
name: aether-gateway-${{ matrix.arch }}
path: target/${{ matrix.target }}/release/aether-gateway
if-no-files-found: error
retention-days: 1
docker:
name: Docker multi-arch
needs: [frontend, build]
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-amd64/aether-gateway dist/aether-gateway-amd64
cp artifacts/aether-gateway-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
- 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 }}
@@ -124,13 +222,13 @@ jobs:
username: ${{ secrets.DOCKERHUB_USERNAME }}
password: ${{ secrets.DOCKERHUB_TOKEN }}
- name: Extract metadata
- name: Extract metadata for app image
id: meta
uses: docker/metadata-action@v5
with:
images: |
${{ env.REGISTRY }}/${{ env.GHCR_IMAGE }}
docker.io/${{ env.DOCKERHUB_IMAGE }}
${{ env.REGISTRY }}/${{ env.APP_IMAGE_NAME }}
docker.io/fawney19/aether
tags: |
type=semver,pattern={{version}}
type=semver,pattern={{major}}.{{minor}}
@@ -140,12 +238,90 @@ jobs:
flavor: |
latest=auto
- name: Build and push
- 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
push: true
tags: ${{ steps.meta.outputs.tags }}
labels: ${{ steps.meta.outputs.labels }}
platforms: linux/amd64,linux/arm64
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
-136
View File
@@ -1,136 +0,0 @@
name: Rust CI
on:
push:
branches:
- master
- aether-rust-pioneer
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_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
- name: Show Rust toolchain
run: rustup show active-toolchain
- name: Install rustfmt component
run: rustup component add rustfmt
- name: Format
run: cargo fmt --all --check
clippy:
name: Clippy
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: Install clippy component
run: rustup component add 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 --all-targets -- -D warnings
- 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
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: Test
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
run: cargo test --workspace
- name: Show sccache stats
if: always()
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
run: sccache --show-stats
check:
name: check
runs-on: ubuntu-latest
needs:
- fmt
- clippy
- test
if: ${{ always() }}
steps:
- name: Verify required jobs
run: |
if [ "${{ needs.fmt.result }}" != "success" ] || \
[ "${{ needs.clippy.result }}" != "success" ] || \
[ "${{ needs.test.result }}" != "success" ]; then
echo "Rust CI failed"
exit 1
fi
+2 -5
View File
@@ -204,6 +204,7 @@ logs/
# Git backup
.git.backup/
.worktrees/
# Database backups
backups/
@@ -213,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
@@ -242,4 +239,4 @@ src/_version.py
# Analysis folder (third-party code for reference)
analysis/
new-api/
apps/aether-proxy/aether-proxy.toml
/aether-proxy/target/
+1
View File
@@ -0,0 +1 @@
3.13
Generated
-5356
View File
File diff suppressed because it is too large Load Diff
-95
View File
@@ -1,95 +0,0 @@
[workspace]
members = [
"apps/aether-proxy",
"crates/aether-admin",
"crates/aether-ai-pipeline",
"crates/aether-data-contracts",
"crates/aether-cache",
"crates/aether-billing",
"crates/aether-wallet",
"crates/aether-crypto",
"crates/aether-contracts",
"crates/aether-data",
"crates/aether-model-fetch",
"crates/aether-provider-transport",
"crates/aether-scheduler-core",
"crates/aether-usage-runtime",
"crates/aether-video-tasks-core",
"apps/aether-gateway",
"crates/aether-http",
"crates/aether-runtime",
"crates/aether-testkit",
]
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-ai-pipeline = { path = "crates/aether-ai-pipeline" }
aether-data-contracts = { path = "crates/aether-data-contracts" }
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" }
aether-model-fetch = { path = "crates/aether-model-fetch" }
aether-provider-transport = { path = "crates/aether-provider-transport" }
aether-scheduler-core = { path = "crates/aether-scheduler-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" }
aether-testkit = { path = "crates/aether-testkit" }
aes = "0.8"
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"
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"] }
regex = "1"
rustls = { version = "0.23", features = ["ring"] }
serde = { version = "1", features = ["derive"] }
serde_json = { version = "1", features = ["preserve_order"] }
serde_path_to_error = "0.1"
sha2 = "0.10"
sqlx = { version = "0.8", default-features = false, features = ["postgres", "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"
url = "2"
[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
+284 -24
View File
@@ -1,31 +1,291 @@
# syntax=docker/dockerfile:1
# Aether Gateway 运行时镜像(交叉编译方案)
# 二进制和前端产物均由 CI 预先构建,此 Dockerfile 仅做打包
# 用法: docker buildx build --platform linux/amd64,linux/arm64 -f Dockerfile.app .
#
# 构建上下文中须包含:
# dist/aether-gateway-amd64 (x86_64-unknown-linux-musl 交叉编译产物)
# dist/aether-gateway-arm64 (aarch64-unknown-linux-musl 交叉编译产物)
# dist/frontend/ (npm run build 产物)
# 运行镜像:从 base 提取产物到精简运行时
# 构建命令: docker build -f Dockerfile.app -t aether-app:latest .
# 用于 GitHub Actions CI(官方源)
FROM gcr.io/distroless/static-debian12
# TARGETARCH 由 buildx 自动注入: amd64 或 arm64
ARG TARGETARCH
COPY dist/aether-gateway-${TARGETARCH} /usr/local/bin/aether-gateway
COPY dist/frontend/ /srv/frontend
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
ENV RUST_LOG=aether_gateway=info \
APP_PORT=8084 \
AETHER_GATEWAY_STATIC_DIR=/srv/frontend
EXPOSE 8084
ARG HUB_RELEASE_REPO=fawney19/Aether
ARG HUB_TAG
ARG TARGETARCH
ARG GITHUB_TOKEN
# 运行时依赖(无 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 ["/usr/local/bin/aether-gateway", "--healthcheck"]
USER root
ENTRYPOINT ["/usr/local/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"]
+299 -110
View File
@@ -1,131 +1,320 @@
# syntax=docker/dockerfile:1
# Aether 运行镜像:Rust gateway 直接服务 API + 前端静态文件(国内镜像源版本)
# 运行镜像:从 base 提取产物到精简运行时(国内镜像源版本)
# 构建命令: docker build -f Dockerfile.app.local -t aether-app:latest .
# 用于本地/国内服务器部署
# ==================== 前端构建 ====================
FROM node:22-slim AS frontend-builder
WORKDIR /app/frontend
COPY frontend/package*.json ./
RUN npm config set registry https://registry.npmmirror.com && npm ci
COPY frontend/ ./
RUN npm run build
FROM aether-base:latest AS builder
# ==================== Rust gateway 构建 ====================
FROM rust:1.94.1-slim AS gateway-base
WORKDIR /build
WORKDIR /app
# 本地镜像优先缩短构建时间,保留 release 语义,但改用更快的 thin LTO。
ENV CARGO_REGISTRIES_CRATES_IO_PROTOCOL=sparse \
CARGO_PROFILE_RELEASE_LTO=thin \
CARGO_PROFILE_RELEASE_CODEGEN_UNITS=16
# 复制前端源码并构建
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
# 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 \
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
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 --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 \
cargo build --release --locked -p aether-gateway && \
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; \
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; \
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; \
if [ -f /etc/nsswitch.conf ]; then \
cp /etc/nsswitch.conf /runtime-root/etc/nsswitch.conf; \
fi
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; \
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; \
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_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"]
+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
+50 -67
View File
@@ -43,18 +43,12 @@ cd Aether
# 2. 配置环境变量
cp .env.example .env
./generate_keys.sh # 生成密钥, 并将生成的密钥填入 .env
python generate_keys.py # 生成密钥, 并将生成的密钥填入 .env
# 3. 首次部署 / 更新
# 3. 部署 / 更新(自动执行数据库迁移)
docker compose pull && docker compose up -d
# 4. 默认会在 app 启动前自动执行挂起的 migration / backfill
# 如需手工控制,可在 .env 中设 AETHER_GATEWAY_AUTO_PREPARE_DATABASE=false
# 然后按需执行:
docker compose run --rm app --migrate
docker compose run --rm app --apply-backfills
# 5. 升级前备份 (可选)
# 4. 升级前备份 (可选)
docker compose exec postgres pg_dump -U postgres aether | gzip > backup_$(date +%Y%m%d_%H%M%S).sql.gz
```
@@ -67,9 +61,9 @@ cd Aether
# 2. 配置环境变量
cp .env.example .env
./generate_keys.sh # 生成密钥, 并将生成的密钥填入 .env
python generate_keys.py # 生成密钥, 并将生成的密钥填入 .env
# 3. 部署 / 更新(自动构建并启动)
# 3. 部署 / 更新(自动构建、启动、迁移)
git pull
./deploy.sh
```
@@ -80,72 +74,46 @@ git pull
# 启动依赖
docker compose -f docker-compose.build.yml up -d postgres redis
# 数据库迁移(仅在已有数据库引入新 migration 时需要)
./dev.sh --migrate
# 数据回填
./dev.sh --apply-backfills
# 后端
uv sync
./dev.sh
# 前端
cd frontend && npm install && npm run dev
```
`./dev.sh` 现在只保留一种本地模式:
| 角色 | 本地地址 | 说明 |
|------|----------|------|
| Rust frontdoor | 默认 `http://localhost:8084` | `aether-gateway`,本地唯一公开入口;实际端口由 `APP_PORT` 控制 |
本地默认链路是:
```text
client -> rust frontdoor (aether-gateway) -> execution_runtime/provider transport
```
其中:
- `aether-gateway` 负责公开入口、健康检查、格式转换、本地执行 runtime,以及当前已迁到 Rust 的 frontdoor/control/background 路径。
- `./dev.sh` 不再启动 Python 宿主;未下沉到 Rust 的 legacy 路由会直接失败。
- `./dev.sh --migrate` 会复用 `.env` 里的数据库配置,显式执行一次数据库迁移后退出。
- `./dev.sh` 默认把 `AETHER_GATEWAY_VIDEO_TASK_TRUTH_SOURCE_MODE` 设为 `rust-authoritative`,避免本地还依赖 Python sync report 语义。
- 空库首次启动会自动初始化到当前 baseline。
- `aether-gateway` 默认启动不会自动应用后续 schema migration;如果数据库版本落后,服务会拒绝启动,并提示先执行 `aether-gateway --migrate`。
- 仓库自带的 `docker-compose.yml` 和 `docker-compose.build.yml` 都已把 `AETHER_GATEWAY_AUTO_PREPARE_DATABASE` 设为默认开启,因此无论是预构建镜像部署还是 `./deploy.sh` / 本地构建 compose,常规启动都会在监听端口前自动执行挂起的 migration 和 backfill。
## Aether Proxy (可选)
Aether Proxy 是配套的正向代理节点,部署在海外 VPS 上,为墙内的 Aether 实例中转 API 流量。或者部署在其他服务器为指定的提供商、账号、Key使用不同的节点访问。支持 TUI 向导一键配置、systemd 服务管理、TLS 加密、DNS 缓存及连接池调优。
- Docker Compose 部署或下载预编译二进制直接运行
- 通过 `aether-proxy setup` 完成交互式配置,自动注册为系统服务
- 详细文档见 [apps/aether-proxy/README.md](apps/aether-proxy/README.md)
- 详细文档见 [aether-proxy/README.md](aether-proxy/README.md)
## 环境变量
部署建议直接参考对应示例文件:
### 必需配置
- Docker Compose:根目录 [`.env.example`](.env.example)
- systemd 二进制部署:[deploy/systemd/aether-gateway.env.example](deploy/systemd/aether-gateway.env.example)
| 变量 | 说明 |
|------|------|
| `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`:`aether-gateway` 唯一监听端口,固定绑定 `0.0.0.0:${APP_PORT}`
- `DATABASE_URL` / `REDIS_URL`:`aether-gateway` 直接读取的共享后端连接串
- `AETHER_GATEWAY_AUTO_PREPARE_DATABASE`:常规启动前自动执行挂起的 schema migration 和 backfill;仓库自带的 `docker-compose.yml` 和 `docker-compose.build.yml` 默认开启
- `JWT_SECRET_KEY` / `ENCRYPTION_KEY`:认证和敏感数据加密所需密钥
- `API_KEY_PREFIX`:用户和管理员新建 API Key 时使用的前缀,默认 `sk`
- `PAYMENT_CALLBACK_SECRET`:支付回调公开入口的共享密钥;未配置时相关路由保持禁用
- `ADMIN_USERNAME` / `ADMIN_PASSWORD` / `ADMIN_EMAIL`:首次启动时自举首个本地管理员
- `CORS_ORIGINS` / `CORS_ALLOW_CREDENTIALS`:前端跨域来源控制;如果要跨域带登录 Cookie,`CORS_ORIGINS` 不能写 `*`
- `AETHER_GATEWAY_DEPLOYMENT_TOPOLOGY=single-node|multi-node`
- `AETHER_GATEWAY_NODE_ROLE=all|frontdoor|background`
- `RUST_LOG`:Rust 日志过滤,例如 `aether_gateway=info`、`aether_gateway=debug,sqlx=warn`
- 如果使用仓库内置的数据栈 compose,再额外配置 `DB_PASSWORD` / `REDIS_PASSWORD`
systemd 的 `.env` 必须保持简单 `KEY=VALUE` 形式,不要写 `export`、`${VAR}` 或命令替换。
| 变量 | 默认值 | 说明 |
|------|--------|------|
| `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
@@ -164,22 +132,37 @@ systemd 的 `.env` 必须保持简单 `KEY=VALUE` 形式,不要写 `export`、
**有备份的情况(推荐):**
```bash
# Docker Compose:
# 1. 切回旧镜像 tag / digest
# 2. 恢复 Postgres 备份
# 3. 再启动 app
# 1. 停止应用
docker compose stop app
# systemd:
# 1. 把 /opt/aether/current 切回旧 release
# 2. systemctl restart aether-gateway
# 3. 如果升级包含数据库结构变更,再恢复 Postgres 备份
# 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,方便回滚时使用。
**没有备份的情况:**
当前不应该再依赖旧的 `alembic downgrade` 路线。空库首次启动会自动初始化;如果开启 `AETHER_GATEWAY_AUTO_PREPARE_DATABASE=true`(仓库自带的两份 compose 默认都如此),常规服务启动也会自动应用挂起的 migration/backfill。无论是否自动执行,只要本次发布带来了不可逆的数据结构变化,没有备份就不能保证安全回滚。因此升级前强烈建议先备份 `Postgres`。
```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,如果迁移涉及不可逆的数据变更(如删除列),可能无法完全恢复数据。因此强烈建议升级前备份。
---
-24
View File
@@ -1,24 +0,0 @@
# file generated by vcs-versioning
# don't change, don't track in version control
from __future__ import annotations
__all__ = [
"__version__",
"__version_tuple__",
"version",
"version_tuple",
"__commit_id__",
"commit_id",
]
version: str
__version__: str
__version_tuple__: tuple[int | str, ...]
version_tuple: tuple[int | str, ...]
commit_id: str | None
__commit_id__: str | None
__version__ = version = '0.6.4.dev6+gddf18fed9.d20260331'
__version_tuple__ = version_tuple = (0, 6, 4, 'dev6', 'gddf18fed9.d20260331')
__commit_id__ = commit_id = None
-334
View File
@@ -1,334 +0,0 @@
"""Admin API routers.
The admin surface remains Python-only host/control-plane scope. It is not part
of the compatibility frontdoor manifest that Rust is preparing to absorb.
"""
from __future__ import annotations
from fastapi import APIRouter
from starlette.routing import BaseRoute
from .adaptive import router as adaptive_router
from .api_keys import router as api_keys_router
from .billing import router as billing_router
from .endpoints import router as endpoints_router
from .models import router as models_router
from .modules import router as modules_router
from .monitoring import router as monitoring_router
from .payments import router as payments_router
from .pool import router as pool_router
from .provider_oauth import router as provider_oauth_router
from .provider_ops import router as provider_ops_router
from .provider_query import router as provider_query_router
from .provider_strategy import router as provider_strategy_router
from .providers import router as providers_router
from .security import router as security_router
from .stats import router as stats_router
from .system import router as system_router
from .usage import router as usage_router
from .users import router as users_router
from .video_tasks import router as video_tasks_router
from .wallets import router as wallets_router
_RUST_OWNED_ADMIN_ROUTE_SIGNATURES = frozenset(
{
("GET", "/api/admin/modules/status"),
("GET", "/api/admin/modules/status/{module_name}"),
("PUT", "/api/admin/modules/status/{module_name}/enabled"),
("GET", "/api/admin/system/version"),
("GET", "/api/admin/system/check-update"),
("GET", "/api/admin/system/aws-regions"),
("GET", "/api/admin/system/stats"),
("GET", "/api/admin/system/settings"),
("GET", "/api/admin/system/config/export"),
("GET", "/api/admin/system/users/export"),
("POST", "/api/admin/system/config/import"),
("POST", "/api/admin/system/users/import"),
("POST", "/api/admin/system/smtp/test"),
("POST", "/api/admin/system/cleanup"),
("POST", "/api/admin/system/purge/config"),
("POST", "/api/admin/system/purge/users"),
("POST", "/api/admin/system/purge/usage"),
("POST", "/api/admin/system/purge/audit-logs"),
("POST", "/api/admin/system/purge/request-bodies"),
("POST", "/api/admin/system/purge/stats"),
("PUT", "/api/admin/system/settings"),
("GET", "/api/admin/system/configs"),
("GET", "/api/admin/system/configs/{key}"),
("PUT", "/api/admin/system/configs/{key}"),
("DELETE", "/api/admin/system/configs/{key}"),
("GET", "/api/admin/system/api-formats"),
("GET", "/api/admin/system/email/templates"),
("GET", "/api/admin/system/email/templates/{template_type}"),
("PUT", "/api/admin/system/email/templates/{template_type}"),
("POST", "/api/admin/system/email/templates/{template_type}/preview"),
("POST", "/api/admin/system/email/templates/{template_type}/reset"),
("GET", "/api/admin/providers/"),
("POST", "/api/admin/providers/"),
("PATCH", "/api/admin/providers/{provider_id}"),
("DELETE", "/api/admin/providers/{provider_id}"),
("GET", "/api/admin/providers/summary"),
("GET", "/api/admin/providers/{provider_id}/summary"),
("GET", "/api/admin/providers/{provider_id}/health-monitor"),
("GET", "/api/admin/providers/{provider_id}/mapping-preview"),
("GET", "/api/admin/providers/{provider_id}/delete-task/{task_id}"),
("GET", "/api/admin/providers/{provider_id}/pool-status"),
("POST", "/api/admin/providers/{provider_id}/pool/clear-cooldown/{key_id}"),
("POST", "/api/admin/providers/{provider_id}/pool/reset-cost/{key_id}"),
("GET", "/api/admin/providers/{provider_id}/models"),
("POST", "/api/admin/providers/{provider_id}/models"),
("GET", "/api/admin/providers/{provider_id}/models/{model_id}"),
("PATCH", "/api/admin/providers/{provider_id}/models/{model_id}"),
("DELETE", "/api/admin/providers/{provider_id}/models/{model_id}"),
("POST", "/api/admin/providers/{provider_id}/models/batch"),
("GET", "/api/admin/providers/{provider_id}/available-source-models"),
("POST", "/api/admin/providers/{provider_id}/assign-global-models"),
("POST", "/api/admin/providers/{provider_id}/import-from-upstream"),
("GET", "/api/admin/endpoints/providers/{provider_id}/endpoints"),
("POST", "/api/admin/endpoints/providers/{provider_id}/endpoints"),
("GET", "/api/admin/endpoints/defaults/{api_format}/body-rules"),
("GET", "/api/admin/endpoints/{endpoint_id}"),
("PUT", "/api/admin/endpoints/{endpoint_id}"),
("DELETE", "/api/admin/endpoints/{endpoint_id}"),
("PUT", "/api/admin/endpoints/keys/{key_id}"),
("GET", "/api/admin/endpoints/keys/grouped-by-format"),
("GET", "/api/admin/endpoints/keys/{key_id}/reveal"),
("GET", "/api/admin/endpoints/keys/{key_id}/export"),
("DELETE", "/api/admin/endpoints/keys/{key_id}"),
("POST", "/api/admin/endpoints/keys/batch-delete"),
("POST", "/api/admin/endpoints/keys/{key_id}/clear-oauth-invalid"),
("GET", "/api/admin/endpoints/providers/{provider_id}/keys"),
("POST", "/api/admin/endpoints/providers/{provider_id}/keys"),
("POST", "/api/admin/endpoints/providers/{provider_id}/refresh-quota"),
("GET", "/api/admin/endpoints/rpm/key/{key_id}"),
("DELETE", "/api/admin/endpoints/rpm/key/{key_id}"),
("GET", "/api/admin/endpoints/health/summary"),
("GET", "/api/admin/endpoints/health/status"),
("GET", "/api/admin/endpoints/health/api-formats"),
("GET", "/api/admin/endpoints/health/key/{key_id}"),
("PATCH", "/api/admin/endpoints/health/keys/{key_id}"),
("PATCH", "/api/admin/endpoints/health/keys"),
("GET", "/api/admin/provider-oauth/supported-types"),
("POST", "/api/admin/provider-oauth/keys/{key_id}/start"),
("POST", "/api/admin/provider-oauth/keys/{key_id}/complete"),
("POST", "/api/admin/provider-oauth/keys/{key_id}/refresh"),
("POST", "/api/admin/provider-oauth/providers/{provider_id}/start"),
("POST", "/api/admin/provider-oauth/providers/{provider_id}/complete"),
("POST", "/api/admin/provider-oauth/providers/{provider_id}/import-refresh-token"),
("POST", "/api/admin/provider-oauth/providers/{provider_id}/device-authorize"),
("POST", "/api/admin/provider-oauth/providers/{provider_id}/device-poll"),
("POST", "/api/admin/provider-oauth/providers/{provider_id}/batch-import"),
("POST", "/api/admin/provider-oauth/providers/{provider_id}/batch-import/tasks"),
("GET", "/api/admin/provider-oauth/providers/{provider_id}/batch-import/tasks/{task_id}"),
("GET", "/api/admin/adaptive/keys"),
("PATCH", "/api/admin/adaptive/keys/{key_id}/mode"),
("GET", "/api/admin/adaptive/keys/{key_id}/stats"),
("DELETE", "/api/admin/adaptive/keys/{key_id}/learning"),
("PATCH", "/api/admin/adaptive/keys/{key_id}/limit"),
("GET", "/api/admin/adaptive/summary"),
("GET", "/api/admin/provider-ops/architectures"),
("GET", "/api/admin/provider-ops/architectures/{architecture_id}"),
("GET", "/api/admin/provider-ops/providers/{provider_id}/status"),
("GET", "/api/admin/provider-ops/providers/{provider_id}/config"),
("PUT", "/api/admin/provider-ops/providers/{provider_id}/config"),
("DELETE", "/api/admin/provider-ops/providers/{provider_id}/config"),
("POST", "/api/admin/provider-ops/providers/{provider_id}/connect"),
("POST", "/api/admin/provider-ops/providers/{provider_id}/disconnect"),
("POST", "/api/admin/provider-ops/providers/{provider_id}/verify"),
("POST", "/api/admin/provider-ops/providers/{provider_id}/actions/{action_type}"),
("GET", "/api/admin/provider-ops/providers/{provider_id}/balance"),
("POST", "/api/admin/provider-ops/providers/{provider_id}/balance"),
("POST", "/api/admin/provider-ops/providers/{provider_id}/checkin"),
("POST", "/api/admin/provider-ops/batch/balance"),
("GET", "/api/admin/billing/presets"),
("POST", "/api/admin/billing/presets/apply"),
("GET", "/api/admin/billing/rules"),
("GET", "/api/admin/billing/rules/{rule_id}"),
("POST", "/api/admin/billing/rules"),
("PUT", "/api/admin/billing/rules/{rule_id}"),
("GET", "/api/admin/billing/collectors"),
("GET", "/api/admin/billing/collectors/{collector_id}"),
("POST", "/api/admin/billing/collectors"),
("PUT", "/api/admin/billing/collectors/{collector_id}"),
("PUT", "/api/admin/provider-strategy/providers/{provider_id}/billing"),
("GET", "/api/admin/provider-strategy/providers/{provider_id}/stats"),
("GET", "/api/admin/provider-strategy/strategies"),
("DELETE", "/api/admin/provider-strategy/providers/{provider_id}/quota"),
("POST", "/api/admin/provider-query/models"),
("POST", "/api/admin/provider-query/test-model"),
("POST", "/api/admin/provider-query/test-model-failover"),
("GET", "/api/admin/payments/orders"),
("GET", "/api/admin/payments/orders/{order_id}"),
("POST", "/api/admin/payments/orders/{order_id}/expire"),
("POST", "/api/admin/payments/orders/{order_id}/credit"),
("POST", "/api/admin/payments/orders/{order_id}/fail"),
("GET", "/api/admin/payments/callbacks"),
("POST", "/api/admin/security/ip/blacklist"),
("DELETE", "/api/admin/security/ip/blacklist/{ip_address}"),
("GET", "/api/admin/security/ip/blacklist/stats"),
("POST", "/api/admin/security/ip/whitelist"),
("DELETE", "/api/admin/security/ip/whitelist/{ip_address}"),
("GET", "/api/admin/security/ip/whitelist"),
("GET", "/api/admin/stats/providers/quota-usage"),
("GET", "/api/admin/stats/comparison"),
("GET", "/api/admin/stats/errors/distribution"),
("GET", "/api/admin/stats/performance/percentiles"),
("GET", "/api/admin/stats/cost/forecast"),
("GET", "/api/admin/stats/cost/savings"),
("GET", "/api/admin/stats/leaderboard/api-keys"),
("GET", "/api/admin/stats/leaderboard/models"),
("GET", "/api/admin/stats/leaderboard/users"),
("GET", "/api/admin/stats/time-series"),
("GET", "/api/admin/monitoring/audit-logs"),
("GET", "/api/admin/monitoring/system-status"),
("GET", "/api/admin/monitoring/suspicious-activities"),
("GET", "/api/admin/monitoring/user-behavior/{user_id}"),
("GET", "/api/admin/monitoring/resilience-status"),
("GET", "/api/admin/monitoring/resilience/circuit-history"),
("DELETE", "/api/admin/monitoring/resilience/error-stats"),
("GET", "/api/admin/monitoring/trace/{request_id}"),
("GET", "/api/admin/monitoring/trace/stats/provider/{provider_id}"),
("GET", "/api/admin/monitoring/cache/stats"),
("GET", "/api/admin/monitoring/cache/affinity/{user_identifier}"),
("GET", "/api/admin/monitoring/cache/affinities"),
("DELETE", "/api/admin/monitoring/cache/users/{user_identifier}"),
(
"DELETE",
"/api/admin/monitoring/cache/affinity/{affinity_key}/{endpoint_id}/{model_id}/{api_format}",
),
("DELETE", "/api/admin/monitoring/cache"),
("DELETE", "/api/admin/monitoring/cache/providers/{provider_id}"),
("GET", "/api/admin/monitoring/cache/config"),
("GET", "/api/admin/monitoring/cache/metrics"),
("GET", "/api/admin/monitoring/cache/model-mapping/stats"),
("DELETE", "/api/admin/monitoring/cache/model-mapping"),
("DELETE", "/api/admin/monitoring/cache/model-mapping/{model_name}"),
(
"DELETE",
"/api/admin/monitoring/cache/model-mapping/provider/{provider_id}/{global_model_id}",
),
("GET", "/api/admin/monitoring/cache/redis-keys"),
("DELETE", "/api/admin/monitoring/cache/redis-keys/{category}"),
("GET", "/api/admin/usage/aggregation/stats"),
("GET", "/api/admin/usage/stats"),
("GET", "/api/admin/usage/heatmap"),
("GET", "/api/admin/usage/records"),
("GET", "/api/admin/usage/active"),
("GET", "/api/admin/usage/cache-affinity/hit-analysis"),
("GET", "/api/admin/usage/cache-affinity/interval-timeline"),
("GET", "/api/admin/usage/cache-affinity/ttl-analysis"),
("GET", "/api/admin/usage/{usage_id}/curl"),
("GET", "/api/admin/usage/{usage_id}"),
("POST", "/api/admin/usage/{usage_id}/replay"),
("GET", "/api/admin/video-tasks"),
("GET", "/api/admin/video-tasks/stats"),
("GET", "/api/admin/video-tasks/{task_id}"),
("POST", "/api/admin/video-tasks/{task_id}/cancel"),
("GET", "/api/admin/video-tasks/{task_id}/video"),
("GET", "/api/admin/wallets"),
("GET", "/api/admin/wallets/ledger"),
("GET", "/api/admin/wallets/refund-requests"),
("GET", "/api/admin/wallets/{wallet_id}"),
("GET", "/api/admin/wallets/{wallet_id}/transactions"),
("GET", "/api/admin/wallets/{wallet_id}/refunds"),
("POST", "/api/admin/wallets/{wallet_id}/adjust"),
("POST", "/api/admin/wallets/{wallet_id}/recharge"),
("POST", "/api/admin/wallets/{wallet_id}/refunds/{refund_id}/process"),
("POST", "/api/admin/wallets/{wallet_id}/refunds/{refund_id}/complete"),
("POST", "/api/admin/wallets/{wallet_id}/refunds/{refund_id}/fail"),
("GET", "/api/admin/api-keys"),
("POST", "/api/admin/api-keys"),
("GET", "/api/admin/api-keys/{key_id}"),
("PUT", "/api/admin/api-keys/{key_id}"),
("PATCH", "/api/admin/api-keys/{key_id}"),
("DELETE", "/api/admin/api-keys/{key_id}"),
("GET", "/api/admin/users"),
("POST", "/api/admin/users"),
("GET", "/api/admin/users/{user_id}"),
("PUT", "/api/admin/users/{user_id}"),
("DELETE", "/api/admin/users/{user_id}"),
("GET", "/api/admin/users/{user_id}/sessions"),
("DELETE", "/api/admin/users/{user_id}/sessions"),
("DELETE", "/api/admin/users/{user_id}/sessions/{session_id}"),
("GET", "/api/admin/users/{user_id}/api-keys"),
("POST", "/api/admin/users/{user_id}/api-keys"),
("DELETE", "/api/admin/users/{user_id}/api-keys/{key_id}"),
("PUT", "/api/admin/users/{user_id}/api-keys/{key_id}"),
("PATCH", "/api/admin/users/{user_id}/api-keys/{key_id}/lock"),
("GET", "/api/admin/users/{user_id}/api-keys/{key_id}/full-key"),
("GET", "/api/admin/pool/overview"),
("GET", "/api/admin/pool/scheduling-presets"),
("GET", "/api/admin/pool/{provider_id}/keys"),
("GET", "/api/admin/pool/{provider_id}/keys/batch-delete-task/{task_id}"),
("POST", "/api/admin/pool/{provider_id}/keys/batch-action"),
("POST", "/api/admin/pool/{provider_id}/keys/batch-import"),
("POST", "/api/admin/pool/{provider_id}/keys/cleanup-banned"),
("POST", "/api/admin/pool/{provider_id}/keys/resolve-selection"),
("GET", "/api/admin/proxy-nodes"),
("GET", "/api/admin/models/catalog"),
("GET", "/api/admin/models/external"),
("DELETE", "/api/admin/models/external/cache"),
("GET", "/api/admin/models/global"),
("POST", "/api/admin/models/global"),
("GET", "/api/admin/models/global/{global_model_id}"),
("PATCH", "/api/admin/models/global/{global_model_id}"),
("DELETE", "/api/admin/models/global/{global_model_id}"),
("POST", "/api/admin/models/global/batch-delete"),
("POST", "/api/admin/models/global/{global_model_id}/assign-to-providers"),
("GET", "/api/admin/models/global/{global_model_id}/providers"),
("GET", "/api/admin/models/global/{global_model_id}/routing"),
}
)
def _route_is_rust_owned(route: BaseRoute) -> bool:
path = getattr(route, "path", None)
methods = getattr(route, "methods", None)
if not isinstance(path, str) or not methods:
return False
return any(
(method, path) in _RUST_OWNED_ADMIN_ROUTE_SIGNATURES
for method in methods
if method not in {"HEAD", "OPTIONS"}
)
def _build_python_admin_router() -> APIRouter:
"""Admin/control-plane routes that still require the Python host."""
admin_router = APIRouter()
admin_router.include_router(system_router)
admin_router.include_router(users_router)
admin_router.include_router(providers_router)
admin_router.include_router(api_keys_router)
admin_router.include_router(billing_router)
admin_router.include_router(usage_router)
admin_router.include_router(monitoring_router)
admin_router.include_router(payments_router)
admin_router.include_router(endpoints_router)
admin_router.include_router(provider_strategy_router)
admin_router.include_router(provider_oauth_router)
admin_router.include_router(adaptive_router)
admin_router.include_router(models_router)
admin_router.include_router(security_router)
admin_router.include_router(stats_router)
admin_router.include_router(provider_query_router)
admin_router.include_router(modules_router)
admin_router.include_router(pool_router)
admin_router.include_router(provider_ops_router)
admin_router.include_router(video_tasks_router)
admin_router.include_router(wallets_router)
admin_router.routes = [route for route in admin_router.routes if not _route_is_rust_owned(route)]
return admin_router
# Admin/control-plane 在本轮 frontdoor cutover 后仍保留在 Python 宿主。
python_admin_router = _build_python_admin_router()
router = python_admin_router
# 注意:以下路由已迁移到模块系统,由 ModuleRegistry 动态注册
# - ldap_router: 当 LDAP_AVAILABLE=true 时注册
# - management_tokens_router: 当 MANAGEMENT_TOKENS_AVAILABLE=true 时注册
# - proxy_nodes_router: 当 PROXY_NODES_AVAILABLE=true 时注册
__all__ = ["python_admin_router", "router"]
@@ -1,383 +0,0 @@
"""
Gemini Files 管理 API
提供文件映射管理与能力查询;上传入口已收成 Rust-only 兼容壳。
"""
from __future__ import annotations
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Any
from fastapi import APIRouter, Depends, File, HTTPException, Query, Request, UploadFile
from pydantic import BaseModel
from sqlalchemy import delete, func
from sqlalchemy.orm import Session, load_only
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext
from src.api.base.pipeline import get_pipeline
from src.database import get_db
from src.models.database import GeminiFileMapping, ProviderAPIKey, User
from src.services.gemini_files_mapping import delete_file_key_mapping
router = APIRouter(prefix="/api/admin/gemini-files", tags=["Gemini Files Management"])
pipeline = get_pipeline()
_RUST_UPLOADER_DETAIL = "Admin Gemini file upload requires Rust uploader"
class FileMappingResponse(BaseModel):
id: str
file_name: str
key_id: str
key_name: str | None = None
user_id: str | None = None
username: str | None = None
display_name: str | None = None
mime_type: str | None = None
created_at: datetime
expires_at: datetime
is_expired: bool
class FileMappingListResponse(BaseModel):
items: list[FileMappingResponse]
total: int
page: int
page_size: int
class FileMappingStatsResponse(BaseModel):
total_mappings: int
active_mappings: int
expired_mappings: int
by_mime_type: dict[str, int]
capable_keys_count: int
class CapableKeyResponse(BaseModel):
id: str
name: str
provider_name: str | None = None
class UploadResultItem(BaseModel):
key_id: str
key_name: str | None = None
success: bool
file_name: str | None = None
error: str | None = None
class UploadResponse(BaseModel):
display_name: str
mime_type: str
size_bytes: int
results: list[UploadResultItem]
success_count: int
fail_count: int
async def _list_file_mappings_response(
*,
db: Session,
page: int,
page_size: int,
include_expired: bool,
search: str | None,
) -> FileMappingListResponse:
now = datetime.now(timezone.utc)
query = db.query(GeminiFileMapping)
count_query = db.query(func.count(GeminiFileMapping.id))
if not include_expired:
active_filter = GeminiFileMapping.expires_at > now
query = query.filter(active_filter)
count_query = count_query.filter(active_filter)
if search:
search_pattern = f"%{search}%"
search_filter = (GeminiFileMapping.file_name.ilike(search_pattern)) | (
GeminiFileMapping.display_name.ilike(search_pattern)
)
query = query.filter(search_filter)
count_query = count_query.filter(search_filter)
total = int(count_query.scalar() or 0)
offset = (page - 1) * page_size
mappings = (
query.options(
load_only(
GeminiFileMapping.id,
GeminiFileMapping.file_name,
GeminiFileMapping.key_id,
GeminiFileMapping.user_id,
GeminiFileMapping.display_name,
GeminiFileMapping.mime_type,
GeminiFileMapping.created_at,
GeminiFileMapping.expires_at,
)
)
.order_by(GeminiFileMapping.created_at.desc())
.offset(offset)
.limit(page_size)
.all()
)
key_ids = {m.key_id for m in mappings}
user_ids = {m.user_id for m in mappings if m.user_id}
keys_map: dict[str, str | None] = {}
if key_ids:
keys = (
db.query(ProviderAPIKey)
.options(load_only(ProviderAPIKey.id, ProviderAPIKey.name))
.filter(ProviderAPIKey.id.in_(key_ids))
.all()
)
keys_map = {str(k.id): k.name for k in keys}
users_map: dict[str, str | None] = {}
if user_ids:
users = (
db.query(User)
.options(load_only(User.id, User.username))
.filter(User.id.in_(user_ids))
.all()
)
users_map = {str(u.id): u.username for u in users}
return FileMappingListResponse(
items=[
FileMappingResponse(
id=str(m.id),
file_name=m.file_name,
key_id=str(m.key_id),
key_name=keys_map.get(str(m.key_id)),
user_id=str(m.user_id) if m.user_id else None,
username=users_map.get(str(m.user_id)) if m.user_id else None,
display_name=m.display_name,
mime_type=m.mime_type,
created_at=m.created_at,
expires_at=m.expires_at,
is_expired=m.expires_at <= now,
)
for m in mappings
],
total=total,
page=page,
page_size=page_size,
)
async def _get_file_mapping_stats_response(*, db: Session) -> FileMappingStatsResponse:
now = datetime.now(timezone.utc)
total_mappings = db.query(func.count(GeminiFileMapping.id)).scalar() or 0
active_mappings = (
db.query(func.count(GeminiFileMapping.id))
.filter(GeminiFileMapping.expires_at > now)
.scalar()
or 0
)
expired_mappings = total_mappings - active_mappings
mime_stats = (
db.query(GeminiFileMapping.mime_type, func.count(GeminiFileMapping.id))
.filter(GeminiFileMapping.expires_at > now)
.group_by(GeminiFileMapping.mime_type)
.all()
)
by_mime_type = {(mime_type or "unknown"): count for mime_type, count in mime_stats}
keys = db.query(ProviderAPIKey.capabilities).filter(ProviderAPIKey.is_active.is_(True)).all()
capable_keys_count = sum(
1
for (capabilities,) in keys
if isinstance(capabilities, dict) and capabilities.get("gemini_files", False)
)
return FileMappingStatsResponse(
total_mappings=total_mappings,
active_mappings=active_mappings,
expired_mappings=expired_mappings,
by_mime_type=by_mime_type,
capable_keys_count=capable_keys_count,
)
async def _delete_mapping_response(*, db: Session, mapping_id: str) -> dict[str, Any]:
mapping = db.query(GeminiFileMapping).filter(GeminiFileMapping.id == mapping_id).first()
if not mapping:
raise HTTPException(status_code=404, detail="Mapping not found")
file_name = mapping.file_name
db.delete(mapping)
db.commit()
await delete_file_key_mapping(file_name)
return {"message": "Mapping deleted successfully", "file_name": file_name}
async def _cleanup_expired_mappings_response(*, db: Session) -> dict[str, Any]:
now = datetime.now(timezone.utc)
result = db.execute(delete(GeminiFileMapping).where(GeminiFileMapping.expires_at <= now))
db.commit()
deleted_count = result.rowcount
return {
"message": f"Cleaned up {deleted_count} expired mappings",
"deleted_count": deleted_count,
}
async def _list_capable_keys_response(*, db: Session) -> list[CapableKeyResponse]:
from src.models.database import Provider
key_rows = (
db.query(
ProviderAPIKey.id,
ProviderAPIKey.name,
ProviderAPIKey.provider_id,
ProviderAPIKey.capabilities,
)
.filter(ProviderAPIKey.is_active.is_(True))
.all()
)
capable_keys = [
key
for key in key_rows
if isinstance(key.capabilities, dict) and key.capabilities.get("gemini_files", False)
]
provider_ids = {key.provider_id for key in capable_keys if key.provider_id}
provider_map: dict[str, str] = {}
if provider_ids:
providers = db.query(Provider.id, Provider.name).filter(Provider.id.in_(provider_ids)).all()
provider_map = {str(provider_id): provider_name for provider_id, provider_name in providers}
return [
CapableKeyResponse(
id=str(key.id),
name=key.name,
provider_name=provider_map.get(str(key.provider_id)),
)
for key in capable_keys
]
async def _upload_file_response(*, file: UploadFile, key_ids: str) -> Any:
del file, key_ids
raise HTTPException(status_code=503, detail=_RUST_UPLOADER_DETAIL)
@dataclass
class AdminGeminiFilesListMappingsAdapter(AdminApiAdapter):
page: int
page_size: int
include_expired: bool
search: str | None
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
return await _list_file_mappings_response(
db=context.db,
page=self.page,
page_size=self.page_size,
include_expired=self.include_expired,
search=self.search,
)
class AdminGeminiFilesStatsAdapter(AdminApiAdapter):
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
return await _get_file_mapping_stats_response(db=context.db)
@dataclass
class AdminGeminiFilesDeleteMappingAdapter(AdminApiAdapter):
mapping_id: str
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
return await _delete_mapping_response(db=context.db, mapping_id=self.mapping_id)
class AdminGeminiFilesCleanupMappingsAdapter(AdminApiAdapter):
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
return await _cleanup_expired_mappings_response(db=context.db)
class AdminGeminiFilesCapableKeysAdapter(AdminApiAdapter):
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
return await _list_capable_keys_response(db=context.db)
@dataclass
class AdminGeminiFilesUploadAdapter(AdminApiAdapter):
file: UploadFile
key_ids: str
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
del context
return await _upload_file_response(file=self.file, key_ids=self.key_ids)
@router.get("/mappings", response_model=FileMappingListResponse)
async def list_file_mappings(
request: Request,
db: Session = Depends(get_db),
page: int = Query(1, ge=1),
page_size: int = Query(20, ge=1, le=100),
include_expired: bool = Query(False),
search: str | None = Query(None),
) -> Any:
adapter = AdminGeminiFilesListMappingsAdapter(
page=page,
page_size=page_size,
include_expired=include_expired,
search=search,
)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.get("/stats", response_model=FileMappingStatsResponse)
async def get_file_mapping_stats(
request: Request,
db: Session = Depends(get_db),
) -> Any:
adapter = AdminGeminiFilesStatsAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.delete("/mappings/{mapping_id}")
async def delete_mapping(
mapping_id: str,
request: Request,
db: Session = Depends(get_db),
) -> Any:
adapter = AdminGeminiFilesDeleteMappingAdapter(mapping_id=mapping_id)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.delete("/mappings")
async def cleanup_expired_mappings(
request: Request,
db: Session = Depends(get_db),
) -> Any:
adapter = AdminGeminiFilesCleanupMappingsAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.get("/capable-keys", response_model=list[CapableKeyResponse])
async def list_capable_keys(
request: Request,
db: Session = Depends(get_db),
) -> Any:
adapter = AdminGeminiFilesCapableKeysAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.post("/upload", response_model=UploadResponse)
async def upload_file(
request: Request,
file: UploadFile = File(...),
key_ids: str = Query(..., description="逗号分隔的 Key ID 列表"),
db: Session = Depends(get_db),
) -> Any:
adapter = AdminGeminiFilesUploadAdapter(file=file, key_ids=key_ids)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@@ -1,22 +0,0 @@
"""Stats admin routes export."""
from fastapi import APIRouter
from .comparison import router as comparison_router
from .cost import router as cost_router
from .errors import router as errors_router
from .leaderboard import router as leaderboard_router
from .performance import router as performance_router
from .quota import router as quota_router
from .time_series import router as time_series_router
router = APIRouter(prefix="/api/admin/stats", tags=["Admin - Stats"])
router.include_router(leaderboard_router)
router.include_router(time_series_router)
router.include_router(cost_router)
router.include_router(quota_router)
router.include_router(performance_router)
router.include_router(errors_router)
router.include_router(comparison_router)
__all__ = ["router"]
@@ -1,196 +0,0 @@
"""Shared helpers for admin stats routes."""
from __future__ import annotations
import hashlib
import json
from datetime import date, datetime, time, timedelta, timezone
from typing import Any, Literal
from fastapi import HTTPException
from sqlalchemy import and_, or_
from src.api.base.pipeline import get_pipeline
from src.config.settings import config
from src.models.database import Usage
from src.services.system.time_range import TimeRangeParams
pipeline = get_pipeline()
def _apply_admin_default_range(
params: TimeRangeParams | None,
) -> TimeRangeParams | None:
"""Apply a default range to avoid unbounded scans."""
if params is not None:
return params
days = int(getattr(config, "admin_usage_default_days", 0) or 0)
if days <= 0:
return None
today = datetime.now(timezone.utc).date()
start_date = today - timedelta(days=days - 1)
return TimeRangeParams(
start_date=start_date,
end_date=today,
timezone="UTC",
tz_offset_minutes=0,
).validate_and_resolve()
def _build_time_range_params(
start_date: date | None,
end_date: date | None,
preset: str | None,
timezone_name: str | None,
tz_offset_minutes: int | None,
) -> TimeRangeParams | None:
if not preset and start_date is None and end_date is None:
return None
try:
return TimeRangeParams(
start_date=start_date,
end_date=end_date,
preset=preset,
timezone=timezone_name,
tz_offset_minutes=tz_offset_minutes or 0,
).validate_and_resolve()
except Exception as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
def _hash_filters(filters: dict[str, Any]) -> str:
raw = json.dumps(filters, sort_keys=True, ensure_ascii=False, default=str)
return hashlib.sha1(raw.encode("utf-8")).hexdigest()[:16]
def _build_time_range_from_days(
days: int, timezone_name: str | None, tz_offset_minutes: int | None
) -> TimeRangeParams:
base = TimeRangeParams(
preset="today",
timezone=timezone_name,
tz_offset_minutes=tz_offset_minutes or 0,
).validate_and_resolve()
user_today = base.start_date
start_date = user_today - timedelta(days=days - 1)
return TimeRangeParams(
start_date=start_date,
end_date=user_today,
timezone=timezone_name,
tz_offset_minutes=tz_offset_minutes or 0,
).validate_and_resolve()
def _linear_regression(values: list[float]) -> tuple[float, float]:
n = len(values)
if n <= 1:
return 0.0, values[0] if values else 0.0
xs = list(range(n))
sum_x = sum(xs)
sum_y = sum(values)
sum_x2 = sum(x * x for x in xs)
sum_xy = sum(x * y for x, y in zip(xs, values))
denom = n * sum_x2 - sum_x * sum_x
if denom == 0:
return 0.0, values[-1]
slope = (n * sum_xy - sum_x * sum_y) / denom
intercept = (sum_y - slope * sum_x) / n
return slope, intercept
def _build_cache_key(
leaderboard_type: str,
metric: str,
time_range: TimeRangeParams | None,
filters: dict[str, Any],
) -> str:
start_value = time_range.start_date.isoformat() if time_range else "all"
end_value = time_range.end_date.isoformat() if time_range else "all"
tz_value = time_range.timezone if time_range else "utc"
offset_value = time_range.tz_offset_minutes if time_range else 0
return (
f"leaderboard:{leaderboard_type}:{metric}:{start_value}:{end_value}:"
f"{tz_value}:{offset_value}:{_hash_filters(filters)}"
)
def _is_today_range(time_range: TimeRangeParams | None) -> bool:
if not time_range:
return False
try:
user_today = time_range._get_user_today()
except Exception:
return False
return time_range.end_date == user_today
def _split_daily_and_usage_segments(
time_range: TimeRangeParams | None,
use_daily: bool,
) -> tuple[tuple[datetime, datetime] | None, list[tuple[datetime, datetime]] | None]:
if not time_range:
return None, None
start_utc, end_utc = time_range.to_utc_datetime_range()
if not use_daily:
return None, [(start_utc, end_utc)]
complete_dates, head_boundary, tail_boundary = time_range.get_complete_utc_dates()
daily_range = None
if complete_dates:
daily_start = datetime.combine(complete_dates[0], time.min, tzinfo=timezone.utc)
daily_end = datetime.combine(
complete_dates[-1] + timedelta(days=1), time.min, tzinfo=timezone.utc
)
daily_range = (daily_start, daily_end)
usage_segments: list[tuple[datetime, datetime]] = []
if head_boundary:
usage_segments.append(head_boundary)
if tail_boundary:
usage_segments.append(tail_boundary)
if not daily_range and not usage_segments:
usage_segments = [(start_utc, end_utc)]
return daily_range, usage_segments
def _apply_usage_time_segments(
query: Any, segments: list[tuple[datetime, datetime]] | None
) -> Any | None:
if segments is None:
return query
if not segments:
return None
conditions = []
for start_utc, end_utc in segments:
if start_utc >= end_utc:
continue
conditions.append(and_(Usage.created_at >= start_utc, Usage.created_at < end_utc))
if not conditions:
return None
return query.filter(or_(*conditions))
def _union_queries(queries: list[Any]) -> Any | None:
base = None
for query in queries:
if query is None:
continue
if base is None:
base = query
else:
base = base.union_all(query)
return base
def _metric_order(
metric: Literal["requests", "tokens", "cost"], order: Literal["asc", "desc"], expr: Any
) -> Any:
return expr.asc() if order == "asc" else expr.desc()
@@ -1,142 +0,0 @@
"""Admin comparison stats routes."""
from __future__ import annotations
from datetime import date, timedelta
from typing import Any, Literal
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from sqlalchemy.orm import Session
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext
from src.config.constants import CacheTTL
from src.database import get_db
from src.services.system.stats_aggregator import AggregatedStats, StatsFilter, query_stats_hybrid
from src.services.system.time_range import TimeRangeParams
from src.utils.cache_decorator import cache_result
from .common import pipeline
router = APIRouter()
class AdminComparisonAdapter(AdminApiAdapter):
def __init__(
self,
current_start: date,
current_end: date,
comparison_type: Literal["period", "year"],
timezone_name: str | None,
tz_offset_minutes: int | None,
) -> None:
self.current_start = current_start
self.current_end = current_end
self.comparison_type = comparison_type
self.timezone_name = timezone_name
self.tz_offset_minutes = tz_offset_minutes
@cache_result(
key_prefix="admin:stats:comparison",
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
user_specific=False,
vary_by=[
"current_start",
"current_end",
"comparison_type",
"timezone_name",
"tz_offset_minutes",
],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
if self.current_start > self.current_end:
raise HTTPException(status_code=400, detail="current_start must be <= current_end")
days = (self.current_end - self.current_start).days + 1
def _safe_year_shift(value: date) -> date:
try:
return value.replace(year=value.year - 1)
except ValueError:
return value.replace(year=value.year - 1, day=28)
if self.comparison_type == "period":
comparison_end = self.current_start - timedelta(days=1)
comparison_start = comparison_end - timedelta(days=days - 1)
else:
comparison_start = _safe_year_shift(self.current_start)
comparison_end = _safe_year_shift(self.current_end)
current_range = TimeRangeParams(
start_date=self.current_start,
end_date=self.current_end,
timezone=self.timezone_name,
tz_offset_minutes=self.tz_offset_minutes or 0,
).validate_and_resolve()
comparison_range = TimeRangeParams(
start_date=comparison_start,
end_date=comparison_end,
timezone=self.timezone_name,
tz_offset_minutes=self.tz_offset_minutes or 0,
).validate_and_resolve()
current_stats = query_stats_hybrid(context.db, current_range, filters=StatsFilter())
comparison_stats = query_stats_hybrid(context.db, comparison_range, filters=StatsFilter())
def _stats_payload(stats: AggregatedStats) -> dict[str, Any]:
total_tokens = (
stats.input_tokens
+ stats.output_tokens
+ stats.cache_creation_tokens
+ stats.cache_read_tokens
)
return {
"total_requests": stats.total_requests,
"total_tokens": total_tokens,
"total_cost": float(stats.total_cost),
"actual_total_cost": float(stats.actual_total_cost),
"avg_response_time_ms": float(stats.avg_response_time_ms),
"error_requests": stats.error_requests,
}
def _pct_change(current: float, previous: float) -> float | None:
if previous == 0:
return None if current != 0 else 0.0
return round((current - previous) / previous * 100, 2)
current_payload = _stats_payload(current_stats)
comparison_payload = _stats_payload(comparison_stats)
changes = {
key: _pct_change(float(current_payload[key]), float(comparison_payload[key]))
for key in current_payload.keys()
}
return {
"current": current_payload,
"comparison": comparison_payload,
"change_percent": changes,
"current_start": self.current_start.isoformat(),
"current_end": self.current_end.isoformat(),
"comparison_start": comparison_start.isoformat(),
"comparison_end": comparison_end.isoformat(),
}
@router.get("/comparison")
async def get_comparison(
request: Request,
db: Session = Depends(get_db),
current_start: date = Query(...),
current_end: date = Query(...),
comparison_type: Literal["period", "year"] = Query("period"),
timezone_name: str | None = Query(None, alias="timezone"),
tz_offset_minutes: int | None = Query(0),
) -> Any:
adapter = AdminComparisonAdapter(
current_start=current_start,
current_end=current_end,
comparison_type=comparison_type,
timezone_name=timezone_name,
tz_offset_minutes=tz_offset_minutes,
)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
-215
View File
@@ -1,215 +0,0 @@
"""Admin cost stats routes."""
from __future__ import annotations
from datetime import date, timedelta
from typing import Any
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from sqlalchemy import func
from sqlalchemy.orm import Session
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext
from src.config.constants import CacheTTL
from src.database import get_db
from src.models.database import Usage
from src.services.system.stats_aggregator import query_time_series
from src.services.system.time_range import TimeRangeParams
from src.utils.cache_decorator import cache_result
from .common import (
_apply_admin_default_range,
_build_time_range_from_days,
_build_time_range_params,
_linear_regression,
pipeline,
)
router = APIRouter()
class AdminCostForecastAdapter(AdminApiAdapter):
def __init__(
self,
time_range: TimeRangeParams | None,
days: int,
forecast_days: int,
timezone_name: str | None,
tz_offset_minutes: int | None,
) -> None:
self.time_range = time_range
self.days = days
self.forecast_days = forecast_days
self.timezone_name = timezone_name
self.fallback_tz_offset_minutes = tz_offset_minutes
@cache_result(
key_prefix="admin:stats:cost:forecast",
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
user_specific=False,
vary_by=[
"time_range.start_date",
"time_range.end_date",
"time_range.preset",
"time_range.timezone",
"time_range.tz_offset_minutes",
"days",
"forecast_days",
"timezone_name",
"fallback_tz_offset_minutes",
],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
time_range = self.time_range or _build_time_range_from_days(
self.days, self.timezone_name, self.fallback_tz_offset_minutes
)
time_range.granularity = "day"
try:
series = query_time_series(context.db, time_range)
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
history = [
{"date": item["date"], "total_cost": float(item.get("total_cost", 0.0))}
for item in series
]
values = [item["total_cost"] for item in history]
slope, intercept = _linear_regression(values)
forecast = []
if history:
last_date = date.fromisoformat(history[-1]["date"])
else:
last_date = time_range.end_date
for i in range(self.forecast_days):
idx = len(values) + i
predicted = max(0.0, slope * idx + intercept)
forecast.append(
{
"date": (last_date + timedelta(days=i + 1)).isoformat(),
"total_cost": round(predicted, 4),
}
)
return {
"history": history,
"forecast": forecast,
"slope": round(slope, 6),
"intercept": round(intercept, 6),
"start_date": time_range.start_date.isoformat(),
"end_date": time_range.end_date.isoformat(),
}
@router.get("/cost/forecast")
async def get_cost_forecast(
request: Request,
db: Session = Depends(get_db),
start_date: date | None = Query(None),
end_date: date | None = Query(None),
preset: str | None = Query(None),
timezone_name: str | None = Query(None, alias="timezone"),
tz_offset_minutes: int | None = Query(0),
days: int = Query(30, ge=7, le=365),
forecast_days: int = Query(7, ge=1, le=90),
) -> Any:
time_range = _build_time_range_params(
start_date, end_date, preset, timezone_name, tz_offset_minutes
)
adapter = AdminCostForecastAdapter(
time_range=time_range,
days=days,
forecast_days=forecast_days,
timezone_name=timezone_name,
tz_offset_minutes=tz_offset_minutes,
)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
class AdminCostSavingsAdapter(AdminApiAdapter):
def __init__(
self,
time_range: TimeRangeParams | None,
provider_name: str | None,
model: str | None,
) -> None:
self.time_range = _apply_admin_default_range(time_range)
self.provider_name = provider_name
self.model = model
@cache_result(
key_prefix="admin:stats:cost:savings",
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
user_specific=False,
vary_by=[
"time_range.start_date",
"time_range.end_date",
"time_range.preset",
"time_range.timezone",
"time_range.tz_offset_minutes",
"provider_name",
"model",
],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
if not self.time_range:
return {
"cache_read_tokens": 0,
"cache_read_cost": 0.0,
"cache_creation_cost": 0.0,
"estimated_full_cost": 0.0,
"cache_savings": 0.0,
}
start_utc, end_utc = self.time_range.to_utc_datetime_range()
query = context.db.query(
func.sum(Usage.cache_read_input_tokens).label("cache_read_tokens"),
func.sum(Usage.cache_read_cost_usd).label("cache_read_cost"),
func.sum(Usage.cache_creation_cost_usd).label("cache_creation_cost"),
func.sum(
func.coalesce(Usage.output_price_per_1m, 0) * Usage.cache_read_input_tokens
).label("estimated_full_cost_raw"),
).filter(Usage.created_at >= start_utc, Usage.created_at < end_utc)
if self.provider_name:
query = query.filter(Usage.provider_name == self.provider_name)
if self.model:
query = query.filter(Usage.model == self.model)
row = query.first()
cache_read_tokens = int(getattr(row, "cache_read_tokens", 0) or 0)
cache_read_cost = float(getattr(row, "cache_read_cost", 0) or 0.0)
cache_creation_cost = float(getattr(row, "cache_creation_cost", 0) or 0.0)
estimated_full_cost = float(getattr(row, "estimated_full_cost_raw", 0) or 0.0) / 1_000_000
if estimated_full_cost <= 0 and cache_read_cost > 0:
estimated_full_cost = cache_read_cost * 10
cache_savings = estimated_full_cost - cache_read_cost
return {
"cache_read_tokens": cache_read_tokens,
"cache_read_cost": round(cache_read_cost, 6),
"cache_creation_cost": round(cache_creation_cost, 6),
"estimated_full_cost": round(estimated_full_cost, 6),
"cache_savings": round(cache_savings, 6),
}
@router.get("/cost/savings")
async def get_cost_savings(
request: Request,
db: Session = Depends(get_db),
start_date: date | None = Query(None),
end_date: date | None = Query(None),
preset: str | None = Query(None),
timezone_name: str | None = Query(None, alias="timezone"),
tz_offset_minutes: int | None = Query(0),
provider_name: str | None = Query(None),
model: str | None = Query(None),
) -> Any:
time_range = _build_time_range_params(
start_date, end_date, preset, timezone_name, tz_offset_minutes
)
adapter = AdminCostSavingsAdapter(
time_range=time_range, provider_name=provider_name, model=model
)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@@ -1,125 +0,0 @@
"""Admin error stats routes."""
from __future__ import annotations
from datetime import date, timezone
from typing import Any
from fastapi import APIRouter, Depends, Query, Request
from sqlalchemy import func
from sqlalchemy.orm import Session
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext
from src.config.constants import CacheTTL
from src.database import get_db
from src.models.database import StatsDailyError, Usage
from src.services.system.time_range import TimeRangeParams
from src.utils.cache_decorator import cache_result
from .common import _apply_admin_default_range, _build_time_range_params, pipeline
router = APIRouter()
class AdminErrorDistributionAdapter(AdminApiAdapter):
def __init__(self, time_range: TimeRangeParams | None) -> None:
self.time_range = _apply_admin_default_range(time_range)
@cache_result(
key_prefix="admin:stats:errors:distribution",
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
user_specific=False,
vary_by=[
"time_range.start_date",
"time_range.end_date",
"time_range.preset",
"time_range.timezone",
"time_range.tz_offset_minutes",
],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
if not self.time_range:
return {"distribution": [], "trend": []}
time_range = self.time_range
is_utc = (time_range.timezone in {None, "UTC"}) and time_range.tz_offset_minutes == 0
distribution: dict[str, int] = {}
trend: dict[str, dict[str, int]] = {}
if is_utc:
start_utc, end_utc = time_range.to_utc_datetime_range()
rows = (
context.db.query(StatsDailyError)
.filter(StatsDailyError.date >= start_utc, StatsDailyError.date < end_utc)
.all()
)
for row in rows:
date_str = (
row.date.astimezone(timezone.utc).date().isoformat()
if row.date.tzinfo
else row.date.date().isoformat()
)
distribution[row.error_category] = distribution.get(row.error_category, 0) + int(
row.count or 0
)
trend.setdefault(date_str, {})
trend[date_str][row.error_category] = trend[date_str].get(
row.error_category, 0
) + int(row.count or 0)
else:
for local_date, day_start_utc, day_end_utc in time_range.get_local_day_hours():
rows = (
context.db.query(
Usage.error_category,
func.count(Usage.id).label("count"),
)
.filter(
Usage.created_at >= day_start_utc,
Usage.created_at < day_end_utc,
Usage.error_category.isnot(None),
)
.group_by(Usage.error_category)
.all()
)
date_str = local_date.isoformat()
for row in rows:
distribution[row.error_category] = distribution.get(
row.error_category, 0
) + int(row.count or 0)
trend.setdefault(date_str, {})
trend[date_str][row.error_category] = trend[date_str].get(
row.error_category, 0
) + int(row.count or 0)
trend_items = []
for day in sorted(trend.keys()):
counts = trend[day]
total = sum(counts.values())
trend_items.append({"date": day, "total": total, "categories": counts})
distribution_items = [
{"category": category, "count": count}
for category, count in sorted(
distribution.items(), key=lambda item: item[1], reverse=True
)
]
return {"distribution": distribution_items, "trend": trend_items}
@router.get("/errors/distribution")
async def get_error_distribution(
request: Request,
db: Session = Depends(get_db),
start_date: date | None = Query(None),
end_date: date | None = Query(None),
preset: str | None = Query(None),
timezone_name: str | None = Query(None, alias="timezone"),
tz_offset_minutes: int | None = Query(0),
) -> Any:
time_range = _build_time_range_params(
start_date, end_date, preset, timezone_name, tz_offset_minutes
)
adapter = AdminErrorDistributionAdapter(time_range=time_range)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@@ -1,759 +0,0 @@
"""Admin leaderboard stats routes."""
from __future__ import annotations
import json
from datetime import date
from typing import Any, Literal
from fastapi import APIRouter, Depends, Query, Request
from sqlalchemy import func
from sqlalchemy.orm import Session
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext
from src.clients.redis_client import get_redis_client_sync
from src.config.constants import CacheTTL
from src.core.enums import UserRole
from src.database import get_db
from src.models.database import (
ApiKey,
StatsDailyApiKey,
StatsDailyModel,
StatsUserDaily,
Usage,
User,
)
from src.services.system.time_range import TimeRangeParams
from .common import (
_apply_admin_default_range,
_apply_usage_time_segments,
_build_cache_key,
_build_time_range_params,
_is_today_range,
_metric_order,
_split_daily_and_usage_segments,
_union_queries,
pipeline,
)
router = APIRouter()
class AdminUserLeaderboardAdapter(AdminApiAdapter):
def __init__(
self,
time_range: TimeRangeParams | None,
metric: Literal["requests", "tokens", "cost"],
order: Literal["asc", "desc"],
limit: int,
offset: int,
provider_name: str | None,
model: str | None,
include_inactive: bool,
exclude_admin: bool,
) -> None:
self.time_range = _apply_admin_default_range(time_range)
self.metric = metric
self.order = order
self.limit = limit
self.offset = offset
self.provider_name = provider_name
self.model = model
self.include_inactive = include_inactive
self.exclude_admin = exclude_admin
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
cacheable = not _is_today_range(self.time_range)
redis_client = get_redis_client_sync()
cache_key = None
if cacheable and redis_client:
cache_key = _build_cache_key(
"users",
self.metric,
self.time_range,
{
"order": self.order,
"limit": self.limit,
"offset": self.offset,
"provider_name": self.provider_name,
"model": self.model,
"include_inactive": self.include_inactive,
"exclude_admin": self.exclude_admin,
},
)
cached = await redis_client.get(cache_key)
if cached:
try:
return json.loads(cached)
except Exception:
pass
use_daily = self.time_range is not None and not self.provider_name and not self.model
daily_range, usage_segments = _split_daily_and_usage_segments(self.time_range, use_daily)
daily_query = None
if daily_range:
daily_query = (
db.query(
StatsUserDaily.user_id.label("entity_id"),
func.sum(StatsUserDaily.total_requests).label("requests"),
func.sum(
StatsUserDaily.input_tokens
+ StatsUserDaily.output_tokens
+ StatsUserDaily.cache_creation_tokens
+ StatsUserDaily.cache_read_tokens
).label("tokens"),
func.sum(StatsUserDaily.total_cost).label("cost"),
)
.filter(StatsUserDaily.date >= daily_range[0], StatsUserDaily.date < daily_range[1])
.group_by(StatsUserDaily.user_id)
)
usage_query = db.query(
Usage.user_id.label("entity_id"),
func.count(Usage.id).label("requests"),
func.sum(
Usage.input_tokens
+ Usage.output_tokens
+ Usage.cache_creation_input_tokens
+ Usage.cache_read_input_tokens
).label("tokens"),
func.sum(Usage.total_cost_usd).label("cost"),
).filter(
Usage.user_id.isnot(None),
Usage.status.notin_(["pending", "streaming"]),
Usage.provider_name.notin_(["unknown", "pending"]),
)
if self.provider_name:
usage_query = usage_query.filter(Usage.provider_name == self.provider_name)
if self.model:
usage_query = usage_query.filter(Usage.model == self.model)
usage_query = _apply_usage_time_segments(usage_query, usage_segments)
if usage_query is not None:
usage_query = usage_query.group_by(Usage.user_id)
union_query = _union_queries([daily_query, usage_query])
if union_query is None:
return {
"items": [],
"total": 0,
"metric": self.metric,
"start_date": self.time_range.start_date.isoformat() if self.time_range else None,
"end_date": self.time_range.end_date.isoformat() if self.time_range else None,
}
union_subq = union_query.subquery()
agg_subq = (
db.query(
union_subq.c.entity_id.label("entity_id"),
func.sum(union_subq.c.requests).label("requests"),
func.sum(union_subq.c.tokens).label("tokens"),
func.sum(union_subq.c.cost).label("cost"),
)
.group_by(union_subq.c.entity_id)
.subquery()
)
base_query = (
db.query(
User.id.label("id"),
User.username,
User.email,
agg_subq.c.requests,
agg_subq.c.tokens,
agg_subq.c.cost,
)
.join(agg_subq, agg_subq.c.entity_id == User.id)
.filter(User.is_deleted.is_(False))
)
if not self.include_inactive:
base_query = base_query.filter(User.is_active.is_(True))
if self.exclude_admin:
base_query = base_query.filter(User.role != UserRole.ADMIN)
metric_expr = {
"requests": agg_subq.c.requests,
"tokens": agg_subq.c.tokens,
"cost": agg_subq.c.cost,
}[self.metric]
order_expr = _metric_order(self.metric, self.order, metric_expr)
rank_expr = func.dense_rank().over(order_by=order_expr).label("rank")
total = db.query(func.count()).select_from(base_query.subquery()).scalar() or 0
rows = (
base_query.add_columns(rank_expr, metric_expr.label("metric_value"))
.order_by(order_expr)
.offset(self.offset)
.limit(self.limit)
.all()
)
items = []
for row in rows:
name = row.username or row.email or str(row.id)
value = row.metric_value or 0
if self.metric in {"requests", "tokens"}:
value = int(value)
else:
value = float(value)
items.append(
{
"rank": int(row.rank),
"id": row.id,
"name": name,
"value": value,
"requests": int(row.requests or 0),
"tokens": int(row.tokens or 0),
"cost": float(row.cost or 0.0),
}
)
context.add_audit_metadata(
action="leaderboard_users",
start_date=self.time_range.start_date.isoformat() if self.time_range else None,
end_date=self.time_range.end_date.isoformat() if self.time_range else None,
preset=self.time_range.preset if self.time_range else None,
timezone=self.time_range.timezone if self.time_range else None,
metric=self.metric,
order=self.order,
limit=self.limit,
offset=self.offset,
provider_name=self.provider_name,
model=self.model,
include_inactive=self.include_inactive,
exclude_admin=self.exclude_admin,
result_count=len(items),
total=total,
)
result = {
"items": items,
"total": total,
"metric": self.metric,
"start_date": self.time_range.start_date.isoformat() if self.time_range else None,
"end_date": self.time_range.end_date.isoformat() if self.time_range else None,
}
if cacheable and redis_client and cache_key:
try:
await redis_client.setex(
cache_key, CacheTTL.ADMIN_LEADERBOARD, json.dumps(result, ensure_ascii=False)
)
except Exception:
pass
return result
class AdminApiKeyLeaderboardAdapter(AdminApiAdapter):
def __init__(
self,
time_range: TimeRangeParams | None,
metric: Literal["requests", "tokens", "cost"],
order: Literal["asc", "desc"],
limit: int,
offset: int,
provider_name: str | None,
model: str | None,
include_inactive: bool,
exclude_admin: bool,
) -> None:
self.time_range = _apply_admin_default_range(time_range)
self.metric = metric
self.order = order
self.limit = limit
self.offset = offset
self.provider_name = provider_name
self.model = model
self.include_inactive = include_inactive
self.exclude_admin = exclude_admin
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
cacheable = not _is_today_range(self.time_range)
redis_client = get_redis_client_sync()
cache_key = None
if cacheable and redis_client:
cache_key = _build_cache_key(
"api_keys",
self.metric,
self.time_range,
{
"order": self.order,
"limit": self.limit,
"offset": self.offset,
"provider_name": self.provider_name,
"model": self.model,
"include_inactive": self.include_inactive,
"exclude_admin": self.exclude_admin,
},
)
cached = await redis_client.get(cache_key)
if cached:
try:
return json.loads(cached)
except Exception:
pass
use_daily = self.time_range is not None and not self.provider_name and not self.model
daily_range, usage_segments = _split_daily_and_usage_segments(self.time_range, use_daily)
daily_query = None
if daily_range:
daily_query = (
db.query(
StatsDailyApiKey.api_key_id.label("entity_id"),
func.sum(StatsDailyApiKey.total_requests).label("requests"),
func.sum(
StatsDailyApiKey.input_tokens
+ StatsDailyApiKey.output_tokens
+ StatsDailyApiKey.cache_creation_tokens
+ StatsDailyApiKey.cache_read_tokens
).label("tokens"),
func.sum(StatsDailyApiKey.total_cost).label("cost"),
)
.filter(
StatsDailyApiKey.date >= daily_range[0], StatsDailyApiKey.date < daily_range[1]
)
.group_by(StatsDailyApiKey.api_key_id)
)
usage_query = db.query(
Usage.api_key_id.label("entity_id"),
func.count(Usage.id).label("requests"),
func.sum(
Usage.input_tokens
+ Usage.output_tokens
+ Usage.cache_creation_input_tokens
+ Usage.cache_read_input_tokens
).label("tokens"),
func.sum(Usage.total_cost_usd).label("cost"),
).filter(
Usage.api_key_id.isnot(None),
Usage.status.notin_(["pending", "streaming"]),
Usage.provider_name.notin_(["unknown", "pending"]),
)
if self.provider_name:
usage_query = usage_query.filter(Usage.provider_name == self.provider_name)
if self.model:
usage_query = usage_query.filter(Usage.model == self.model)
usage_query = _apply_usage_time_segments(usage_query, usage_segments)
if usage_query is None:
return {
"items": [],
"total": 0,
"metric": self.metric,
"start_date": self.time_range.start_date.isoformat() if self.time_range else None,
"end_date": self.time_range.end_date.isoformat() if self.time_range else None,
}
usage_query = usage_query.group_by(Usage.api_key_id)
union_query = _union_queries([daily_query, usage_query])
if union_query is None:
return {
"items": [],
"total": 0,
"metric": self.metric,
"start_date": self.time_range.start_date.isoformat() if self.time_range else None,
"end_date": self.time_range.end_date.isoformat() if self.time_range else None,
}
union_subq = union_query.subquery()
agg_subq = (
db.query(
union_subq.c.entity_id.label("entity_id"),
func.sum(union_subq.c.requests).label("requests"),
func.sum(union_subq.c.tokens).label("tokens"),
func.sum(union_subq.c.cost).label("cost"),
)
.group_by(union_subq.c.entity_id)
.subquery()
)
base_query = (
db.query(
ApiKey,
User,
agg_subq.c.requests,
agg_subq.c.tokens,
agg_subq.c.cost,
)
.join(agg_subq, agg_subq.c.entity_id == ApiKey.id)
.join(User, User.id == ApiKey.user_id)
.filter(User.is_deleted.is_(False))
)
if not self.include_inactive:
base_query = base_query.filter(ApiKey.is_active.is_(True))
if self.exclude_admin:
base_query = base_query.filter(User.role != UserRole.ADMIN)
metric_expr = {
"requests": agg_subq.c.requests,
"tokens": agg_subq.c.tokens,
"cost": agg_subq.c.cost,
}[self.metric]
order_expr = _metric_order(self.metric, self.order, metric_expr)
rank_expr = func.dense_rank().over(order_by=order_expr).label("rank")
total = db.query(func.count()).select_from(base_query.subquery()).scalar() or 0
rows = (
base_query.add_columns(rank_expr, metric_expr.label("metric_value"))
.order_by(order_expr)
.offset(self.offset)
.limit(self.limit)
.all()
)
items = []
for row in rows:
api_key = row.ApiKey
name = api_key.name or api_key.get_display_key()
value = row.metric_value or 0
if self.metric in {"requests", "tokens"}:
value = int(value)
else:
value = float(value)
items.append(
{
"rank": int(row.rank),
"id": api_key.id,
"name": name,
"value": value,
"requests": int(row.requests or 0),
"tokens": int(row.tokens or 0),
"cost": float(row.cost or 0.0),
}
)
context.add_audit_metadata(
action="leaderboard_api_keys",
start_date=self.time_range.start_date.isoformat() if self.time_range else None,
end_date=self.time_range.end_date.isoformat() if self.time_range else None,
preset=self.time_range.preset if self.time_range else None,
timezone=self.time_range.timezone if self.time_range else None,
metric=self.metric,
order=self.order,
limit=self.limit,
offset=self.offset,
provider_name=self.provider_name,
model=self.model,
include_inactive=self.include_inactive,
exclude_admin=self.exclude_admin,
result_count=len(items),
total=total,
)
result = {
"items": items,
"total": total,
"metric": self.metric,
"start_date": self.time_range.start_date.isoformat() if self.time_range else None,
"end_date": self.time_range.end_date.isoformat() if self.time_range else None,
}
if cacheable and redis_client and cache_key:
try:
await redis_client.setex(
cache_key, CacheTTL.ADMIN_LEADERBOARD, json.dumps(result, ensure_ascii=False)
)
except Exception:
pass
return result
class AdminModelLeaderboardAdapter(AdminApiAdapter):
def __init__(
self,
time_range: TimeRangeParams | None,
metric: Literal["requests", "tokens", "cost"],
order: Literal["asc", "desc"],
limit: int,
offset: int,
provider_name: str | None,
model: str | None,
) -> None:
self.time_range = _apply_admin_default_range(time_range)
self.metric = metric
self.order = order
self.limit = limit
self.offset = offset
self.provider_name = provider_name
self.model = model
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
cacheable = not _is_today_range(self.time_range)
redis_client = get_redis_client_sync()
cache_key = None
if cacheable and redis_client:
cache_key = _build_cache_key(
"models",
self.metric,
self.time_range,
{
"order": self.order,
"limit": self.limit,
"offset": self.offset,
"provider_name": self.provider_name,
"model": self.model,
},
)
cached = await redis_client.get(cache_key)
if cached:
try:
return json.loads(cached)
except Exception:
pass
use_daily = self.time_range is not None and not self.provider_name
daily_range, usage_segments = _split_daily_and_usage_segments(self.time_range, use_daily)
daily_query = None
if daily_range:
daily_query = (
db.query(
StatsDailyModel.model.label("entity_id"),
func.sum(StatsDailyModel.total_requests).label("requests"),
func.sum(
StatsDailyModel.input_tokens
+ StatsDailyModel.output_tokens
+ StatsDailyModel.cache_creation_tokens
+ StatsDailyModel.cache_read_tokens
).label("tokens"),
func.sum(StatsDailyModel.total_cost).label("cost"),
)
.filter(
StatsDailyModel.date >= daily_range[0], StatsDailyModel.date < daily_range[1]
)
.group_by(StatsDailyModel.model)
)
if self.model:
daily_query = daily_query.filter(StatsDailyModel.model == self.model)
usage_query = db.query(
Usage.model.label("entity_id"),
func.count(Usage.id).label("requests"),
func.sum(
Usage.input_tokens
+ Usage.output_tokens
+ Usage.cache_creation_input_tokens
+ Usage.cache_read_input_tokens
).label("tokens"),
func.sum(Usage.total_cost_usd).label("cost"),
).filter(
Usage.status.notin_(["pending", "streaming"]),
Usage.provider_name.notin_(["unknown", "pending"]),
)
if self.provider_name:
usage_query = usage_query.filter(Usage.provider_name == self.provider_name)
if self.model:
usage_query = usage_query.filter(Usage.model == self.model)
usage_query = _apply_usage_time_segments(usage_query, usage_segments)
if usage_query is not None:
usage_query = usage_query.group_by(Usage.model)
union_query = _union_queries([daily_query, usage_query])
if union_query is None:
return {
"items": [],
"total": 0,
"metric": self.metric,
"start_date": self.time_range.start_date.isoformat() if self.time_range else None,
"end_date": self.time_range.end_date.isoformat() if self.time_range else None,
}
union_subq = union_query.subquery()
agg_subq = (
db.query(
union_subq.c.entity_id.label("entity_id"),
func.sum(union_subq.c.requests).label("requests"),
func.sum(union_subq.c.tokens).label("tokens"),
func.sum(union_subq.c.cost).label("cost"),
)
.group_by(union_subq.c.entity_id)
.subquery()
)
base_query = db.query(
agg_subq.c.entity_id.label("id"),
agg_subq.c.entity_id.label("name"),
agg_subq.c.requests,
agg_subq.c.tokens,
agg_subq.c.cost,
)
metric_expr = {
"requests": agg_subq.c.requests,
"tokens": agg_subq.c.tokens,
"cost": agg_subq.c.cost,
}[self.metric]
order_expr = _metric_order(self.metric, self.order, metric_expr)
rank_expr = func.dense_rank().over(order_by=order_expr).label("rank")
total = db.query(func.count()).select_from(base_query.subquery()).scalar() or 0
rows = (
base_query.add_columns(rank_expr, metric_expr.label("metric_value"))
.order_by(order_expr)
.offset(self.offset)
.limit(self.limit)
.all()
)
items = []
for row in rows:
value = row.metric_value or 0
if self.metric in {"requests", "tokens"}:
value = int(value)
else:
value = float(value)
items.append(
{
"rank": int(row.rank),
"id": row.id,
"name": row.name,
"value": value,
"requests": int(row.requests or 0),
"tokens": int(row.tokens or 0),
"cost": float(row.cost or 0.0),
}
)
context.add_audit_metadata(
action="leaderboard_models",
start_date=self.time_range.start_date.isoformat() if self.time_range else None,
end_date=self.time_range.end_date.isoformat() if self.time_range else None,
preset=self.time_range.preset if self.time_range else None,
timezone=self.time_range.timezone if self.time_range else None,
metric=self.metric,
order=self.order,
limit=self.limit,
offset=self.offset,
provider_name=self.provider_name,
model=self.model,
result_count=len(items),
total=total,
)
result = {
"items": items,
"total": total,
"metric": self.metric,
"start_date": self.time_range.start_date.isoformat() if self.time_range else None,
"end_date": self.time_range.end_date.isoformat() if self.time_range else None,
}
if cacheable and redis_client and cache_key:
try:
await redis_client.setex(
cache_key, CacheTTL.ADMIN_LEADERBOARD, json.dumps(result, ensure_ascii=False)
)
except Exception:
pass
return result
@router.get("/leaderboard/users")
async def get_user_leaderboard(
request: Request,
db: Session = Depends(get_db),
start_date: date | None = Query(None),
end_date: date | None = Query(None),
preset: str | None = Query(None),
timezone_name: str | None = Query(None, alias="timezone"),
tz_offset_minutes: int | None = Query(0),
metric: Literal["requests", "tokens", "cost"] = Query("requests"),
order: Literal["desc", "asc"] = Query("desc"),
limit: int = Query(10, ge=1, le=100),
offset: int = Query(0, ge=0),
provider_name: str | None = Query(None),
model: str | None = Query(None),
include_inactive: bool = Query(False),
exclude_admin: bool = Query(False),
) -> Any:
time_range = _build_time_range_params(
start_date, end_date, preset, timezone_name, tz_offset_minutes
)
adapter = AdminUserLeaderboardAdapter(
time_range=time_range,
metric=metric,
order=order,
limit=limit,
offset=offset,
provider_name=provider_name,
model=model,
include_inactive=include_inactive,
exclude_admin=exclude_admin,
)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.get("/leaderboard/api-keys")
async def get_api_key_leaderboard(
request: Request,
db: Session = Depends(get_db),
start_date: date | None = Query(None),
end_date: date | None = Query(None),
preset: str | None = Query(None),
timezone_name: str | None = Query(None, alias="timezone"),
tz_offset_minutes: int | None = Query(0),
metric: Literal["requests", "tokens", "cost"] = Query("requests"),
order: Literal["desc", "asc"] = Query("desc"),
limit: int = Query(10, ge=1, le=100),
offset: int = Query(0, ge=0),
provider_name: str | None = Query(None),
model: str | None = Query(None),
include_inactive: bool = Query(False),
exclude_admin: bool = Query(False),
) -> Any:
time_range = _build_time_range_params(
start_date, end_date, preset, timezone_name, tz_offset_minutes
)
adapter = AdminApiKeyLeaderboardAdapter(
time_range=time_range,
metric=metric,
order=order,
limit=limit,
offset=offset,
provider_name=provider_name,
model=model,
include_inactive=include_inactive,
exclude_admin=exclude_admin,
)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.get("/leaderboard/models")
async def get_model_leaderboard(
request: Request,
db: Session = Depends(get_db),
start_date: date | None = Query(None),
end_date: date | None = Query(None),
preset: str | None = Query(None),
timezone_name: str | None = Query(None, alias="timezone"),
tz_offset_minutes: int | None = Query(0),
metric: Literal["requests", "tokens", "cost"] = Query("requests"),
order: Literal["desc", "asc"] = Query("desc"),
limit: int = Query(10, ge=1, le=100),
offset: int = Query(0, ge=0),
provider_name: str | None = Query(None),
model: str | None = Query(None),
) -> Any:
time_range = _build_time_range_params(
start_date, end_date, preset, timezone_name, tz_offset_minutes
)
adapter = AdminModelLeaderboardAdapter(
time_range=time_range,
metric=metric,
order=order,
limit=limit,
offset=offset,
provider_name=provider_name,
model=model,
)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@@ -1,93 +0,0 @@
"""Admin performance stats routes."""
from __future__ import annotations
from datetime import date, timezone
from typing import Any
from fastapi import APIRouter, Depends, Query, Request
from sqlalchemy.orm import Session
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext
from src.config.constants import CacheTTL
from src.database import get_db
from src.models.database import StatsDaily
from src.services.system.stats_aggregator import StatsAggregatorService
from src.services.system.time_range import TimeRangeParams
from src.utils.cache_decorator import cache_result
from .common import _apply_admin_default_range, _build_time_range_params, pipeline
router = APIRouter()
class AdminPercentilesAdapter(AdminApiAdapter):
def __init__(self, time_range: TimeRangeParams | None) -> None:
self.time_range = _apply_admin_default_range(time_range)
@cache_result(
key_prefix="admin:stats:performance:percentiles",
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
user_specific=False,
vary_by=[
"time_range.start_date",
"time_range.end_date",
"time_range.preset",
"time_range.timezone",
"time_range.tz_offset_minutes",
],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
if not self.time_range:
return []
time_range = self.time_range
is_utc = (time_range.timezone in {None, "UTC"}) and time_range.tz_offset_minutes == 0
if is_utc:
start_utc, end_utc = time_range.to_utc_datetime_range()
rows = (
context.db.query(StatsDaily)
.filter(StatsDaily.date >= start_utc, StatsDaily.date < end_utc)
.order_by(StatsDaily.date.asc())
.all()
)
result = []
for row in rows:
date_str = (
row.date.astimezone(timezone.utc).date().isoformat()
if row.date.tzinfo
else row.date.date().isoformat()
)
result.append(
{
"date": date_str,
"p50_response_time_ms": row.p50_response_time_ms,
"p90_response_time_ms": row.p90_response_time_ms,
"p99_response_time_ms": row.p99_response_time_ms,
"p50_first_byte_time_ms": row.p50_first_byte_time_ms,
"p90_first_byte_time_ms": row.p90_first_byte_time_ms,
"p99_first_byte_time_ms": row.p99_first_byte_time_ms,
}
)
return result
return StatsAggregatorService.compute_percentiles_by_local_day(context.db, time_range)
@router.get("/performance/percentiles")
async def get_percentiles(
request: Request,
db: Session = Depends(get_db),
start_date: date | None = Query(None),
end_date: date | None = Query(None),
preset: str | None = Query(None),
timezone_name: str | None = Query(None, alias="timezone"),
tz_offset_minutes: int | None = Query(0),
) -> Any:
time_range = _build_time_range_params(
start_date, end_date, preset, timezone_name, tz_offset_minutes
)
adapter = AdminPercentilesAdapter(time_range=time_range)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@@ -1,90 +0,0 @@
"""Admin quota usage stats routes."""
from __future__ import annotations
from datetime import datetime, timedelta, timezone
from typing import Any
from fastapi import APIRouter, Depends, Request
from sqlalchemy.orm import Session
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext
from src.config.constants import CacheTTL
from src.core.enums import ProviderBillingType
from src.database import get_db
from src.models.database import Provider
from src.utils.cache_decorator import cache_result
from .common import pipeline
router = APIRouter()
class AdminQuotaUsageAdapter(AdminApiAdapter):
@cache_result(
key_prefix="admin:stats:providers:quota_usage",
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
user_specific=False,
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
providers = (
db.query(Provider)
.filter(
(Provider.billing_type == ProviderBillingType.MONTHLY_QUOTA)
| (Provider.monthly_quota_usd.isnot(None))
)
.all()
)
now = datetime.now(timezone.utc)
result = []
for provider in providers:
quota = float(provider.monthly_quota_usd or 0)
used = float(provider.monthly_used_usd or 0)
remaining = max(quota - used, 0.0)
usage_percent = round((used / quota) * 100, 2) if quota > 0 else 0.0
reset_at = provider.quota_last_reset_at
if reset_at:
days_elapsed = max(1, (now - reset_at).days)
else:
days_elapsed = max(1, now.day - 1)
daily_rate = used / days_elapsed if used > 0 else 0.0
estimated_exhaust_at = None
if daily_rate > 0 and remaining > 0:
estimated_exhaust_at = now + timedelta(days=remaining / daily_rate)
if provider.quota_expires_at:
if not estimated_exhaust_at or provider.quota_expires_at < estimated_exhaust_at:
estimated_exhaust_at = provider.quota_expires_at
result.append(
{
"id": provider.id,
"name": provider.name,
"quota_usd": float(quota),
"used_usd": float(used),
"remaining_usd": float(remaining),
"usage_percent": usage_percent,
"quota_expires_at": (
provider.quota_expires_at.isoformat() if provider.quota_expires_at else None
),
"estimated_exhaust_at": (
estimated_exhaust_at.isoformat() if estimated_exhaust_at else None
),
}
)
result.sort(key=lambda x: x["usage_percent"], reverse=True)
return {"providers": result}
@router.get("/providers/quota-usage")
async def get_quota_usage(
request: Request,
db: Session = Depends(get_db),
) -> Any:
adapter = AdminQuotaUsageAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@@ -1,5 +0,0 @@
"""Admin stats routes (compat export)."""
from . import router
__all__ = ["router"]
@@ -1,93 +0,0 @@
"""Admin time series stats routes."""
from __future__ import annotations
from datetime import date
from typing import Any, Literal
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from sqlalchemy.orm import Session
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext
from src.config.constants import CacheTTL
from src.database import get_db
from src.services.system.stats_aggregator import TimeSeriesFilter, query_time_series
from src.services.system.time_range import TimeRangeParams
from src.utils.cache_decorator import cache_result
from .common import _apply_admin_default_range, _build_time_range_params, pipeline
router = APIRouter()
class AdminTimeSeriesAdapter(AdminApiAdapter):
def __init__(
self,
time_range: TimeRangeParams | None,
user_id: str | None,
model: str | None,
provider_name: str | None,
) -> None:
self.time_range = _apply_admin_default_range(time_range)
self.user_id = user_id
self.model = model
self.provider_name = provider_name
@cache_result(
key_prefix="admin:stats:time_series",
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
user_specific=False,
vary_by=[
"time_range.start_date",
"time_range.end_date",
"time_range.preset",
"time_range.timezone",
"time_range.tz_offset_minutes",
"time_range.granularity",
"user_id",
"model",
"provider_name",
],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
if not self.time_range:
return []
try:
return query_time_series(
context.db,
self.time_range,
filters=TimeSeriesFilter(
user_id=self.user_id, model=self.model, provider_name=self.provider_name
),
)
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
@router.get("/time-series")
async def get_time_series(
request: Request,
db: Session = Depends(get_db),
start_date: date | None = Query(None),
end_date: date | None = Query(None),
preset: str | None = Query(None),
granularity: Literal["hour", "day", "week", "month"] = Query("day"),
timezone_name: str | None = Query(None, alias="timezone"),
tz_offset_minutes: int | None = Query(0),
user_id: str | None = Query(None),
model: str | None = Query(None),
provider_name: str | None = Query(None),
) -> Any:
time_range = _build_time_range_params(
start_date, end_date, preset, timezone_name, tz_offset_minutes
)
if time_range:
time_range.granularity = granularity
adapter = AdminTimeSeriesAdapter(
time_range=time_range,
user_id=user_id,
model=model,
provider_name=provider_name,
)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
File diff suppressed because it is too large Load Diff
@@ -1,46 +0,0 @@
"""Announcement system routers."""
from __future__ import annotations
from fastapi import APIRouter
from starlette.routing import BaseRoute
from .routes import router as announcement_router
_RUST_OWNED_ANNOUNCEMENT_ROUTE_SIGNATURES = frozenset(
{
("GET", "/api/announcements"),
("GET", "/api/announcements/active"),
("POST", "/api/announcements"),
("GET", "/api/announcements/{announcement_id}"),
("PATCH", "/api/announcements/{announcement_id}/read-status"),
("PUT", "/api/announcements/{announcement_id}"),
("DELETE", "/api/announcements/{announcement_id}"),
("GET", "/api/announcements/users/me/unread-count"),
}
)
def _route_is_rust_owned(route: BaseRoute) -> bool:
path = getattr(route, "path", None)
methods = getattr(route, "methods", None)
if not isinstance(path, str) or not methods:
return False
return any(
(method, path) in _RUST_OWNED_ANNOUNCEMENT_ROUTE_SIGNATURES
for method in methods
if method not in {"HEAD", "OPTIONS"}
)
def _build_python_announcement_router() -> APIRouter:
router = APIRouter()
router.include_router(announcement_router)
router.routes = [route for route in router.routes if not _route_is_rust_owned(route)]
return router
python_announcement_router = _build_python_announcement_router()
router = python_announcement_router
__all__ = ["python_announcement_router", "router"]
-48
View File
@@ -1,48 +0,0 @@
"""Authentication route group."""
from __future__ import annotations
from fastapi import APIRouter
from starlette.routing import BaseRoute
from .routes import router as auth_router
_RUST_OWNED_AUTH_ROUTE_SIGNATURES = frozenset(
{
("GET", "/api/auth/registration-settings"),
("GET", "/api/auth/settings"),
("POST", "/api/auth/login"),
("POST", "/api/auth/refresh"),
("POST", "/api/auth/register"),
("GET", "/api/auth/me"),
("POST", "/api/auth/logout"),
("POST", "/api/auth/send-verification-code"),
("POST", "/api/auth/verify-email"),
("POST", "/api/auth/verification-status"),
}
)
def _route_is_rust_owned(route: BaseRoute) -> bool:
path = getattr(route, "path", None)
methods = getattr(route, "methods", None)
if not isinstance(path, str) or not methods:
return False
return any(
(method, path) in _RUST_OWNED_AUTH_ROUTE_SIGNATURES
for method in methods
if method not in {"HEAD", "OPTIONS"}
)
def _build_python_auth_router() -> APIRouter:
router = APIRouter()
router.include_router(auth_router)
router.routes = [route for route in router.routes if not _route_is_rust_owned(route)]
return router
python_auth_router = _build_python_auth_router()
router = python_auth_router
__all__ = ["python_auth_router", "router"]
@@ -1,42 +0,0 @@
"""Dashboard API routers."""
from __future__ import annotations
from fastapi import APIRouter
from starlette.routing import BaseRoute
from .routes import router as dashboard_router
_RUST_OWNED_DASHBOARD_ROUTE_SIGNATURES = frozenset(
{
("GET", "/api/dashboard/stats"),
("GET", "/api/dashboard/recent-requests"),
("GET", "/api/dashboard/provider-status"),
("GET", "/api/dashboard/daily-stats"),
}
)
def _route_is_rust_owned(route: BaseRoute) -> bool:
path = getattr(route, "path", None)
methods = getattr(route, "methods", None)
if not isinstance(path, str) or not methods:
return False
return any(
(method, path) in _RUST_OWNED_DASHBOARD_ROUTE_SIGNATURES
for method in methods
if method not in {"HEAD", "OPTIONS"}
)
def _build_python_dashboard_router() -> APIRouter:
router = APIRouter()
router.include_router(dashboard_router)
router.routes = [route for route in router.routes if not _route_is_rust_owned(route)]
return router
python_dashboard_router = _build_python_dashboard_router()
router = python_dashboard_router
__all__ = ["python_dashboard_router", "router"]
File diff suppressed because it is too large Load Diff
@@ -1,68 +0,0 @@
"""
Handler 基类模块
提供 Adapter、Handler 的抽象基类,以及请求构建器和响应解析器。
注意:Handler 基类(ChatHandlerBase, CliMessageHandlerBase 等)不在这里导出,
因为它们依赖 services.usage.stream,而后者又需要导入 response_parser,
会形成循环导入。请直接从具体模块导入 Handler 基类。
"""
# Chat Adapter 基类(不会引起循环导入)
from src.api.handlers.base.chat_adapter_base import (
ChatAdapterBase,
get_adapter_class,
get_adapter_instance,
list_registered_formats,
register_adapter,
)
# CLI Adapter 基类
from src.api.handlers.base.cli_adapter_base import (
CliAdapterBase,
get_cli_adapter_class,
get_cli_adapter_instance,
list_registered_cli_formats,
register_cli_adapter,
)
# 请求构建器
from src.api.handlers.base.request_builder import (
SENSITIVE_HEADERS,
PassthroughRequestBuilder,
RequestBuilder,
build_passthrough_request,
)
# 响应解析器
from src.api.handlers.base.response_parser import (
ParsedChunk,
ParsedResponse,
ResponseParser,
StreamStats,
)
__all__ = [
# Chat Adapter
"ChatAdapterBase",
"register_adapter",
"get_adapter_class",
"get_adapter_instance",
"list_registered_formats",
# CLI Adapter
"CliAdapterBase",
"register_cli_adapter",
"get_cli_adapter_class",
"get_cli_adapter_instance",
"list_registered_cli_formats",
# 请求构建器
"RequestBuilder",
"PassthroughRequestBuilder",
"build_passthrough_request",
"SENSITIVE_HEADERS",
# 响应解析器
"ResponseParser",
"ParsedChunk",
"ParsedResponse",
"StreamStats",
]
@@ -1,19 +0,0 @@
"""Internal routers still surfaced by the Python host."""
from fastapi import APIRouter
def _build_python_internal_router() -> APIRouter:
"""Internal APIs that still belong to the Python host/runtime."""
return APIRouter()
# Legacy internal gateway bridge 与 internal tunnel 模块仍保留在 `src.api.internal.*`
# 里给测试与过渡逻辑复用,但 Python host 已不再公开任何 `/api/internal/*` 路由。
python_internal_router = _build_python_internal_router()
router = python_internal_router
__all__ = [
"python_internal_router",
"router",
]
-14
View File
@@ -1,14 +0,0 @@
from __future__ import annotations
import ipaddress
from fastapi import HTTPException, Request
def ensure_loopback(request: Request) -> None:
host = request.client.host if request.client else ""
try:
if not ipaddress.ip_address(host).is_loopback:
raise ValueError(host)
except ValueError as exc:
raise HTTPException(status_code=403, detail="loopback access only") from exc
File diff suppressed because it is too large Load Diff
@@ -1,43 +0,0 @@
"""Compatibility re-export layer for gateway chat builders."""
from __future__ import annotations
from .gateway_chat_decision import (
_build_chat_stream_decision,
_build_chat_sync_decision,
_build_claude_chat_stream_decision,
_build_claude_chat_sync_decision,
_build_gemini_chat_stream_decision,
_build_gemini_chat_sync_decision,
_build_openai_chat_stream_decision,
_build_openai_chat_sync_decision,
)
from .gateway_chat_plan import (
_build_chat_stream_plan,
_build_chat_sync_plan,
_build_claude_chat_stream_plan,
_build_claude_chat_sync_plan,
_build_gemini_chat_stream_plan,
_build_gemini_chat_sync_plan,
_build_openai_chat_stream_plan,
_build_openai_chat_sync_plan,
)
__all__ = [
"_build_chat_sync_decision",
"_build_openai_chat_sync_decision",
"_build_chat_stream_decision",
"_build_openai_chat_stream_decision",
"_build_claude_chat_stream_decision",
"_build_gemini_chat_stream_decision",
"_build_claude_chat_sync_decision",
"_build_gemini_chat_sync_decision",
"_build_chat_sync_plan",
"_build_openai_chat_sync_plan",
"_build_chat_stream_plan",
"_build_openai_chat_stream_plan",
"_build_claude_chat_stream_plan",
"_build_gemini_chat_stream_plan",
"_build_claude_chat_sync_plan",
"_build_gemini_chat_sync_plan",
]
@@ -1,945 +0,0 @@
from __future__ import annotations
import base64
import time
import uuid
from dataclasses import asdict
from datetime import datetime, timezone
from typing import Any
from fastapi import BackgroundTasks, HTTPException, Request
from fastapi.responses import JSONResponse, Response
from sqlalchemy.orm import Session
from src.api.base.context import ApiRequestContext
from src.api.base.pipeline import get_pipeline
from src.core.api_format.headers import extract_client_api_key_for_endpoint_with_query
from src.core.api_format.metadata import get_auth_config_for_endpoint
from src.core.crypto import crypto_service
from src.core.http_compression import normalize_content_encoding
from src.core.logger import logger
from src.database import create_session, get_db
from src.models.database import ApiKey, RequestCandidate, User
from src.services.auth.service import AuthService
from src.utils.async_utils import safe_create_task
from .common import ensure_loopback
from .gateway_contract import (
_GEMINI_FILES_DOWNLOAD_ROUTE_RE,
_GEMINI_FILES_RESOURCE_ROUTE_RE,
_GEMINI_MODEL_OPERATION_CANCEL_RE,
_GEMINI_OPERATION_CANCEL_RE,
_GEMINI_SYNC_ROUTE_RE,
_GEMINI_VIDEO_CREATE_ROUTE_RE,
_GEMINI_VIDEO_MODEL_OPERATION_ANY_RE,
_OPENAI_VIDEO_CANCEL_ROUTE_RE,
_OPENAI_VIDEO_CONTENT_ROUTE_RE,
_OPENAI_VIDEO_REMIX_ROUTE_RE,
_OPENAI_VIDEO_TASK_ROUTE_RE,
CONTROL_ACTION_HEADER,
CONTROL_ACTION_PROXY_PUBLIC,
CONTROL_EXECUTED_HEADER,
GatewayAuthContext,
GatewayExecuteRequest,
GatewayExecutionDecisionResponse,
GatewayExecutionPlanResponse,
GatewayResolveRequest,
GatewayRouteDecision,
GatewayStreamReportRequest,
GatewaySyncReportRequest,
classify_gateway_route,
)
class _GatewayProxy:
def __getattr__(self, name: str) -> Any:
from . import gateway as gateway_module
return getattr(gateway_module, name)
gateway_module = _GatewayProxy()
async def _build_chat_sync_decision(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
expected_api_format: str,
decision_kind: str,
report_kind: str,
finalize_kind: str,
) -> GatewayExecutionDecisionResponse | None:
from src.api.handlers.base.chat_adapter_base import ChatAdapterBase
from src.api.handlers.base.chat_sync_executor import ChatSyncExecutor
from src.services.task.request_state import MutableRequestBodyState
if str(payload.method or "").strip().upper() != "POST":
return None
adapter, path_params = gateway_module._resolve_gateway_sync_adapter(decision, payload.path)
if not isinstance(adapter, ChatAdapterBase):
return None
if gateway_module._is_stream_request_payload(payload.body_json, path_params):
return None
user, api_key = gateway_module._load_gateway_auth_models(db, auth_context)
pipeline = gateway_module.get_pipeline()
await pipeline._check_user_rate_limit(request, db, user, api_key)
context = gateway_module._build_gateway_request_context(
request=request,
payload=payload,
db=db,
user=user,
api_key=api_key,
adapter=adapter,
path_params=path_params,
balance_remaining=auth_context.balance_remaining,
)
authorize_result = adapter.authorize(context)
if hasattr(authorize_result, "__await__"):
await authorize_result
original_request_body = await context.ensure_json_body_async()
if context.path_params:
original_request_body = adapter._merge_path_params(
original_request_body,
context.path_params,
)
request_obj = adapter._validate_request_body(original_request_body, context.path_params)
if isinstance(request_obj, JSONResponse):
return None
handler = adapter._create_handler(
db=db,
user=user,
api_key=api_key,
request_id=context.request_id,
client_ip=context.client_ip,
user_agent=context.user_agent,
start_time=context.start_time,
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
)
await handler._convert_request(request_obj)
model = handler.extract_model_from_request(
original_request_body,
context.path_params,
)
api_format = handler.allowed_api_formats[0]
capability_requirements = handler._resolve_capability_requirements(
model_name=model,
request_headers=context.original_headers,
request_body=original_request_body,
)
preferred_key_ids = await handler._resolve_preferred_key_ids(
model_name=model,
request_body=original_request_body,
)
candidate = await gateway_module._select_gateway_direct_candidate(
db=db,
redis_client=getattr(handler, "redis", None),
api_format=str(api_format),
model_name=str(model or "unknown"),
user_api_key=api_key,
request_id=context.request_id,
is_stream=False,
capability_requirements=capability_requirements or None,
preferred_key_ids=preferred_key_ids or None,
request_body=original_request_body,
)
if candidate is None:
return None
sync_executor = ChatSyncExecutor(handler)
prepared_plan = await sync_executor._build_sync_execution_plan(
candidate.provider,
candidate.endpoint,
candidate.key,
candidate,
model=str(model or "unknown"),
api_format=api_format,
original_headers=context.original_headers,
request_state=MutableRequestBodyState(original_request_body),
query_params=context.query_params,
client_content_encoding=context.client_content_encoding,
)
contract = prepared_plan.contract
contract_provider_api_format = str(contract.provider_api_format or "").strip().lower()
contract_client_api_format = str(contract.client_api_format or "").strip().lower()
if not prepared_plan.remote_eligible or contract_client_api_format != expected_api_format:
return None
provider_request_headers = {
str(header_name).strip().lower(): str(header_value).strip()
for header_name, header_value in dict(prepared_plan.headers or {}).items()
if str(header_name).strip() and str(header_value).strip()
}
provider_request_body = dict(prepared_plan.payload or {})
auth_header, auth_value = gateway_module._extract_gateway_upstream_auth(
provider_request_headers,
provider_api_format=contract_provider_api_format,
key=candidate.key,
)
prompt_cache_key = str(provider_request_body.get("prompt_cache_key") or "").strip() or None
mapped_model = getattr(sync_executor._ctx, "mapped_model_result", None)
has_envelope = prepared_plan.envelope is not None
needs_conversion = bool(prepared_plan.needs_conversion)
selected_report_kind = (
finalize_kind
if (
prepared_plan.upstream_is_stream
or needs_conversion
or has_envelope
or contract_provider_api_format != expected_api_format
)
else report_kind
)
decision_extra_headers: dict[str, str] = {}
decision_provider_request_headers = provider_request_headers or None
decision_provider_request_body = provider_request_body
report_context = {
"user_id": str(user.id),
"api_key_id": str(api_key.id),
"request_id": str(context.request_id),
"model": str(model or "unknown"),
"provider_name": str(contract.provider_name or "unknown"),
"provider_id": str(contract.provider_id or ""),
"endpoint_id": str(contract.endpoint_id or ""),
"key_id": str(contract.key_id or ""),
"candidate_id": str(contract.candidate_id or "") or None,
"provider_api_format": str(contract.provider_api_format or ""),
"client_api_format": str(contract.client_api_format or ""),
"mapped_model": str(mapped_model or "").strip() or None,
"original_headers": dict(context.original_headers),
"original_request_body": original_request_body,
"proxy_info": prepared_plan.proxy_info,
"has_envelope": has_envelope,
"envelope_name": gateway_module._gateway_report_context_envelope_name(
prepared_plan.envelope
),
"needs_conversion": needs_conversion,
}
report_context["provider_request_headers"] = provider_request_headers
report_context["provider_request_body"] = provider_request_body
return GatewayExecutionDecisionResponse(
action="executor_sync_decision",
decision_kind=decision_kind,
request_id=str(contract.request_id or context.request_id or ""),
candidate_id=str(contract.candidate_id or "").strip() or None,
provider_name=str(contract.provider_name or candidate.provider.name or ""),
provider_id=str(contract.provider_id or candidate.provider.id or ""),
endpoint_id=str(contract.endpoint_id or candidate.endpoint.id or ""),
key_id=str(contract.key_id or candidate.key.id or ""),
upstream_base_url=(
str(
prepared_plan.selected_base_url or getattr(candidate.endpoint, "base_url", "") or ""
).strip()
),
upstream_url=str(contract.url or "").strip() or None,
auth_header=auth_header,
auth_value=str(auth_value or "").strip(),
provider_api_format=contract_provider_api_format,
client_api_format=contract_client_api_format,
model_name=str(contract.model_name or model or "unknown"),
mapped_model=str(mapped_model or "").strip() or None,
prompt_cache_key=prompt_cache_key,
extra_headers=decision_extra_headers,
provider_request_headers=decision_provider_request_headers,
provider_request_body=decision_provider_request_body,
content_type=(
str(contract.content_type or provider_request_headers.get("content-type") or "").strip()
or "application/json"
),
proxy=gateway_module._serialize_gateway_sync_proxy(contract.proxy),
tls_profile=str(contract.tls_profile or "").strip() or None,
timeouts=gateway_module._serialize_gateway_sync_timeouts(contract.timeouts),
upstream_is_stream=prepared_plan.upstream_is_stream or None,
report_kind=selected_report_kind,
report_context=report_context,
auth_context=auth_context,
)
async def _build_openai_chat_sync_decision(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionDecisionResponse | None:
from src.api.handlers.base.chat_adapter_base import ChatAdapterBase
from src.config.settings import config
from src.services.provider.auth import get_provider_auth
from src.services.provider.behavior import get_provider_behavior
from src.services.provider.prompt_cache import maybe_patch_request_with_prompt_cache_key
from src.services.provider.stream_policy import (
get_upstream_stream_policy,
resolve_upstream_is_stream,
)
from src.services.provider.upstream_headers import build_upstream_extra_headers
from src.services.proxy_node.resolver import (
build_proxy_url_async,
get_system_proxy_config_async,
resolve_delegate_config_async,
resolve_effective_proxy,
resolve_proxy_info_async,
)
from src.services.request.execution_runtime_plan import (
ExecutionPlanTimeouts,
ExecutionProxySnapshot,
)
if str(payload.method or "").strip().upper() != "POST":
return None
adapter, path_params = gateway_module._resolve_gateway_sync_adapter(decision, payload.path)
if not isinstance(adapter, ChatAdapterBase):
return None
if gateway_module._is_stream_request_payload(payload.body_json, path_params):
return None
if not isinstance(payload.body_json, dict):
return None
user, api_key = gateway_module._load_gateway_auth_models(db, auth_context)
pipeline = gateway_module.get_pipeline()
await pipeline._check_user_rate_limit(request, db, user, api_key)
context = gateway_module._build_gateway_request_context(
request=request,
payload=payload,
db=db,
user=user,
api_key=api_key,
adapter=adapter,
path_params=path_params,
balance_remaining=auth_context.balance_remaining,
)
authorize_result = adapter.authorize(context)
if hasattr(authorize_result, "__await__"):
await authorize_result
original_request_body = await context.ensure_json_body_async()
if context.path_params:
original_request_body = adapter._merge_path_params(
original_request_body,
context.path_params,
)
request_obj = adapter._validate_request_body(original_request_body, context.path_params)
if isinstance(request_obj, JSONResponse):
return None
handler = adapter._create_handler(
db=db,
user=user,
api_key=api_key,
request_id=context.request_id,
client_ip=context.client_ip,
user_agent=context.user_agent,
start_time=context.start_time,
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
)
model = handler.extract_model_from_request(
original_request_body,
context.path_params,
)
client_api_format = str(handler.allowed_api_formats[0] or "").strip().lower()
if client_api_format != "openai:chat":
return None
capability_requirements = handler._resolve_capability_requirements(
model_name=model,
request_headers=context.original_headers,
request_body=original_request_body,
)
preferred_key_ids = await handler._resolve_preferred_key_ids(
model_name=model,
request_body=original_request_body,
)
candidate = await gateway_module._select_gateway_direct_candidate(
db=db,
redis_client=getattr(handler, "redis", None),
api_format=client_api_format,
model_name=str(model or "unknown"),
user_api_key=api_key,
request_id=context.request_id,
is_stream=False,
capability_requirements=capability_requirements or None,
preferred_key_ids=preferred_key_ids or None,
request_body=original_request_body,
)
if candidate is None:
return None
provider = candidate.provider
endpoint = candidate.endpoint
key = candidate.key
provider_api_format = str(endpoint.api_format or client_api_format or "").strip().lower()
if provider_api_format != "openai:chat":
return None
if getattr(endpoint, "custom_path", None):
return None
if getattr(endpoint, "header_rules", None) or getattr(endpoint, "body_rules", None):
return None
provider_type = str(getattr(provider, "provider_type", "") or "").strip().lower()
behavior = get_provider_behavior(
provider_type=provider_type,
endpoint_sig=provider_api_format,
)
if (
behavior.envelope is not None
or behavior.same_format_variant is not None
or behavior.cross_format_variant is not None
):
return None
upstream_policy = get_upstream_stream_policy(
endpoint,
provider_type=provider_type,
endpoint_sig=provider_api_format,
)
upstream_is_stream = resolve_upstream_is_stream(
client_is_stream=False,
policy=upstream_policy,
)
mapped_model = candidate.mapping_matched_model if candidate else None
if not mapped_model:
mapped_model = await handler._get_mapped_model(
source_model=model,
provider_id=str(provider.id),
api_format=provider_api_format,
)
provider_request_body = dict(original_request_body)
if mapped_model:
provider_request_body = handler.apply_mapped_model(provider_request_body, mapped_model)
if upstream_is_stream:
provider_request_body["stream"] = True
provider_request_body = maybe_patch_request_with_prompt_cache_key(
provider_request_body,
provider_api_format=provider_api_format,
provider_type=provider_type,
base_url=getattr(endpoint, "base_url", None),
user_api_key_id=str(getattr(api_key, "id", "") or ""),
request_headers=context.original_headers,
)
prompt_cache_key = str(provider_request_body.get("prompt_cache_key") or "").strip() or None
auth_info = await get_provider_auth(endpoint, key)
if auth_info is not None:
auth_header, auth_value = auth_info.as_tuple()
decrypted_auth_config = auth_info.decrypted_auth_config
else:
auth_header, auth_type = get_auth_config_for_endpoint(provider_api_format)
decrypted_key = crypto_service.decrypt(key.api_key)
auth_value = f"Bearer {decrypted_key}" if auth_type == "bearer" else decrypted_key
decrypted_auth_config = None
extra_headers = build_upstream_extra_headers(
provider_type=provider_type,
endpoint_sig=provider_api_format,
request_body=provider_request_body,
original_headers=context.original_headers,
decrypted_auth_config=decrypted_auth_config,
)
effective_proxy = resolve_effective_proxy(provider.proxy, getattr(key, "proxy", None))
proxy_info = await resolve_proxy_info_async(effective_proxy)
delegate_cfg = await resolve_delegate_config_async(effective_proxy)
is_tunnel_delegate = bool(delegate_cfg and delegate_cfg.get("tunnel"))
effective_proxy_for_contract = effective_proxy
if not effective_proxy_for_contract or not effective_proxy_for_contract.get("enabled", True):
effective_proxy_for_contract = await get_system_proxy_config_async()
proxy_url: str | None = None
if effective_proxy_for_contract and not is_tunnel_delegate:
proxy_url = await build_proxy_url_async(effective_proxy_for_contract)
proxy_snapshot = ExecutionProxySnapshot.from_proxy_info(
proxy_info,
proxy_url=proxy_url,
mode_override="tunnel" if is_tunnel_delegate else None,
node_id_override=(
str(delegate_cfg.get("node_id") or "").strip() or None if is_tunnel_delegate else None
),
)
request_timeout = provider.request_timeout or config.http_request_timeout
timeouts = ExecutionPlanTimeouts(
connect_ms=int(config.http_connect_timeout * 1000),
read_ms=int(config.http_read_timeout * 1000),
write_ms=int(config.http_write_timeout * 1000),
pool_ms=int(config.http_pool_timeout * 1000),
total_ms=int(request_timeout * 1000),
)
return GatewayExecutionDecisionResponse(
action="executor_sync_decision",
decision_kind="openai_chat_sync",
request_id=str(context.request_id),
candidate_id=str(
getattr(candidate, "request_candidate_id", "") or getattr(candidate, "id", "") or ""
)
or None,
provider_name=str(provider.name),
provider_id=str(provider.id),
endpoint_id=str(endpoint.id),
key_id=str(key.id),
upstream_base_url=str(getattr(endpoint, "base_url", "") or "").strip(),
auth_header=str(auth_header or "").strip() or "authorization",
auth_value=str(auth_value or "").strip(),
provider_api_format=provider_api_format,
client_api_format=client_api_format,
model_name=str(model or "unknown"),
mapped_model=str(mapped_model or "").strip() or None,
prompt_cache_key=prompt_cache_key,
extra_headers={str(k): str(v) for k, v in (extra_headers or {}).items()},
content_type="application/json",
proxy=gateway_module._serialize_gateway_sync_proxy(proxy_snapshot),
timeouts=gateway_module._serialize_gateway_sync_timeouts(timeouts),
upstream_is_stream=upstream_is_stream or None,
report_kind=(
"openai_chat_sync_finalize" if upstream_is_stream else "openai_chat_sync_success"
),
report_context={
"user_id": str(user.id),
"api_key_id": str(api_key.id),
"request_id": str(context.request_id),
"candidate_id": str(
getattr(candidate, "request_candidate_id", "") or getattr(candidate, "id", "") or ""
)
or None,
"model": str(model or "unknown"),
"provider_name": str(provider.name or "unknown"),
"provider_id": str(provider.id),
"endpoint_id": str(endpoint.id),
"key_id": str(key.id),
"provider_api_format": provider_api_format,
"client_api_format": client_api_format,
"mapped_model": str(mapped_model or "").strip() or None,
"original_headers": dict(context.original_headers),
"original_request_body": original_request_body,
"proxy_info": proxy_info,
"has_envelope": False,
"needs_conversion": False,
},
auth_context=auth_context,
)
async def _build_chat_stream_decision(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
expected_api_format: str,
decision_kind: str,
report_kind: str,
) -> GatewayExecutionDecisionResponse | None:
from src.api.handlers.base.chat_adapter_base import ChatAdapterBase
from src.config.settings import config
from src.core.api_format.headers import set_accept_if_absent
from src.services.provider.auth import get_provider_auth
from src.services.provider.transport import build_provider_url
from src.services.proxy_node.resolver import (
build_proxy_url_async,
get_system_proxy_config_async,
resolve_delegate_config_async,
resolve_effective_proxy,
resolve_proxy_info_async,
)
from src.services.request.execution_runtime_plan import (
ExecutionPlan,
ExecutionPlanTimeouts,
ExecutionProxySnapshot,
build_execution_plan_body,
is_remote_execution_runtime_contract_eligible,
)
if str(payload.method or "").strip().upper() != "POST":
return None
adapter, path_params = gateway_module._resolve_gateway_sync_adapter(decision, payload.path)
if not isinstance(adapter, ChatAdapterBase):
return None
if not gateway_module._is_stream_request_payload(payload.body_json, path_params):
return None
user, api_key = gateway_module._load_gateway_auth_models(db, auth_context)
pipeline = gateway_module.get_pipeline()
await pipeline._check_user_rate_limit(request, db, user, api_key)
context = gateway_module._build_gateway_request_context(
request=request,
payload=payload,
db=db,
user=user,
api_key=api_key,
adapter=adapter,
path_params=path_params,
balance_remaining=auth_context.balance_remaining,
)
authorize_result = adapter.authorize(context)
if hasattr(authorize_result, "__await__"):
await authorize_result
original_request_body = await context.ensure_json_body_async()
if context.path_params:
original_request_body = adapter._merge_path_params(
original_request_body,
context.path_params,
)
request_obj = adapter._validate_request_body(original_request_body, context.path_params)
if isinstance(request_obj, JSONResponse):
return None
handler = adapter._create_handler(
db=db,
user=user,
api_key=api_key,
request_id=context.request_id,
client_ip=context.client_ip,
user_agent=context.user_agent,
start_time=context.start_time,
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
)
converted_request = await handler._convert_request(request_obj)
model = getattr(converted_request, "model", original_request_body.get("model", "unknown"))
api_format = str(handler.allowed_api_formats[0] or "").strip().lower()
if api_format != expected_api_format:
return None
capability_requirements = handler._resolve_capability_requirements(
model_name=model,
request_headers=context.original_headers,
request_body=original_request_body,
)
preferred_key_ids = await handler._resolve_preferred_key_ids(
model_name=model,
request_body=original_request_body,
)
candidate = await gateway_module._select_gateway_direct_candidate(
db=db,
redis_client=getattr(handler, "redis", None),
api_format=api_format,
model_name=str(model or "unknown"),
user_api_key=api_key,
request_id=context.request_id,
is_stream=True,
capability_requirements=capability_requirements or None,
preferred_key_ids=preferred_key_ids or None,
request_body=original_request_body,
)
if candidate is None:
return None
provider = candidate.provider
endpoint = candidate.endpoint
key = candidate.key
provider_api_format = str(endpoint.api_format or api_format or "").strip().lower()
prep = await handler._prepare_provider_request(
model=str(model or "unknown"),
provider=provider,
endpoint=endpoint,
key=key,
working_request_body=dict(original_request_body),
original_headers=context.original_headers,
client_api_format=api_format,
provider_api_format=provider_api_format,
candidate=candidate,
client_is_stream=True,
)
provider_api_format = str(prep.provider_api_format or "").strip().lower()
client_api_format = str(prep.client_api_format or "").strip().lower()
if (
not prep.upstream_is_stream
or gateway_module._stream_executor_requires_python_rewrite(
envelope=prep.envelope,
needs_conversion=bool(prep.needs_conversion),
provider_api_format=provider_api_format,
client_api_format=client_api_format,
)
or client_api_format != expected_api_format
):
return None
if provider_api_format != expected_api_format and not (
expected_api_format == "openai:chat"
and provider_api_format in {"claude:chat", "gemini:chat"}
):
return None
auth_info = prep.auth_info or await get_provider_auth(endpoint, key)
provider_payload, provider_headers = handler._request_builder.build(
prep.request_body,
context.original_headers,
endpoint,
key,
is_stream=prep.upstream_is_stream,
extra_headers=prep.extra_headers if prep.extra_headers else None,
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
envelope=prep.envelope,
provider_api_format=prep.provider_api_format,
)
if not isinstance(provider_payload, dict):
return None
if prep.upstream_is_stream:
set_accept_if_absent(provider_headers)
provider_request_headers = {
str(k).lower(): str(v)
for k, v in dict(provider_headers or {}).items()
if str(k).strip() and str(v).strip()
}
provider_request_body = dict(provider_payload)
auth_header, auth_value = gateway_module._extract_gateway_upstream_auth(
provider_request_headers,
provider_api_format=provider_api_format,
key=key,
)
upstream_url = build_provider_url(
endpoint,
query_params=context.query_params,
path_params={"model": prep.url_model},
is_stream=prep.upstream_is_stream,
key=key,
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
)
effective_proxy = resolve_effective_proxy(provider.proxy, getattr(key, "proxy", None))
proxy_info = await resolve_proxy_info_async(effective_proxy)
delegate_cfg = await resolve_delegate_config_async(effective_proxy)
effective_proxy_for_contract = effective_proxy
if not effective_proxy_for_contract or not effective_proxy_for_contract.get("enabled", True):
effective_proxy_for_contract = await get_system_proxy_config_async()
is_tunnel_delegate = bool(delegate_cfg and delegate_cfg.get("tunnel"))
proxy_url: str | None = None
if effective_proxy_for_contract and not is_tunnel_delegate:
proxy_url = await build_proxy_url_async(effective_proxy_for_contract)
proxy_snapshot = ExecutionProxySnapshot.from_proxy_info(
proxy_info,
proxy_url=proxy_url,
mode_override="tunnel" if is_tunnel_delegate else None,
node_id_override=(
str(delegate_cfg.get("node_id") or "").strip() or None if is_tunnel_delegate else None
),
)
timeouts = ExecutionPlanTimeouts(
connect_ms=int(config.http_connect_timeout * 1000),
read_ms=int(config.http_read_timeout * 1000),
write_ms=int(config.http_write_timeout * 1000),
pool_ms=int(config.http_pool_timeout * 1000),
total_ms=None,
)
contract = ExecutionPlan(
request_id=str(context.request_id or ""),
candidate_id=str(
getattr(candidate, "request_candidate_id", "") or getattr(candidate, "id", "") or ""
)
or None,
provider_name=str(provider.name),
provider_id=str(provider.id),
endpoint_id=str(endpoint.id),
key_id=str(key.id),
method="POST",
url=str(upstream_url),
headers=dict(provider_request_headers),
body=build_execution_plan_body(
provider_request_body,
content_type=str(provider_request_headers.get("content-type") or "").strip() or None,
),
stream=True,
provider_api_format=provider_api_format,
client_api_format=str(prep.client_api_format or api_format),
model_name=str(model or ""),
content_type=str(provider_request_headers.get("content-type") or "").strip() or None,
content_encoding=context.client_content_encoding,
proxy=proxy_snapshot,
tls_profile=prep.tls_profile,
timeouts=timeouts,
)
if not is_remote_execution_runtime_contract_eligible(contract):
return None
prompt_cache_key = str(provider_request_body.get("prompt_cache_key") or "").strip() or None
has_envelope = prep.envelope is not None
needs_conversion = bool(prep.needs_conversion)
decision_extra_headers: dict[str, str] = {}
decision_provider_request_headers = provider_request_headers or None
decision_provider_request_body = provider_request_body
report_context = {
"user_id": str(user.id),
"api_key_id": str(api_key.id),
"request_id": str(context.request_id),
"candidate_id": str(contract.candidate_id or "") or None,
"model": str(model or "unknown"),
"provider_name": str(contract.provider_name or provider.name or "unknown"),
"provider_id": str(contract.provider_id or provider.id or ""),
"endpoint_id": str(contract.endpoint_id or endpoint.id or ""),
"key_id": str(contract.key_id or key.id or ""),
"provider_api_format": provider_api_format,
"client_api_format": str(prep.client_api_format or api_format),
"mapped_model": str(prep.mapped_model or "").strip() or None,
"original_headers": dict(context.original_headers),
"original_request_body": original_request_body,
"proxy_info": proxy_info,
"has_envelope": has_envelope,
"envelope_name": gateway_module._gateway_report_context_envelope_name(prep.envelope),
"needs_conversion": needs_conversion,
}
report_context["provider_request_headers"] = provider_request_headers
report_context["provider_request_body"] = provider_request_body
return GatewayExecutionDecisionResponse(
action="executor_stream_decision",
decision_kind=decision_kind,
request_id=str(context.request_id),
candidate_id=str(contract.candidate_id or "") or None,
provider_name=str(contract.provider_name or provider.name or ""),
provider_id=str(contract.provider_id or provider.id or ""),
endpoint_id=str(contract.endpoint_id or endpoint.id or ""),
key_id=str(contract.key_id or key.id or ""),
upstream_base_url=str(getattr(endpoint, "base_url", "") or "").strip(),
upstream_url=str(upstream_url),
auth_header=auth_header,
auth_value=str(auth_value or "").strip(),
provider_api_format=provider_api_format,
client_api_format=str(prep.client_api_format or api_format),
model_name=str(model or "unknown"),
mapped_model=str(prep.mapped_model or "").strip() or None,
prompt_cache_key=prompt_cache_key,
extra_headers=decision_extra_headers,
provider_request_headers=decision_provider_request_headers,
provider_request_body=decision_provider_request_body,
content_type=(
str(provider_request_headers.get("content-type") or "").strip() or "application/json"
),
proxy=gateway_module._serialize_gateway_sync_proxy(proxy_snapshot),
tls_profile=str(prep.tls_profile or "").strip() or None,
timeouts=gateway_module._serialize_gateway_sync_timeouts(timeouts),
report_kind=report_kind,
report_context=report_context,
auth_context=auth_context,
)
async def _build_openai_chat_stream_decision(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionDecisionResponse | None:
return await _build_chat_stream_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
expected_api_format="openai:chat",
decision_kind="openai_chat_stream",
report_kind="openai_chat_stream_success",
)
async def _build_claude_chat_stream_decision(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionDecisionResponse | None:
return await _build_chat_stream_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
expected_api_format="claude:chat",
decision_kind="claude_chat_stream",
report_kind="claude_chat_stream_success",
)
async def _build_gemini_chat_stream_decision(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionDecisionResponse | None:
return await _build_chat_stream_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
expected_api_format="gemini:chat",
decision_kind="gemini_chat_stream",
report_kind="gemini_chat_stream_success",
)
async def _build_claude_chat_sync_decision(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionDecisionResponse | None:
return await _build_chat_sync_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
expected_api_format="claude:chat",
decision_kind="claude_chat_sync",
report_kind="claude_chat_sync_success",
finalize_kind="claude_chat_sync_finalize",
)
async def _build_gemini_chat_sync_decision(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionDecisionResponse | None:
return await _build_chat_sync_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
expected_api_format="gemini:chat",
decision_kind="gemini_chat_sync",
report_kind="gemini_chat_sync_success",
finalize_kind="gemini_chat_sync_finalize",
)
@@ -1,591 +0,0 @@
from __future__ import annotations
import base64
import time
import uuid
from dataclasses import asdict
from datetime import datetime, timezone
from typing import Any
from fastapi import BackgroundTasks, HTTPException, Request
from fastapi.responses import JSONResponse, Response
from sqlalchemy.orm import Session
from src.api.base.context import ApiRequestContext
from src.api.base.pipeline import get_pipeline
from src.core.api_format.headers import extract_client_api_key_for_endpoint_with_query
from src.core.api_format.metadata import get_auth_config_for_endpoint
from src.core.crypto import crypto_service
from src.core.http_compression import normalize_content_encoding
from src.core.logger import logger
from src.database import create_session, get_db
from src.models.database import ApiKey, RequestCandidate, User
from src.services.auth.service import AuthService
from src.utils.async_utils import safe_create_task
from .common import ensure_loopback
from .gateway_contract import (
_GEMINI_FILES_DOWNLOAD_ROUTE_RE,
_GEMINI_FILES_RESOURCE_ROUTE_RE,
_GEMINI_MODEL_OPERATION_CANCEL_RE,
_GEMINI_OPERATION_CANCEL_RE,
_GEMINI_SYNC_ROUTE_RE,
_GEMINI_VIDEO_CREATE_ROUTE_RE,
_GEMINI_VIDEO_MODEL_OPERATION_ANY_RE,
_OPENAI_VIDEO_CANCEL_ROUTE_RE,
_OPENAI_VIDEO_CONTENT_ROUTE_RE,
_OPENAI_VIDEO_REMIX_ROUTE_RE,
_OPENAI_VIDEO_TASK_ROUTE_RE,
CONTROL_ACTION_HEADER,
CONTROL_ACTION_PROXY_PUBLIC,
CONTROL_EXECUTED_HEADER,
GatewayAuthContext,
GatewayExecuteRequest,
GatewayExecutionDecisionResponse,
GatewayExecutionPlanResponse,
GatewayResolveRequest,
GatewayRouteDecision,
GatewayStreamReportRequest,
GatewaySyncReportRequest,
classify_gateway_route,
)
class _GatewayProxy:
def __getattr__(self, name: str) -> Any:
from . import gateway as gateway_module
return getattr(gateway_module, name)
gateway_module = _GatewayProxy()
async def _build_chat_sync_plan(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
expected_api_format: str,
plan_kind: str,
report_kind: str,
finalize_kind: str,
) -> GatewayExecutionPlanResponse | None:
from src.api.handlers.base.chat_adapter_base import ChatAdapterBase
from src.api.handlers.base.chat_sync_executor import ChatSyncExecutor
from src.services.task.request_state import MutableRequestBodyState
if str(payload.method or "").strip().upper() != "POST":
return None
adapter, path_params = gateway_module._resolve_gateway_sync_adapter(decision, payload.path)
if not isinstance(adapter, ChatAdapterBase):
return None
if gateway_module._is_stream_request_payload(payload.body_json, path_params):
return None
user, api_key = gateway_module._load_gateway_auth_models(db, auth_context)
pipeline = gateway_module.get_pipeline()
await pipeline._check_user_rate_limit(request, db, user, api_key)
context = gateway_module._build_gateway_request_context(
request=request,
payload=payload,
db=db,
user=user,
api_key=api_key,
adapter=adapter,
path_params=path_params,
balance_remaining=auth_context.balance_remaining,
)
authorize_result = adapter.authorize(context)
if hasattr(authorize_result, "__await__"):
await authorize_result
original_request_body = await context.ensure_json_body_async()
if context.path_params:
original_request_body = adapter._merge_path_params(
original_request_body,
context.path_params,
)
request_obj = adapter._validate_request_body(original_request_body, context.path_params)
if isinstance(request_obj, JSONResponse):
return None
handler = adapter._create_handler(
db=db,
user=user,
api_key=api_key,
request_id=context.request_id,
client_ip=context.client_ip,
user_agent=context.user_agent,
start_time=context.start_time,
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
)
await handler._convert_request(request_obj)
model = handler.extract_model_from_request(
original_request_body,
context.path_params,
)
api_format = handler.allowed_api_formats[0]
capability_requirements = handler._resolve_capability_requirements(
model_name=model,
request_headers=context.original_headers,
request_body=original_request_body,
)
preferred_key_ids = await handler._resolve_preferred_key_ids(
model_name=model,
request_body=original_request_body,
)
candidate = await gateway_module._select_gateway_direct_candidate(
db=db,
redis_client=getattr(handler, "redis", None),
api_format=str(api_format),
model_name=str(model or "unknown"),
user_api_key=api_key,
request_id=context.request_id,
is_stream=False,
capability_requirements=capability_requirements or None,
preferred_key_ids=preferred_key_ids or None,
request_body=original_request_body,
)
if candidate is None:
return None
sync_executor = ChatSyncExecutor(handler)
prepared_plan = await sync_executor._build_sync_execution_plan(
candidate.provider,
candidate.endpoint,
candidate.key,
candidate,
model=str(model or "unknown"),
api_format=api_format,
original_headers=context.original_headers,
request_state=MutableRequestBodyState(original_request_body),
query_params=context.query_params,
client_content_encoding=context.client_content_encoding,
)
contract_provider_api_format = (
str(prepared_plan.contract.provider_api_format or "").strip().lower()
)
contract_client_api_format = str(prepared_plan.contract.client_api_format or "").strip().lower()
if not prepared_plan.remote_eligible or contract_client_api_format != expected_api_format:
return None
selected_report_kind = (
finalize_kind
if (
prepared_plan.upstream_is_stream
or prepared_plan.needs_conversion
or prepared_plan.envelope is not None
or contract_provider_api_format != expected_api_format
)
else report_kind
)
return GatewayExecutionPlanResponse(
action="executor_sync",
plan_kind=plan_kind,
plan=prepared_plan.contract.to_payload(),
report_kind=selected_report_kind,
report_context={
"user_id": str(user.id),
"api_key_id": str(api_key.id),
"request_id": str(context.request_id),
"model": str(model or "unknown"),
"provider_name": str(prepared_plan.contract.provider_name or "unknown"),
"provider_id": str(prepared_plan.contract.provider_id or ""),
"endpoint_id": str(prepared_plan.contract.endpoint_id or ""),
"key_id": str(prepared_plan.contract.key_id or ""),
"provider_api_format": str(prepared_plan.contract.provider_api_format or ""),
"client_api_format": str(prepared_plan.contract.client_api_format or ""),
"mapped_model": sync_executor._ctx.mapped_model_result,
"original_headers": dict(context.original_headers),
"original_request_body": original_request_body,
"provider_request_headers": dict(prepared_plan.headers),
"provider_request_body": prepared_plan.payload,
"proxy_info": prepared_plan.proxy_info,
"has_envelope": prepared_plan.envelope is not None,
"envelope_name": gateway_module._gateway_report_context_envelope_name(
prepared_plan.envelope
),
"needs_conversion": bool(prepared_plan.needs_conversion),
},
)
async def _build_openai_chat_sync_plan(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionPlanResponse | None:
return await _build_chat_sync_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
expected_api_format="openai:chat",
plan_kind="openai_chat_sync",
report_kind="openai_chat_sync_success",
finalize_kind="openai_chat_sync_finalize",
)
async def _build_chat_stream_plan(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
expected_api_format: str,
plan_kind: str,
report_kind: str,
) -> GatewayExecutionPlanResponse | None:
from src.api.handlers.base.chat_adapter_base import ChatAdapterBase
from src.config.settings import config
from src.core.api_format.headers import set_accept_if_absent
from src.services.provider.auth import get_provider_auth
from src.services.provider.transport import build_provider_url
from src.services.proxy_node.resolver import (
build_proxy_url_async,
get_system_proxy_config_async,
resolve_delegate_config_async,
resolve_effective_proxy,
resolve_proxy_info_async,
)
from src.services.request.execution_runtime_plan import (
ExecutionPlan,
ExecutionPlanTimeouts,
ExecutionProxySnapshot,
build_execution_plan_body,
is_remote_execution_runtime_contract_eligible,
)
if str(payload.method or "").strip().upper() != "POST":
return None
adapter, path_params = gateway_module._resolve_gateway_sync_adapter(decision, payload.path)
if not isinstance(adapter, ChatAdapterBase):
return None
if not gateway_module._is_stream_request_payload(payload.body_json, path_params):
return None
user, api_key = gateway_module._load_gateway_auth_models(db, auth_context)
pipeline = gateway_module.get_pipeline()
await pipeline._check_user_rate_limit(request, db, user, api_key)
context = gateway_module._build_gateway_request_context(
request=request,
payload=payload,
db=db,
user=user,
api_key=api_key,
adapter=adapter,
path_params=path_params,
balance_remaining=auth_context.balance_remaining,
)
authorize_result = adapter.authorize(context)
if hasattr(authorize_result, "__await__"):
await authorize_result
original_request_body = await context.ensure_json_body_async()
if context.path_params:
original_request_body = adapter._merge_path_params(
original_request_body,
context.path_params,
)
request_obj = adapter._validate_request_body(original_request_body, context.path_params)
if isinstance(request_obj, JSONResponse):
return None
handler = adapter._create_handler(
db=db,
user=user,
api_key=api_key,
request_id=context.request_id,
client_ip=context.client_ip,
user_agent=context.user_agent,
start_time=context.start_time,
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
)
converted_request = await handler._convert_request(request_obj)
model = getattr(converted_request, "model", original_request_body.get("model", "unknown"))
api_format = handler.allowed_api_formats[0]
capability_requirements = handler._resolve_capability_requirements(
model_name=model,
request_headers=context.original_headers,
request_body=original_request_body,
)
preferred_key_ids = await handler._resolve_preferred_key_ids(
model_name=model,
request_body=original_request_body,
)
candidate = await gateway_module._select_gateway_direct_candidate(
db=db,
redis_client=getattr(handler, "redis", None),
api_format=str(api_format),
model_name=str(model or "unknown"),
user_api_key=api_key,
request_id=context.request_id,
is_stream=True,
capability_requirements=capability_requirements or None,
preferred_key_ids=preferred_key_ids or None,
request_body=original_request_body,
)
if candidate is None:
return None
provider = candidate.provider
endpoint = candidate.endpoint
key = candidate.key
provider_api_format = str(endpoint.api_format or api_format or "")
prep = await handler._prepare_provider_request(
model=str(model or "unknown"),
provider=provider,
endpoint=endpoint,
key=key,
working_request_body=dict(original_request_body),
original_headers=context.original_headers,
client_api_format=str(api_format),
provider_api_format=provider_api_format,
candidate=candidate,
client_is_stream=True,
)
provider_api_format = str(prep.provider_api_format or "")
auth_info = prep.auth_info or await get_provider_auth(endpoint, key)
provider_payload, provider_headers = handler._request_builder.build(
prep.request_body,
context.original_headers,
endpoint,
key,
is_stream=prep.upstream_is_stream,
extra_headers=prep.extra_headers if prep.extra_headers else None,
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
envelope=prep.envelope,
provider_api_format=prep.provider_api_format,
)
if prep.upstream_is_stream:
set_accept_if_absent(provider_headers)
upstream_url = build_provider_url(
endpoint,
query_params=context.query_params,
path_params={"model": prep.url_model},
is_stream=prep.upstream_is_stream,
key=key,
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
)
effective_proxy = resolve_effective_proxy(provider.proxy, getattr(key, "proxy", None))
proxy_info = await resolve_proxy_info_async(effective_proxy)
delegate_cfg = await resolve_delegate_config_async(effective_proxy)
effective_proxy_for_contract = effective_proxy
if not effective_proxy_for_contract or not effective_proxy_for_contract.get("enabled", True):
effective_proxy_for_contract = await get_system_proxy_config_async()
is_tunnel_delegate = bool(delegate_cfg and delegate_cfg.get("tunnel"))
proxy_url: str | None = None
if effective_proxy_for_contract and not is_tunnel_delegate:
proxy_url = await build_proxy_url_async(effective_proxy_for_contract)
contract = ExecutionPlan(
request_id=str(context.request_id or ""),
candidate_id=str(
getattr(candidate, "request_candidate_id", "") or getattr(candidate, "id", "") or ""
)
or None,
provider_name=str(provider.name),
provider_id=str(provider.id),
endpoint_id=str(endpoint.id),
key_id=str(key.id),
method="POST",
url=upstream_url,
headers=dict(provider_headers),
body=build_execution_plan_body(
provider_payload,
content_type=str(provider_headers.get("content-type") or "").strip() or None,
),
stream=True,
provider_api_format=provider_api_format,
client_api_format=str(api_format),
model_name=str(model or ""),
content_type=str(provider_headers.get("content-type") or "").strip() or None,
content_encoding=context.client_content_encoding,
proxy=ExecutionProxySnapshot.from_proxy_info(
proxy_info,
proxy_url=proxy_url,
mode_override="tunnel" if is_tunnel_delegate else None,
node_id_override=(
str(delegate_cfg.get("node_id") or "").strip() or None
if is_tunnel_delegate
else None
),
),
tls_profile=prep.tls_profile,
timeouts=ExecutionPlanTimeouts(
connect_ms=int(config.http_connect_timeout * 1000),
read_ms=int(config.http_read_timeout * 1000),
write_ms=int(config.http_write_timeout * 1000),
pool_ms=int(config.http_pool_timeout * 1000),
total_ms=None,
),
)
if (
not prep.upstream_is_stream
or gateway_module._stream_executor_requires_python_rewrite(
envelope=prep.envelope,
needs_conversion=bool(prep.needs_conversion),
provider_api_format=provider_api_format,
client_api_format=str(contract.client_api_format or ""),
)
or not is_remote_execution_runtime_contract_eligible(contract)
or provider_api_format.strip().lower() != expected_api_format
or str(contract.client_api_format or "").strip().lower() != expected_api_format
):
return None
return GatewayExecutionPlanResponse(
action="executor_stream",
plan_kind=plan_kind,
plan=contract.to_payload(),
report_kind=report_kind,
report_context={
"user_id": str(user.id),
"api_key_id": str(api_key.id),
"request_id": str(context.request_id),
"model": str(model or "unknown"),
"provider_name": str(contract.provider_name or "unknown"),
"provider_id": str(contract.provider_id or ""),
"endpoint_id": str(contract.endpoint_id or ""),
"key_id": str(contract.key_id or ""),
"candidate_id": str(contract.candidate_id or ""),
"provider_api_format": str(contract.provider_api_format or ""),
"client_api_format": str(contract.client_api_format or ""),
"mapped_model": prep.mapped_model,
"original_headers": dict(context.original_headers),
"original_request_body": original_request_body,
"provider_request_headers": dict(provider_headers),
"provider_request_body": provider_payload,
"proxy_info": proxy_info,
"has_envelope": prep.envelope is not None,
"envelope_name": gateway_module._gateway_report_context_envelope_name(prep.envelope),
"needs_conversion": bool(prep.needs_conversion),
},
)
async def _build_openai_chat_stream_plan(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionPlanResponse | None:
return await _build_chat_stream_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
expected_api_format="openai:chat",
plan_kind="openai_chat_stream",
report_kind="openai_chat_stream_success",
)
async def _build_claude_chat_stream_plan(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionPlanResponse | None:
return await _build_chat_stream_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
expected_api_format="claude:chat",
plan_kind="claude_chat_stream",
report_kind="claude_chat_stream_success",
)
async def _build_gemini_chat_stream_plan(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionPlanResponse | None:
return await _build_chat_stream_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
expected_api_format="gemini:chat",
plan_kind="gemini_chat_stream",
report_kind="gemini_chat_stream_success",
)
async def _build_claude_chat_sync_plan(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionPlanResponse | None:
return await _build_chat_sync_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
expected_api_format="claude:chat",
plan_kind="claude_chat_sync",
report_kind="claude_chat_sync_success",
finalize_kind="claude_chat_sync_finalize",
)
async def _build_gemini_chat_sync_plan(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionPlanResponse | None:
return await _build_chat_sync_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
expected_api_format="gemini:chat",
plan_kind="gemini_chat_sync",
report_kind="gemini_chat_sync_success",
finalize_kind="gemini_chat_sync_finalize",
)
@@ -1,43 +0,0 @@
"""Compatibility re-export layer for gateway CLI builders."""
from __future__ import annotations
from .gateway_cli_decision import (
_build_claude_cli_stream_decision,
_build_claude_cli_sync_decision,
_build_cli_stream_decision,
_build_cli_sync_decision,
_build_gemini_cli_stream_decision,
_build_gemini_cli_sync_decision,
_build_openai_cli_stream_decision,
_build_openai_cli_sync_decision,
)
from .gateway_cli_plan import (
_build_claude_cli_stream_plan,
_build_claude_cli_sync_plan,
_build_cli_stream_plan,
_build_cli_sync_plan,
_build_gemini_cli_stream_plan,
_build_gemini_cli_sync_plan,
_build_openai_cli_stream_plan,
_build_openai_cli_sync_plan,
)
__all__ = [
"_build_cli_sync_decision",
"_build_cli_stream_decision",
"_build_openai_cli_sync_decision",
"_build_openai_cli_stream_decision",
"_build_claude_cli_sync_decision",
"_build_claude_cli_stream_decision",
"_build_gemini_cli_sync_decision",
"_build_gemini_cli_stream_decision",
"_build_cli_stream_plan",
"_build_cli_sync_plan",
"_build_openai_cli_stream_plan",
"_build_claude_cli_stream_plan",
"_build_gemini_cli_stream_plan",
"_build_openai_cli_sync_plan",
"_build_claude_cli_sync_plan",
"_build_gemini_cli_sync_plan",
]
@@ -1,945 +0,0 @@
from __future__ import annotations
import base64
import time
import uuid
from dataclasses import asdict
from datetime import datetime, timezone
from typing import Any
from fastapi import BackgroundTasks, HTTPException, Request
from fastapi.responses import JSONResponse, Response
from sqlalchemy.orm import Session
from src.api.base.context import ApiRequestContext
from src.api.base.pipeline import get_pipeline
from src.core.api_format.headers import extract_client_api_key_for_endpoint_with_query
from src.core.api_format.metadata import get_auth_config_for_endpoint
from src.core.crypto import crypto_service
from src.core.http_compression import normalize_content_encoding
from src.core.logger import logger
from src.database import create_session, get_db
from src.models.database import ApiKey, RequestCandidate, User
from src.services.auth.service import AuthService
from src.utils.async_utils import safe_create_task
from .common import ensure_loopback
from .gateway_contract import (
_GEMINI_FILES_DOWNLOAD_ROUTE_RE,
_GEMINI_FILES_RESOURCE_ROUTE_RE,
_GEMINI_MODEL_OPERATION_CANCEL_RE,
_GEMINI_OPERATION_CANCEL_RE,
_GEMINI_SYNC_ROUTE_RE,
_GEMINI_VIDEO_CREATE_ROUTE_RE,
_GEMINI_VIDEO_MODEL_OPERATION_ANY_RE,
_OPENAI_VIDEO_CANCEL_ROUTE_RE,
_OPENAI_VIDEO_CONTENT_ROUTE_RE,
_OPENAI_VIDEO_REMIX_ROUTE_RE,
_OPENAI_VIDEO_TASK_ROUTE_RE,
CONTROL_ACTION_HEADER,
CONTROL_ACTION_PROXY_PUBLIC,
CONTROL_EXECUTED_HEADER,
GatewayAuthContext,
GatewayExecuteRequest,
GatewayExecutionDecisionResponse,
GatewayExecutionPlanResponse,
GatewayResolveRequest,
GatewayRouteDecision,
GatewayStreamReportRequest,
GatewaySyncReportRequest,
classify_gateway_route,
)
class _GatewayProxy:
def __getattr__(self, name: str) -> Any:
from . import gateway as gateway_module
return getattr(gateway_module, name)
gateway_module = _GatewayProxy()
async def _build_cli_sync_decision(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
expected_api_format: str,
decision_kind: str,
report_kind: str,
finalize_kind: str,
) -> GatewayExecutionDecisionResponse | None:
from src.api.handlers.base.cli_adapter_base import CliAdapterBase
from src.config.settings import config
from src.services.proxy_node.resolver import (
build_proxy_url_async,
get_system_proxy_config_async,
resolve_delegate_config_async,
resolve_effective_proxy,
resolve_proxy_info_async,
)
from src.services.request.execution_runtime_plan import (
ExecutionPlanTimeouts,
ExecutionProxySnapshot,
should_bypass_remote_execution_runtime_url,
)
if str(payload.method or "").strip().upper() != "POST":
return None
adapter, path_params = gateway_module._resolve_gateway_sync_adapter(decision, payload.path)
if not isinstance(adapter, CliAdapterBase):
return None
if gateway_module._is_stream_request_payload(payload.body_json, path_params):
return None
if not isinstance(payload.body_json, dict):
return None
user, api_key = gateway_module._load_gateway_auth_models(db, auth_context)
pipeline = gateway_module.get_pipeline()
await pipeline._check_user_rate_limit(request, db, user, api_key)
context = gateway_module._build_gateway_request_context(
request=request,
payload=payload,
db=db,
user=user,
api_key=api_key,
adapter=adapter,
path_params=path_params,
balance_remaining=auth_context.balance_remaining,
)
authorize_result = adapter.authorize(context)
if hasattr(authorize_result, "__await__"):
await authorize_result
original_request_body = await context.ensure_json_body_async()
if context.path_params:
original_request_body = adapter._merge_path_params(
original_request_body,
context.path_params,
)
handler = adapter.HANDLER_CLASS(
db=db,
user=user,
api_key=api_key,
request_id=context.request_id,
client_ip=context.client_ip,
user_agent=context.user_agent,
start_time=context.start_time,
allowed_api_formats=adapter.allowed_api_formats,
adapter_detector=adapter.detect_capability_requirements,
perf_metrics=context.extra.get("perf"),
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
)
model = handler.extract_model_from_request(original_request_body, context.path_params)
client_api_format = str(handler.primary_api_format or "").strip().lower()
if client_api_format != expected_api_format:
return None
capability_requirements = handler._resolve_capability_requirements(
model_name=model,
request_headers=context.original_headers,
request_body=original_request_body,
)
preferred_key_ids = await handler._resolve_preferred_key_ids(
model_name=model,
request_body=original_request_body,
)
candidate = await gateway_module._select_gateway_direct_candidate(
db=db,
redis_client=getattr(handler, "redis", None),
api_format=client_api_format,
model_name=str(model or "unknown"),
user_api_key=api_key,
request_id=context.request_id,
is_stream=False,
capability_requirements=capability_requirements or None,
preferred_key_ids=preferred_key_ids or None,
request_body=original_request_body,
)
if candidate is None:
return None
provider = candidate.provider
endpoint = candidate.endpoint
key = candidate.key
provider_api_format = str(endpoint.api_format or client_api_format or "").strip().lower()
mapped_model = candidate.mapping_matched_model if candidate else None
if not mapped_model:
mapped_model = await handler._get_mapped_model(
source_model=str(model or "unknown"),
provider_id=str(provider.id),
)
upstream_request = await handler._build_upstream_request(
provider=provider,
endpoint=endpoint,
key=key,
request_body=dict(original_request_body),
original_headers=context.original_headers,
query_params=context.query_params,
client_api_format=client_api_format,
provider_api_format=provider_api_format,
fallback_model=str(model or "unknown"),
mapped_model=mapped_model,
client_is_stream=False,
needs_conversion=bool(getattr(candidate, "needs_conversion", False)),
output_limit=getattr(candidate, "output_limit", None),
)
if should_bypass_remote_execution_runtime_url(
upstream_request.url,
provider_api_format=provider_api_format,
client_api_format=client_api_format,
):
return None
upstream_is_stream = bool(upstream_request.upstream_is_stream)
provider_request_body = dict(upstream_request.payload or {})
provider_request_headers = {
str(k).lower(): str(v)
for k, v in dict(upstream_request.headers or {}).items()
if str(k).strip() and str(v).strip()
}
auth_header, auth_value = gateway_module._extract_gateway_upstream_auth(
provider_request_headers,
provider_api_format=provider_api_format,
key=key,
)
prompt_cache_key = str(provider_request_body.get("prompt_cache_key") or "").strip() or None
effective_proxy = resolve_effective_proxy(provider.proxy, getattr(key, "proxy", None))
proxy_info = await resolve_proxy_info_async(effective_proxy)
delegate_cfg = await resolve_delegate_config_async(effective_proxy)
is_tunnel_delegate = bool(delegate_cfg and delegate_cfg.get("tunnel"))
effective_proxy_for_contract = effective_proxy
if not effective_proxy_for_contract or not effective_proxy_for_contract.get("enabled", True):
effective_proxy_for_contract = await get_system_proxy_config_async()
proxy_url: str | None = None
if effective_proxy_for_contract and not is_tunnel_delegate:
proxy_url = await build_proxy_url_async(effective_proxy_for_contract)
proxy_snapshot = ExecutionProxySnapshot.from_proxy_info(
proxy_info,
proxy_url=proxy_url,
mode_override="tunnel" if is_tunnel_delegate else None,
node_id_override=(
str(delegate_cfg.get("node_id") or "").strip() or None if is_tunnel_delegate else None
),
)
request_timeout = provider.request_timeout or config.http_request_timeout
timeouts = ExecutionPlanTimeouts(
connect_ms=int(config.http_connect_timeout * 1000),
read_ms=int(config.http_read_timeout * 1000),
write_ms=int(config.http_write_timeout * 1000),
pool_ms=int(config.http_pool_timeout * 1000),
total_ms=int(request_timeout * 1000),
)
requires_finalize = (
upstream_is_stream
or bool(getattr(candidate, "needs_conversion", False))
or provider_api_format != client_api_format
or upstream_request.envelope is not None
)
has_envelope = upstream_request.envelope is not None
needs_conversion = bool(getattr(candidate, "needs_conversion", False))
decision_extra_headers: dict[str, str] = {}
decision_provider_request_headers = provider_request_headers or None
decision_provider_request_body = provider_request_body
report_context = {
"user_id": str(user.id),
"api_key_id": str(api_key.id),
"request_id": str(context.request_id),
"candidate_id": str(
getattr(candidate, "request_candidate_id", "") or getattr(candidate, "id", "") or ""
)
or None,
"model": str(model or "unknown"),
"provider_name": str(provider.name or "unknown"),
"provider_id": str(provider.id),
"endpoint_id": str(endpoint.id),
"key_id": str(key.id),
"provider_api_format": provider_api_format,
"client_api_format": client_api_format,
"mapped_model": str(mapped_model or "").strip() or None,
"original_headers": dict(context.original_headers),
"original_request_body": original_request_body,
"proxy_info": proxy_info,
"has_envelope": has_envelope,
"envelope_name": gateway_module._gateway_report_context_envelope_name(
upstream_request.envelope
),
"needs_conversion": needs_conversion,
}
report_context["provider_request_headers"] = provider_request_headers
report_context["provider_request_body"] = provider_request_body
return GatewayExecutionDecisionResponse(
action="executor_sync_decision",
decision_kind=decision_kind,
request_id=str(context.request_id),
candidate_id=str(
getattr(candidate, "request_candidate_id", "") or getattr(candidate, "id", "") or ""
)
or None,
provider_name=str(provider.name),
provider_id=str(provider.id),
endpoint_id=str(endpoint.id),
key_id=str(key.id),
upstream_base_url=str(getattr(endpoint, "base_url", "") or "").strip(),
upstream_url=str(upstream_request.url or "").strip() or None,
auth_header=str(auth_header or "").strip() or "authorization",
auth_value=str(auth_value or "").strip(),
provider_api_format=provider_api_format,
client_api_format=client_api_format,
model_name=str(model or "unknown"),
mapped_model=str(mapped_model or "").strip() or None,
prompt_cache_key=prompt_cache_key,
extra_headers=decision_extra_headers,
provider_request_headers=decision_provider_request_headers,
provider_request_body=decision_provider_request_body,
content_type=(
str(provider_request_headers.get("content-type") or "").strip() or "application/json"
),
proxy=gateway_module._serialize_gateway_sync_proxy(proxy_snapshot),
tls_profile=str(upstream_request.tls_profile or "").strip() or None,
timeouts=gateway_module._serialize_gateway_sync_timeouts(timeouts),
upstream_is_stream=upstream_is_stream or None,
report_kind=finalize_kind if requires_finalize else report_kind,
report_context=report_context,
auth_context=auth_context,
)
async def _build_cli_stream_decision(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
expected_api_format: str,
decision_kind: str,
report_kind: str,
) -> GatewayExecutionDecisionResponse | None:
from src.api.handlers.base.cli_adapter_base import CliAdapterBase
from src.config.settings import config
from src.services.proxy_node.resolver import (
build_proxy_url_async,
get_system_proxy_config_async,
resolve_delegate_config_async,
resolve_effective_proxy,
resolve_proxy_info_async,
)
from src.services.request.execution_runtime_plan import (
ExecutionPlanTimeouts,
ExecutionProxySnapshot,
should_bypass_remote_execution_runtime_url,
)
if str(payload.method or "").strip().upper() != "POST":
return None
adapter, path_params = gateway_module._resolve_gateway_sync_adapter(decision, payload.path)
if not isinstance(adapter, CliAdapterBase):
return None
if not gateway_module._is_stream_request_payload(payload.body_json, path_params):
return None
if not isinstance(payload.body_json, dict):
return None
user, api_key = gateway_module._load_gateway_auth_models(db, auth_context)
pipeline = gateway_module.get_pipeline()
await pipeline._check_user_rate_limit(request, db, user, api_key)
context = gateway_module._build_gateway_request_context(
request=request,
payload=payload,
db=db,
user=user,
api_key=api_key,
adapter=adapter,
path_params=path_params,
balance_remaining=auth_context.balance_remaining,
)
authorize_result = adapter.authorize(context)
if hasattr(authorize_result, "__await__"):
await authorize_result
original_request_body = await context.ensure_json_body_async()
if context.path_params:
original_request_body = adapter._merge_path_params(
original_request_body,
context.path_params,
)
handler = adapter.HANDLER_CLASS(
db=db,
user=user,
api_key=api_key,
request_id=context.request_id,
client_ip=context.client_ip,
user_agent=context.user_agent,
start_time=context.start_time,
allowed_api_formats=adapter.allowed_api_formats,
adapter_detector=adapter.detect_capability_requirements,
perf_metrics=context.extra.get("perf"),
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
)
model = handler.extract_model_from_request(original_request_body, context.path_params)
client_api_format = str(handler.primary_api_format or "").strip().lower()
if client_api_format != expected_api_format:
return None
capability_requirements = handler._resolve_capability_requirements(
model_name=model,
request_headers=context.original_headers,
request_body=original_request_body,
)
preferred_key_ids = await handler._resolve_preferred_key_ids(
model_name=model,
request_body=original_request_body,
)
candidate = await gateway_module._select_gateway_direct_candidate(
db=db,
redis_client=getattr(handler, "redis", None),
api_format=client_api_format,
model_name=str(model or "unknown"),
user_api_key=api_key,
request_id=context.request_id,
is_stream=True,
capability_requirements=capability_requirements or None,
preferred_key_ids=preferred_key_ids or None,
request_body=original_request_body,
)
if candidate is None:
return None
provider = candidate.provider
endpoint = candidate.endpoint
key = candidate.key
provider_api_format = str(endpoint.api_format or client_api_format or "").strip().lower()
if provider_api_format != expected_api_format and not (
expected_api_format in {"openai:cli", "openai:compact"}
and provider_api_format in {"claude:cli", "gemini:cli"}
):
return None
mapped_model = candidate.mapping_matched_model if candidate else None
if not mapped_model:
mapped_model = await handler._get_mapped_model(
source_model=str(model or "unknown"),
provider_id=str(provider.id),
)
upstream_request = await handler._build_upstream_request(
provider=provider,
endpoint=endpoint,
key=key,
request_body=dict(original_request_body),
original_headers=context.original_headers,
query_params=context.query_params,
client_api_format=client_api_format,
provider_api_format=provider_api_format,
fallback_model=str(model or "unknown"),
mapped_model=mapped_model,
client_is_stream=True,
needs_conversion=bool(getattr(candidate, "needs_conversion", False)),
output_limit=getattr(candidate, "output_limit", None),
)
if should_bypass_remote_execution_runtime_url(
upstream_request.url,
provider_api_format=provider_api_format,
client_api_format=client_api_format,
):
return None
if not bool(upstream_request.upstream_is_stream):
return None
if gateway_module._stream_executor_requires_python_rewrite(
envelope=upstream_request.envelope,
needs_conversion=bool(getattr(candidate, "needs_conversion", False)),
provider_api_format=provider_api_format,
client_api_format=client_api_format,
):
return None
provider_request_body = dict(upstream_request.payload or {})
provider_request_headers = {
str(k).lower(): str(v)
for k, v in dict(upstream_request.headers or {}).items()
if str(k).strip() and str(v).strip()
}
auth_header, auth_value = gateway_module._extract_gateway_upstream_auth(
provider_request_headers,
provider_api_format=provider_api_format,
key=key,
)
prompt_cache_key = str(provider_request_body.get("prompt_cache_key") or "").strip() or None
effective_proxy = resolve_effective_proxy(provider.proxy, getattr(key, "proxy", None))
proxy_info = await resolve_proxy_info_async(effective_proxy)
delegate_cfg = await resolve_delegate_config_async(effective_proxy)
is_tunnel_delegate = bool(delegate_cfg and delegate_cfg.get("tunnel"))
effective_proxy_for_contract = effective_proxy
if not effective_proxy_for_contract or not effective_proxy_for_contract.get("enabled", True):
effective_proxy_for_contract = await get_system_proxy_config_async()
proxy_url: str | None = None
if effective_proxy_for_contract and not is_tunnel_delegate:
proxy_url = await build_proxy_url_async(effective_proxy_for_contract)
proxy_snapshot = ExecutionProxySnapshot.from_proxy_info(
proxy_info,
proxy_url=proxy_url,
mode_override="tunnel" if is_tunnel_delegate else None,
node_id_override=(
str(delegate_cfg.get("node_id") or "").strip() or None if is_tunnel_delegate else None
),
)
timeouts = ExecutionPlanTimeouts(
connect_ms=int(config.http_connect_timeout * 1000),
read_ms=int(config.http_read_timeout * 1000),
write_ms=int(config.http_write_timeout * 1000),
pool_ms=int(config.http_pool_timeout * 1000),
total_ms=None,
)
has_envelope = upstream_request.envelope is not None
needs_conversion = bool(getattr(candidate, "needs_conversion", False))
decision_extra_headers: dict[str, str] = {}
decision_provider_request_headers = provider_request_headers or None
decision_provider_request_body = provider_request_body
report_context = {
"user_id": str(user.id),
"api_key_id": str(api_key.id),
"request_id": str(context.request_id),
"candidate_id": str(
getattr(candidate, "request_candidate_id", "") or getattr(candidate, "id", "") or ""
)
or None,
"model": str(model or "unknown"),
"provider_name": str(provider.name or "unknown"),
"provider_id": str(provider.id),
"endpoint_id": str(endpoint.id),
"key_id": str(key.id),
"provider_api_format": provider_api_format,
"client_api_format": client_api_format,
"mapped_model": str(mapped_model or "").strip() or None,
"original_headers": dict(context.original_headers),
"original_request_body": original_request_body,
"proxy_info": proxy_info,
"has_envelope": has_envelope,
"envelope_name": gateway_module._gateway_report_context_envelope_name(
upstream_request.envelope
),
"needs_conversion": needs_conversion,
}
report_context["provider_request_headers"] = provider_request_headers
report_context["provider_request_body"] = provider_request_body
return GatewayExecutionDecisionResponse(
action="executor_stream_decision",
decision_kind=decision_kind,
request_id=str(context.request_id),
candidate_id=str(
getattr(candidate, "request_candidate_id", "") or getattr(candidate, "id", "") or ""
)
or None,
provider_name=str(provider.name),
provider_id=str(provider.id),
endpoint_id=str(endpoint.id),
key_id=str(key.id),
upstream_base_url=str(getattr(endpoint, "base_url", "") or "").strip(),
upstream_url=str(upstream_request.url or "").strip() or None,
auth_header=str(auth_header or "").strip() or "authorization",
auth_value=str(auth_value or "").strip(),
provider_api_format=provider_api_format,
client_api_format=client_api_format,
model_name=str(model or "unknown"),
mapped_model=str(mapped_model or "").strip() or None,
prompt_cache_key=prompt_cache_key,
extra_headers=decision_extra_headers,
provider_request_headers=decision_provider_request_headers,
provider_request_body=decision_provider_request_body,
content_type=(
str(provider_request_headers.get("content-type") or "").strip() or "application/json"
),
proxy=gateway_module._serialize_gateway_sync_proxy(proxy_snapshot),
tls_profile=str(upstream_request.tls_profile or "").strip() or None,
timeouts=gateway_module._serialize_gateway_sync_timeouts(timeouts),
report_kind=report_kind,
report_context=report_context,
auth_context=auth_context,
)
async def _build_openai_cli_sync_decision(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionDecisionResponse | None:
from src.api.handlers.base.cli_adapter_base import CliAdapterBase
from src.config.settings import config
from src.services.proxy_node.resolver import (
build_proxy_url_async,
get_system_proxy_config_async,
resolve_delegate_config_async,
resolve_effective_proxy,
resolve_proxy_info_async,
)
from src.services.request.execution_runtime_plan import (
ExecutionPlanTimeouts,
ExecutionProxySnapshot,
)
if str(payload.method or "").strip().upper() != "POST":
return None
adapter, path_params = gateway_module._resolve_gateway_sync_adapter(decision, payload.path)
if not isinstance(adapter, CliAdapterBase):
return None
if gateway_module._is_stream_request_payload(payload.body_json, path_params):
return None
if not isinstance(payload.body_json, dict):
return None
user, api_key = gateway_module._load_gateway_auth_models(db, auth_context)
pipeline = gateway_module.get_pipeline()
await pipeline._check_user_rate_limit(request, db, user, api_key)
context = gateway_module._build_gateway_request_context(
request=request,
payload=payload,
db=db,
user=user,
api_key=api_key,
adapter=adapter,
path_params=path_params,
balance_remaining=auth_context.balance_remaining,
)
authorize_result = adapter.authorize(context)
if hasattr(authorize_result, "__await__"):
await authorize_result
original_request_body = await context.ensure_json_body_async()
if context.path_params:
original_request_body = adapter._merge_path_params(
original_request_body,
context.path_params,
)
if decision.route_kind == "compact":
original_request_body.pop("stream", None)
handler = adapter.HANDLER_CLASS(
db=db,
user=user,
api_key=api_key,
request_id=context.request_id,
client_ip=context.client_ip,
user_agent=context.user_agent,
start_time=context.start_time,
allowed_api_formats=adapter.allowed_api_formats,
adapter_detector=adapter.detect_capability_requirements,
perf_metrics=context.extra.get("perf"),
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
)
model = handler.extract_model_from_request(original_request_body, context.path_params)
client_api_format = str(handler.primary_api_format or "").strip().lower()
if client_api_format not in {"openai:cli", "openai:compact"}:
return None
capability_requirements = handler._resolve_capability_requirements(
model_name=model,
request_headers=context.original_headers,
request_body=original_request_body,
)
preferred_key_ids = await handler._resolve_preferred_key_ids(
model_name=model,
request_body=original_request_body,
)
candidate = await gateway_module._select_gateway_direct_candidate(
db=db,
redis_client=getattr(handler, "redis", None),
api_format=client_api_format,
model_name=str(model or "unknown"),
user_api_key=api_key,
request_id=context.request_id,
is_stream=False,
capability_requirements=capability_requirements or None,
preferred_key_ids=preferred_key_ids or None,
request_body=original_request_body,
)
if candidate is None:
return None
provider = candidate.provider
endpoint = candidate.endpoint
key = candidate.key
provider_api_format = str(endpoint.api_format or client_api_format or "").strip().lower()
mapped_model = candidate.mapping_matched_model if candidate else None
if not mapped_model:
mapped_model = await handler._get_mapped_model(
source_model=str(model or "unknown"),
provider_id=str(provider.id),
)
upstream_request = await handler._build_upstream_request(
provider=provider,
endpoint=endpoint,
key=key,
request_body=dict(original_request_body),
original_headers=context.original_headers,
query_params=context.query_params,
client_api_format=client_api_format,
provider_api_format=provider_api_format,
fallback_model=str(model or "unknown"),
mapped_model=mapped_model,
client_is_stream=False,
needs_conversion=bool(getattr(candidate, "needs_conversion", False)),
output_limit=getattr(candidate, "output_limit", None),
)
upstream_is_stream = bool(upstream_request.upstream_is_stream)
provider_request_body = dict(upstream_request.payload or {})
provider_request_headers = {
str(k).lower(): str(v)
for k, v in dict(upstream_request.headers or {}).items()
if str(k).strip() and str(v).strip()
}
prompt_cache_key = str(provider_request_body.get("prompt_cache_key") or "").strip() or None
auth_header = ""
auth_value = ""
for header_name, header_value in provider_request_headers.items():
if header_name in {"authorization", "x-api-key", "x-goog-api-key"}:
auth_header = header_name
auth_value = header_value
break
if not auth_header:
auth_header, auth_type = get_auth_config_for_endpoint(provider_api_format)
decrypted_key = crypto_service.decrypt(key.api_key)
auth_value = f"Bearer {decrypted_key}" if auth_type == "bearer" else decrypted_key
effective_proxy = resolve_effective_proxy(provider.proxy, getattr(key, "proxy", None))
proxy_info = await resolve_proxy_info_async(effective_proxy)
delegate_cfg = await resolve_delegate_config_async(effective_proxy)
is_tunnel_delegate = bool(delegate_cfg and delegate_cfg.get("tunnel"))
effective_proxy_for_contract = effective_proxy
if not effective_proxy_for_contract or not effective_proxy_for_contract.get("enabled", True):
effective_proxy_for_contract = await get_system_proxy_config_async()
proxy_url: str | None = None
if effective_proxy_for_contract and not is_tunnel_delegate:
proxy_url = await build_proxy_url_async(effective_proxy_for_contract)
proxy_snapshot = ExecutionProxySnapshot.from_proxy_info(
proxy_info,
proxy_url=proxy_url,
mode_override="tunnel" if is_tunnel_delegate else None,
node_id_override=(
str(delegate_cfg.get("node_id") or "").strip() or None if is_tunnel_delegate else None
),
)
request_timeout = provider.request_timeout or config.http_request_timeout
timeouts = ExecutionPlanTimeouts(
connect_ms=int(config.http_connect_timeout * 1000),
read_ms=int(config.http_read_timeout * 1000),
write_ms=int(config.http_write_timeout * 1000),
pool_ms=int(config.http_pool_timeout * 1000),
total_ms=int(request_timeout * 1000),
)
requires_finalize = (
upstream_is_stream
or bool(getattr(candidate, "needs_conversion", False))
or provider_api_format != client_api_format
or upstream_request.envelope is not None
)
has_envelope = upstream_request.envelope is not None
needs_conversion = bool(getattr(candidate, "needs_conversion", False))
decision_extra_headers: dict[str, str] = {}
decision_provider_request_headers = provider_request_headers or None
decision_provider_request_body = provider_request_body
report_context = {
"user_id": str(user.id),
"api_key_id": str(api_key.id),
"request_id": str(context.request_id),
"candidate_id": str(
getattr(candidate, "request_candidate_id", "") or getattr(candidate, "id", "") or ""
)
or None,
"model": str(model or "unknown"),
"provider_name": str(provider.name or "unknown"),
"provider_id": str(provider.id),
"endpoint_id": str(endpoint.id),
"key_id": str(key.id),
"provider_api_format": provider_api_format,
"client_api_format": client_api_format,
"mapped_model": str(mapped_model or "").strip() or None,
"original_headers": dict(context.original_headers),
"original_request_body": original_request_body,
"proxy_info": proxy_info,
"has_envelope": has_envelope,
"envelope_name": gateway_module._gateway_report_context_envelope_name(
upstream_request.envelope
),
"needs_conversion": needs_conversion,
}
report_context["provider_request_headers"] = provider_request_headers
report_context["provider_request_body"] = provider_request_body
is_compact = decision.route_kind == "compact"
return GatewayExecutionDecisionResponse(
action="executor_sync_decision",
decision_kind="openai_compact_sync" if is_compact else "openai_cli_sync",
request_id=str(context.request_id),
candidate_id=str(
getattr(candidate, "request_candidate_id", "") or getattr(candidate, "id", "") or ""
)
or None,
provider_name=str(provider.name),
provider_id=str(provider.id),
endpoint_id=str(endpoint.id),
key_id=str(key.id),
upstream_base_url=str(getattr(endpoint, "base_url", "") or "").strip(),
upstream_url=str(upstream_request.url or "").strip() or None,
auth_header=str(auth_header or "").strip() or "authorization",
auth_value=str(auth_value or "").strip(),
provider_api_format=provider_api_format,
client_api_format=client_api_format,
model_name=str(model or "unknown"),
mapped_model=str(mapped_model or "").strip() or None,
prompt_cache_key=prompt_cache_key,
extra_headers=decision_extra_headers,
provider_request_headers=decision_provider_request_headers,
provider_request_body=decision_provider_request_body,
content_type=(
str(provider_request_headers.get("content-type") or "").strip() or "application/json"
),
proxy=gateway_module._serialize_gateway_sync_proxy(proxy_snapshot),
tls_profile=str(upstream_request.tls_profile or "").strip() or None,
timeouts=gateway_module._serialize_gateway_sync_timeouts(timeouts),
upstream_is_stream=upstream_is_stream or None,
report_kind=(
("openai_compact_sync_finalize" if is_compact else "openai_cli_sync_finalize")
if requires_finalize
else "openai_cli_sync_success"
),
report_context=report_context,
auth_context=auth_context,
)
async def _build_openai_cli_stream_decision(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionDecisionResponse | None:
is_compact = decision.route_kind == "compact"
return await _build_cli_stream_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
expected_api_format="openai:compact" if is_compact else "openai:cli",
decision_kind="openai_compact_stream" if is_compact else "openai_cli_stream",
report_kind="openai_cli_stream_success",
)
async def _build_claude_cli_sync_decision(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionDecisionResponse | None:
return await _build_cli_sync_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
expected_api_format="claude:cli",
decision_kind="claude_cli_sync",
report_kind="claude_cli_sync_success",
finalize_kind="claude_cli_sync_finalize",
)
async def _build_claude_cli_stream_decision(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionDecisionResponse | None:
return await _build_cli_stream_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
expected_api_format="claude:cli",
decision_kind="claude_cli_stream",
report_kind="claude_cli_stream_success",
)
async def _build_gemini_cli_sync_decision(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionDecisionResponse | None:
return await _build_cli_sync_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
expected_api_format="gemini:cli",
decision_kind="gemini_cli_sync",
report_kind="gemini_cli_sync_success",
finalize_kind="gemini_cli_sync_finalize",
)
async def _build_gemini_cli_stream_decision(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionDecisionResponse | None:
return await _build_cli_stream_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
expected_api_format="gemini:cli",
decision_kind="gemini_cli_stream",
report_kind="gemini_cli_stream_success",
)
@@ -1,681 +0,0 @@
from __future__ import annotations
import base64
import time
import uuid
from dataclasses import asdict
from datetime import datetime, timezone
from typing import Any
from fastapi import BackgroundTasks, HTTPException, Request
from fastapi.responses import JSONResponse, Response
from sqlalchemy.orm import Session
from src.api.base.context import ApiRequestContext
from src.api.base.pipeline import get_pipeline
from src.core.api_format.headers import extract_client_api_key_for_endpoint_with_query
from src.core.api_format.metadata import get_auth_config_for_endpoint
from src.core.crypto import crypto_service
from src.core.http_compression import normalize_content_encoding
from src.core.logger import logger
from src.database import create_session, get_db
from src.models.database import ApiKey, RequestCandidate, User
from src.services.auth.service import AuthService
from src.utils.async_utils import safe_create_task
from .common import ensure_loopback
from .gateway_contract import (
_GEMINI_FILES_DOWNLOAD_ROUTE_RE,
_GEMINI_FILES_RESOURCE_ROUTE_RE,
_GEMINI_MODEL_OPERATION_CANCEL_RE,
_GEMINI_OPERATION_CANCEL_RE,
_GEMINI_SYNC_ROUTE_RE,
_GEMINI_VIDEO_CREATE_ROUTE_RE,
_GEMINI_VIDEO_MODEL_OPERATION_ANY_RE,
_OPENAI_VIDEO_CANCEL_ROUTE_RE,
_OPENAI_VIDEO_CONTENT_ROUTE_RE,
_OPENAI_VIDEO_REMIX_ROUTE_RE,
_OPENAI_VIDEO_TASK_ROUTE_RE,
CONTROL_ACTION_HEADER,
CONTROL_ACTION_PROXY_PUBLIC,
CONTROL_EXECUTED_HEADER,
GatewayAuthContext,
GatewayExecuteRequest,
GatewayExecutionDecisionResponse,
GatewayExecutionPlanResponse,
GatewayResolveRequest,
GatewayRouteDecision,
GatewayStreamReportRequest,
GatewaySyncReportRequest,
classify_gateway_route,
)
class _GatewayProxy:
def __getattr__(self, name: str) -> Any:
from . import gateway as gateway_module
return getattr(gateway_module, name)
gateway_module = _GatewayProxy()
async def _build_cli_stream_plan(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
expected_provider_formats: set[str],
plan_kind: str,
report_kind: str,
log_label: str,
) -> GatewayExecutionPlanResponse | None:
from loguru import logger
from src.api.handlers.base.cli_adapter_base import CliAdapterBase
from src.config.settings import config
from src.services.provider.transport import redact_url_for_log
from src.services.proxy_node.resolver import (
build_proxy_url_async,
resolve_delegate_config_async,
resolve_effective_proxy,
resolve_proxy_info_async,
)
from src.services.request.execution_runtime_plan import (
ExecutionPlan,
ExecutionPlanTimeouts,
ExecutionProxySnapshot,
build_execution_plan_body,
is_remote_execution_runtime_contract_eligible,
)
if str(payload.method or "").strip().upper() != "POST":
return None
adapter, path_params = gateway_module._resolve_gateway_sync_adapter(decision, payload.path)
if not isinstance(adapter, CliAdapterBase):
return None
if not gateway_module._is_stream_request_payload(payload.body_json, path_params):
return None
user, api_key = gateway_module._load_gateway_auth_models(db, auth_context)
pipeline = gateway_module.get_pipeline()
await pipeline._check_user_rate_limit(request, db, user, api_key)
context = gateway_module._build_gateway_request_context(
request=request,
payload=payload,
db=db,
user=user,
api_key=api_key,
adapter=adapter,
path_params=path_params,
balance_remaining=auth_context.balance_remaining,
)
authorize_result = adapter.authorize(context)
if hasattr(authorize_result, "__await__"):
await authorize_result
original_request_body = await context.ensure_json_body_async()
if context.path_params:
original_request_body = adapter._merge_path_params(
original_request_body,
context.path_params,
)
handler = adapter.HANDLER_CLASS(
db=db,
user=user,
api_key=api_key,
request_id=context.request_id,
client_ip=context.client_ip,
user_agent=context.user_agent,
start_time=context.start_time,
allowed_api_formats=adapter.allowed_api_formats,
adapter_detector=adapter.detect_capability_requirements,
perf_metrics=context.extra.get("perf"),
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
)
model = handler.extract_model_from_request(original_request_body, context.path_params)
client_api_format = handler.primary_api_format
capability_requirements = handler._resolve_capability_requirements(
model_name=model,
request_headers=context.original_headers,
request_body=original_request_body,
)
preferred_key_ids = await handler._resolve_preferred_key_ids(
model_name=model,
request_body=original_request_body,
)
candidate = await gateway_module._select_gateway_direct_candidate(
db=db,
redis_client=getattr(handler, "redis", None),
api_format=str(client_api_format),
model_name=str(model or "unknown"),
user_api_key=api_key,
request_id=context.request_id,
is_stream=True,
capability_requirements=capability_requirements or None,
preferred_key_ids=preferred_key_ids or None,
request_body=original_request_body,
)
if candidate is None:
return None
provider = candidate.provider
endpoint = candidate.endpoint
key = candidate.key
provider_api_format = str(endpoint.api_format or "")
mapped_model = candidate.mapping_matched_model if candidate else None
if not mapped_model:
mapped_model = await handler._get_mapped_model(
source_model=str(model or "unknown"),
provider_id=str(provider.id),
)
needs_conversion = bool(getattr(candidate, "needs_conversion", False))
upstream_request = await handler._build_upstream_request(
provider=provider,
endpoint=endpoint,
key=key,
request_body=dict(original_request_body),
original_headers=context.original_headers,
query_params=context.query_params,
client_api_format=str(client_api_format),
provider_api_format=provider_api_format,
fallback_model=str(model or "unknown"),
mapped_model=mapped_model,
client_is_stream=True,
needs_conversion=needs_conversion,
output_limit=candidate.output_limit if candidate else None,
)
effective_proxy = resolve_effective_proxy(provider.proxy, getattr(key, "proxy", None))
stream_proxy_info = await resolve_proxy_info_async(effective_proxy)
delegate_cfg = await resolve_delegate_config_async(effective_proxy)
is_tunnel_delegate = bool(delegate_cfg and delegate_cfg.get("tunnel"))
proxy_url: str | None = None
if effective_proxy and not is_tunnel_delegate:
proxy_url = await build_proxy_url_async(effective_proxy)
contract = ExecutionPlan(
request_id=str(context.request_id or ""),
candidate_id=str(
getattr(candidate, "request_candidate_id", "") or getattr(candidate, "id", "") or ""
)
or None,
provider_name=str(provider.name),
provider_id=str(provider.id),
endpoint_id=str(endpoint.id),
key_id=str(key.id),
method="POST",
url=upstream_request.url,
headers=dict(upstream_request.headers),
body=build_execution_plan_body(
upstream_request.payload,
content_type=str(upstream_request.headers.get("content-type") or "").strip() or None,
),
stream=True,
provider_api_format=provider_api_format,
client_api_format=str(client_api_format),
model_name=str(model or ""),
content_type=str(upstream_request.headers.get("content-type") or "").strip() or None,
content_encoding=context.client_content_encoding,
proxy=ExecutionProxySnapshot.from_proxy_info(
stream_proxy_info,
proxy_url=proxy_url,
mode_override="tunnel" if is_tunnel_delegate else None,
node_id_override=(
str(delegate_cfg.get("node_id") or "").strip() or None
if is_tunnel_delegate
else None
),
),
tls_profile=upstream_request.tls_profile,
timeouts=ExecutionPlanTimeouts(
connect_ms=int(config.http_connect_timeout * 1000),
read_ms=int(config.http_read_timeout * 1000),
write_ms=int(config.http_write_timeout * 1000),
pool_ms=int(config.http_pool_timeout * 1000),
total_ms=None,
),
)
if (
not upstream_request.upstream_is_stream
or gateway_module._stream_executor_requires_python_rewrite(
envelope=upstream_request.envelope,
needs_conversion=needs_conversion,
provider_api_format=provider_api_format,
client_api_format=str(client_api_format),
)
or not is_remote_execution_runtime_contract_eligible(contract)
or provider_api_format != str(client_api_format)
or provider_api_format not in expected_provider_formats
):
return None
logger.debug(
"[gateway] {} stream direct executor candidate accepted: path={} provider={} url={}",
log_label,
payload.path,
provider.name,
redact_url_for_log(upstream_request.url),
)
return GatewayExecutionPlanResponse(
action="executor_stream",
plan_kind=plan_kind,
plan=contract.to_payload(),
report_kind=report_kind,
report_context={
"user_id": str(user.id),
"api_key_id": str(api_key.id),
"request_id": str(context.request_id),
"model": str(model or "unknown"),
"provider_name": str(contract.provider_name or "unknown"),
"provider_id": str(contract.provider_id or ""),
"endpoint_id": str(contract.endpoint_id or ""),
"key_id": str(contract.key_id or ""),
"candidate_id": str(contract.candidate_id or ""),
"provider_api_format": str(contract.provider_api_format or ""),
"client_api_format": str(contract.client_api_format or ""),
"mapped_model": mapped_model,
"original_headers": dict(context.original_headers),
"original_request_body": original_request_body,
"provider_request_headers": dict(upstream_request.headers),
"provider_request_body": upstream_request.payload,
"proxy_info": stream_proxy_info,
"has_envelope": upstream_request.envelope is not None,
"envelope_name": gateway_module._gateway_report_context_envelope_name(
upstream_request.envelope
),
"needs_conversion": bool(needs_conversion),
},
)
async def _build_cli_sync_plan(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
expected_client_formats: set[str],
plan_kind: str,
report_kind: str,
finalize_kind: str,
log_label: str,
) -> GatewayExecutionPlanResponse | None:
from loguru import logger
from src.api.handlers.base.cli_adapter_base import CliAdapterBase
from src.config.settings import config
from src.services.provider.transport import redact_url_for_log
from src.services.proxy_node.resolver import (
build_proxy_url_async,
resolve_delegate_config_async,
resolve_effective_proxy,
resolve_proxy_info_async,
)
from src.services.request.execution_runtime_plan import (
ExecutionPlan,
ExecutionPlanTimeouts,
ExecutionProxySnapshot,
build_execution_plan_body,
is_remote_execution_runtime_contract_eligible,
)
if str(payload.method or "").strip().upper() != "POST":
return None
adapter, path_params = gateway_module._resolve_gateway_sync_adapter(decision, payload.path)
if not isinstance(adapter, CliAdapterBase):
return None
if gateway_module._is_stream_request_payload(payload.body_json, path_params):
return None
user, api_key = gateway_module._load_gateway_auth_models(db, auth_context)
pipeline = gateway_module.get_pipeline()
await pipeline._check_user_rate_limit(request, db, user, api_key)
context = gateway_module._build_gateway_request_context(
request=request,
payload=payload,
db=db,
user=user,
api_key=api_key,
adapter=adapter,
path_params=path_params,
balance_remaining=auth_context.balance_remaining,
)
authorize_result = adapter.authorize(context)
if hasattr(authorize_result, "__await__"):
await authorize_result
original_request_body = await context.ensure_json_body_async()
if context.path_params:
original_request_body = adapter._merge_path_params(
original_request_body,
context.path_params,
)
if decision.route_kind == "compact":
original_request_body.pop("stream", None)
handler = adapter.HANDLER_CLASS(
db=db,
user=user,
api_key=api_key,
request_id=context.request_id,
client_ip=context.client_ip,
user_agent=context.user_agent,
start_time=context.start_time,
allowed_api_formats=adapter.allowed_api_formats,
adapter_detector=adapter.detect_capability_requirements,
perf_metrics=context.extra.get("perf"),
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
)
model = handler.extract_model_from_request(original_request_body, context.path_params)
client_api_format = handler.primary_api_format
capability_requirements = handler._resolve_capability_requirements(
model_name=model,
request_headers=context.original_headers,
request_body=original_request_body,
)
preferred_key_ids = await handler._resolve_preferred_key_ids(
model_name=model,
request_body=original_request_body,
)
candidate = await gateway_module._select_gateway_direct_candidate(
db=db,
redis_client=getattr(handler, "redis", None),
api_format=str(client_api_format),
model_name=str(model or "unknown"),
user_api_key=api_key,
request_id=context.request_id,
is_stream=False,
capability_requirements=capability_requirements or None,
preferred_key_ids=preferred_key_ids or None,
request_body=original_request_body,
)
if candidate is None:
return None
provider = candidate.provider
endpoint = candidate.endpoint
key = candidate.key
provider_api_format = str(endpoint.api_format or "")
mapped_model = candidate.mapping_matched_model if candidate else None
if not mapped_model:
mapped_model = await handler._get_mapped_model(
source_model=str(model or "unknown"),
provider_id=str(provider.id),
)
needs_conversion = bool(getattr(candidate, "needs_conversion", False))
upstream_request = await handler._build_upstream_request(
provider=provider,
endpoint=endpoint,
key=key,
request_body=dict(original_request_body),
original_headers=context.original_headers,
query_params=context.query_params,
client_api_format=str(client_api_format),
provider_api_format=provider_api_format,
fallback_model=str(model or "unknown"),
mapped_model=mapped_model,
client_is_stream=False,
needs_conversion=needs_conversion,
output_limit=candidate.output_limit if candidate else None,
)
_effective_proxy = resolve_effective_proxy(provider.proxy, getattr(key, "proxy", None))
sync_proxy_info = await resolve_proxy_info_async(_effective_proxy)
delegate_cfg = await resolve_delegate_config_async(_effective_proxy)
is_tunnel_delegate = bool(delegate_cfg and delegate_cfg.get("tunnel"))
proxy_url: str | None = None
if _effective_proxy and not is_tunnel_delegate:
proxy_url = await build_proxy_url_async(_effective_proxy)
contract = ExecutionPlan(
request_id=str(context.request_id or ""),
candidate_id=str(
getattr(candidate, "request_candidate_id", "") or getattr(candidate, "id", "") or ""
)
or None,
provider_name=str(provider.name),
provider_id=str(provider.id),
endpoint_id=str(endpoint.id),
key_id=str(key.id),
method="POST",
url=upstream_request.url,
headers=dict(upstream_request.headers),
body=build_execution_plan_body(
upstream_request.payload,
content_type=str(upstream_request.headers.get("content-type") or "").strip() or None,
),
stream=upstream_request.upstream_is_stream,
provider_api_format=provider_api_format,
client_api_format=str(client_api_format),
model_name=str(model or ""),
content_type=str(upstream_request.headers.get("content-type") or "").strip() or None,
content_encoding=context.client_content_encoding,
proxy=ExecutionProxySnapshot.from_proxy_info(
sync_proxy_info,
proxy_url=proxy_url,
mode_override="tunnel" if is_tunnel_delegate else None,
node_id_override=(
str(delegate_cfg.get("node_id") or "").strip() or None
if is_tunnel_delegate
else None
),
),
tls_profile=upstream_request.tls_profile,
timeouts=ExecutionPlanTimeouts(
connect_ms=int(config.http_connect_timeout * 1000),
read_ms=int(config.http_read_timeout * 1000),
write_ms=int(config.http_write_timeout * 1000),
pool_ms=int(config.http_pool_timeout * 1000),
total_ms=int((provider.request_timeout or config.http_request_timeout) * 1000),
),
)
normalized_client_api_format = str(client_api_format)
if (
not is_remote_execution_runtime_contract_eligible(contract)
or normalized_client_api_format not in expected_client_formats
):
return None
selected_report_kind = (
finalize_kind
if (
upstream_request.upstream_is_stream
or needs_conversion
or upstream_request.envelope is not None
or provider_api_format != normalized_client_api_format
)
else report_kind
)
logger.debug(
"[gateway] {} sync direct executor candidate accepted: path={} provider={} url={}",
log_label,
payload.path,
provider.name,
redact_url_for_log(upstream_request.url),
)
return GatewayExecutionPlanResponse(
action="executor_sync",
plan_kind=plan_kind,
plan=contract.to_payload(),
report_kind=selected_report_kind,
report_context={
"user_id": str(user.id),
"api_key_id": str(api_key.id),
"request_id": str(context.request_id),
"model": str(model or "unknown"),
"provider_name": str(contract.provider_name or "unknown"),
"provider_id": str(contract.provider_id or ""),
"endpoint_id": str(contract.endpoint_id or ""),
"key_id": str(contract.key_id or ""),
"provider_api_format": str(contract.provider_api_format or ""),
"client_api_format": str(contract.client_api_format or ""),
"mapped_model": mapped_model,
"original_headers": dict(context.original_headers),
"original_request_body": original_request_body,
"provider_request_headers": dict(upstream_request.headers),
"provider_request_body": upstream_request.payload,
"proxy_info": sync_proxy_info,
"has_envelope": upstream_request.envelope is not None,
"envelope_name": gateway_module._gateway_report_context_envelope_name(
upstream_request.envelope
),
"needs_conversion": bool(needs_conversion),
},
)
async def _build_openai_cli_stream_plan(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionPlanResponse | None:
is_compact = decision.route_kind == "compact"
return await _build_cli_stream_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
expected_provider_formats={"openai:compact" if is_compact else "openai:cli"},
plan_kind="openai_compact_stream" if is_compact else "openai_cli_stream",
report_kind="openai_cli_stream_success",
log_label="openai compact" if is_compact else "openai cli",
)
async def _build_claude_cli_stream_plan(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionPlanResponse | None:
return await _build_cli_stream_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
expected_provider_formats={"claude:cli"},
plan_kind="claude_cli_stream",
report_kind="claude_cli_stream_success",
log_label="claude cli",
)
async def _build_gemini_cli_stream_plan(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionPlanResponse | None:
return await _build_cli_stream_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
expected_provider_formats={"gemini:cli"},
plan_kind="gemini_cli_stream",
report_kind="gemini_cli_stream_success",
log_label="gemini cli",
)
async def _build_openai_cli_sync_plan(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionPlanResponse | None:
return await _build_cli_sync_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
expected_client_formats={"openai:cli", "openai:compact"},
plan_kind="openai_compact_sync" if decision.route_kind == "compact" else "openai_cli_sync",
report_kind="openai_cli_sync_success",
finalize_kind=(
"openai_compact_sync_finalize"
if decision.route_kind == "compact"
else "openai_cli_sync_finalize"
),
log_label="openai cli",
)
async def _build_claude_cli_sync_plan(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionPlanResponse | None:
return await _build_cli_sync_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
expected_client_formats={"claude:cli"},
plan_kind="claude_cli_sync",
report_kind="claude_cli_sync_success",
finalize_kind="claude_cli_sync_finalize",
log_label="claude cli",
)
async def _build_gemini_cli_sync_plan(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionPlanResponse | None:
return await _build_cli_sync_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
expected_client_formats={"gemini:cli"},
plan_kind="gemini_cli_sync",
report_kind="gemini_cli_sync_success",
finalize_kind="gemini_cli_sync_finalize",
log_label="gemini cli",
)
@@ -1,277 +0,0 @@
from __future__ import annotations
import re
from typing import Any
from pydantic import BaseModel, Field
_GEMINI_MODEL_ROUTE_RE = re.compile(
r"^/(v1|v1beta)/models/[^/]+:(generateContent|streamGenerateContent|predictLongRunning)$"
)
_GEMINI_MODEL_OPERATION_CANCEL_RE = re.compile(r"^/v1beta/models/[^/]+/operations/[^/]+:cancel$")
_GEMINI_OPERATION_CANCEL_RE = re.compile(r"^/v1beta/operations/[^/]+:cancel$")
_GEMINI_MODEL_OPERATION_ROUTE_RE = re.compile(r"^/v1beta/models/[^/]+/operations/[^/]+$")
_GEMINI_FILES_ROUTE_RE = re.compile(r"^/v1beta/files(?:/.+)?$")
_GEMINI_FILES_DOWNLOAD_ROUTE_RE = re.compile(r"^/v1beta/files/(?P<file_id>[^/]+):download$")
_GEMINI_FILES_RESOURCE_ROUTE_RE = re.compile(r"^/v1beta/files/(?P<file_name>.+)$")
_GEMINI_SYNC_ROUTE_RE = re.compile(
r"^/(?P<version>v1|v1beta)/models/(?P<model>[^/]+):(?P<action>generateContent|streamGenerateContent)$"
)
_OPENAI_VIDEO_CANCEL_ROUTE_RE = re.compile(r"^/v1/videos/(?P<task_id>[^/]+)/cancel$")
_OPENAI_VIDEO_REMIX_ROUTE_RE = re.compile(r"^/v1/videos/(?P<task_id>[^/]+)/remix$")
_OPENAI_VIDEO_CONTENT_ROUTE_RE = re.compile(r"^/v1/videos/(?P<task_id>[^/]+)/content$")
_OPENAI_VIDEO_TASK_ROUTE_RE = re.compile(r"^/v1/videos/(?P<task_id>[^/]+)$")
_GEMINI_VIDEO_CREATE_ROUTE_RE = re.compile(r"^/v1beta/models/(?P<model>[^/]+):predictLongRunning$")
_GEMINI_VIDEO_MODEL_OPERATION_ANY_RE = re.compile(
r"^/v1beta/models/(?P<model>[^/]+)/operations/(?P<operation_id>[^/]+)(?P<cancel>:cancel)?$"
)
CONTROL_EXECUTED_HEADER = "x-aether-control-executed"
CONTROL_ACTION_HEADER = "x-aether-control-action"
CONTROL_ACTION_PROXY_PUBLIC = "proxy_public"
class GatewayResolveRequest(BaseModel):
trace_id: str | None = Field(None, max_length=128)
method: str = Field(..., min_length=1, max_length=16)
path: str = Field(..., min_length=1, max_length=2048)
query_string: str | None = Field(None, max_length=8192)
headers: dict[str, str] = Field(default_factory=dict)
has_body: bool = False
content_type: str | None = Field(None, max_length=512)
content_length: int | None = Field(None, ge=0)
class GatewayAuthContextRequest(BaseModel):
trace_id: str | None = Field(None, max_length=128)
query_string: str | None = Field(None, max_length=8192)
headers: dict[str, str] = Field(default_factory=dict)
auth_endpoint_signature: str = Field(..., min_length=1, max_length=128)
class GatewayRouteDecision(BaseModel):
action: str = "proxy_public"
route_class: str
public_path: str
public_query_string: str | None = None
route_family: str | None = None
route_kind: str | None = None
auth_endpoint_signature: str | None = None
executor_candidate: bool = False
auth_context: dict[str, Any] | None = None
class GatewayAuthContext(BaseModel):
user_id: str
api_key_id: str
balance_remaining: float | None = None
access_allowed: bool = True
class GatewayExecuteRequest(BaseModel):
trace_id: str | None = Field(None, max_length=128)
method: str = Field(..., min_length=1, max_length=16)
path: str = Field(..., min_length=1, max_length=2048)
query_string: str | None = Field(None, max_length=8192)
headers: dict[str, str] = Field(default_factory=dict)
body_json: dict[str, Any] = Field(default_factory=dict)
body_base64: str | None = None
auth_context: GatewayAuthContext | None = None
class GatewayExecutionPlanResponse(BaseModel):
action: str
plan_kind: str
plan: dict[str, Any]
report_kind: str | None = None
report_context: dict[str, Any] | None = None
auth_context: GatewayAuthContext | None = None
class GatewayExecutionDecisionResponse(BaseModel):
action: str
decision_kind: str
request_id: str
candidate_id: str | None = None
provider_name: str
provider_id: str
endpoint_id: str
key_id: str
upstream_base_url: str
upstream_url: str | None = None
provider_request_method: str | None = None
auth_header: str
auth_value: str
provider_api_format: str
client_api_format: str
model_name: str
mapped_model: str | None = None
prompt_cache_key: str | None = None
extra_headers: dict[str, str] = Field(default_factory=dict)
provider_request_headers: dict[str, str] | None = None
provider_request_body: dict[str, Any] | None = None
content_type: str | None = None
proxy: dict[str, Any] | None = None
tls_profile: str | None = None
timeouts: dict[str, Any] | None = None
upstream_is_stream: bool | None = None
report_kind: str | None = None
report_context: dict[str, Any] | None = None
auth_context: GatewayAuthContext | None = None
class GatewaySyncReportRequest(BaseModel):
trace_id: str | None = Field(None, max_length=128)
report_kind: str = Field(..., min_length=1, max_length=128)
report_context: dict[str, Any] = Field(default_factory=dict)
status_code: int = Field(..., ge=100, le=599)
headers: dict[str, str] = Field(default_factory=dict)
body_json: Any = None
client_body_json: Any = None
body_base64: str | None = None
telemetry: dict[str, Any] | None = None
class GatewayStreamReportRequest(BaseModel):
trace_id: str | None = Field(None, max_length=128)
report_kind: str = Field(..., min_length=1, max_length=128)
report_context: dict[str, Any] = Field(default_factory=dict)
status_code: int = Field(..., ge=100, le=599)
headers: dict[str, str] = Field(default_factory=dict)
body_base64: str | None = None
telemetry: dict[str, Any] | None = None
def classify_gateway_route(
method: str,
path: str,
headers: dict[str, str] | None = None,
) -> GatewayRouteDecision:
normalized_method = str(method or "").strip().upper() or "GET"
normalized_path = str(path or "").strip() or "/"
normalized_headers = {str(k).lower(): str(v) for k, v in (headers or {}).items()}
if not normalized_path.startswith("/"):
normalized_path = f"/{normalized_path}"
if normalized_method == "POST" and normalized_path == "/v1/chat/completions":
return _ai_route(
normalized_path,
family="openai",
kind="chat",
auth_endpoint_signature="openai:chat",
)
if normalized_method == "POST" and normalized_path in {
"/v1/responses",
"/v1/responses/compact",
}:
route_kind = "compact" if normalized_path.endswith("/compact") else "cli"
auth_endpoint_signature = "openai:compact" if route_kind == "compact" else "openai:cli"
return _ai_route(
normalized_path,
family="openai",
kind=route_kind,
auth_endpoint_signature=auth_endpoint_signature,
)
if normalized_method == "POST" and normalized_path == "/v1/messages":
is_claude_cli = _is_claude_cli_request(normalized_headers)
return _ai_route(
normalized_path,
family="claude",
kind="cli" if is_claude_cli else "chat",
auth_endpoint_signature="claude:cli" if is_claude_cli else "claude:chat",
)
if normalized_path.startswith("/v1/videos"):
return _ai_route(
normalized_path,
family="openai",
kind="video",
auth_endpoint_signature="openai:video",
)
if _GEMINI_MODEL_ROUTE_RE.match(normalized_path):
if normalized_path.endswith(":predictLongRunning"):
return _ai_route(
normalized_path,
family="gemini",
kind="video",
auth_endpoint_signature="gemini:video",
)
is_gemini_cli = _is_gemini_cli_request(normalized_headers)
return _ai_route(
normalized_path,
family="gemini",
kind="cli" if is_gemini_cli else "chat",
auth_endpoint_signature="gemini:cli" if is_gemini_cli else "gemini:chat",
)
if (
_GEMINI_MODEL_OPERATION_CANCEL_RE.match(normalized_path)
or _GEMINI_OPERATION_CANCEL_RE.match(normalized_path)
or _GEMINI_MODEL_OPERATION_ROUTE_RE.match(normalized_path)
or normalized_path == "/v1beta/operations"
or normalized_path.startswith("/v1beta/operations/")
):
return _ai_route(
normalized_path,
family="gemini",
kind="video",
auth_endpoint_signature="gemini:video",
)
if normalized_method == "POST" and normalized_path == "/upload/v1beta/files":
return _ai_route(
normalized_path,
family="gemini",
kind="files",
auth_endpoint_signature="gemini:chat",
)
if _GEMINI_FILES_ROUTE_RE.match(normalized_path):
return _ai_route(
normalized_path,
family="gemini",
kind="files",
auth_endpoint_signature="gemini:chat",
)
return GatewayRouteDecision(
route_class="passthrough",
public_path=normalized_path,
executor_candidate=False,
)
def _ai_route(
path: str,
*,
family: str,
kind: str,
auth_endpoint_signature: str,
) -> GatewayRouteDecision:
return GatewayRouteDecision(
route_class="ai_public",
public_path=path,
route_family=family,
route_kind=kind,
auth_endpoint_signature=auth_endpoint_signature,
executor_candidate=True,
)
def _is_claude_cli_request(headers: dict[str, str]) -> bool:
auth_header = str(headers.get("authorization") or "").strip().lower()
has_bearer = auth_header.startswith("bearer ")
has_api_key = bool(str(headers.get("x-api-key") or "").strip())
return has_bearer and not has_api_key
def _is_gemini_cli_request(headers: dict[str, str]) -> bool:
x_app = str(headers.get("x-app") or "").lower()
if "cli" in x_app:
return True
user_agent = str(headers.get("user-agent") or "").lower()
return "geminicli" in user_agent or "gemini-cli" in user_agent
@@ -1,853 +0,0 @@
from __future__ import annotations
from typing import Any
from urllib.parse import parse_qsl
from fastapi import Request
from sqlalchemy.orm import Session
from src.core.api_format.headers import extract_client_api_key_for_endpoint_with_query
from src.services.auth.service import AuthService
from .gateway_contract import (
_GEMINI_FILES_DOWNLOAD_ROUTE_RE,
_GEMINI_FILES_RESOURCE_ROUTE_RE,
_GEMINI_MODEL_OPERATION_CANCEL_RE,
_GEMINI_OPERATION_CANCEL_RE,
_GEMINI_VIDEO_CREATE_ROUTE_RE,
_OPENAI_VIDEO_CANCEL_ROUTE_RE,
_OPENAI_VIDEO_CONTENT_ROUTE_RE,
_OPENAI_VIDEO_REMIX_ROUTE_RE,
_OPENAI_VIDEO_TASK_ROUTE_RE,
GatewayAuthContext,
GatewayExecuteRequest,
GatewayExecutionDecisionResponse,
GatewayExecutionPlanResponse,
GatewayResolveRequest,
GatewayRouteDecision,
classify_gateway_route,
)
def _gateway_module() -> Any:
from . import gateway as gateway_module
return gateway_module
def _parse_query_string(query_string: str | None) -> dict[str, str]:
return {
str(key): str(value) for key, value in parse_qsl(query_string or "", keep_blank_values=True)
}
async def _resolve_auth_context(
payload: GatewayResolveRequest,
decision: GatewayRouteDecision,
) -> dict[str, Any] | None:
gateway_module = _gateway_module()
return await gateway_module._resolve_auth_context_signature(
headers=payload.headers,
query_string=payload.query_string,
auth_endpoint_signature=str(decision.auth_endpoint_signature or ""),
)
async def _resolve_auth_context_signature(
*,
headers: dict[str, str] | None,
query_string: str | None,
auth_endpoint_signature: str,
) -> dict[str, Any] | None:
normalized_signature = str(auth_endpoint_signature or "").strip().lower()
if not normalized_signature:
return None
client_api_key = extract_client_api_key_for_endpoint_with_query(
headers or {},
_parse_query_string(query_string),
normalized_signature,
)
if not client_api_key:
return None
auth_result = await AuthService.authenticate_api_key_threadsafe(client_api_key)
if not auth_result or not auth_result.user or not auth_result.api_key:
return None
return GatewayAuthContext(
user_id=str(auth_result.user.id),
api_key_id=str(auth_result.api_key.id),
balance_remaining=auth_result.balance_remaining,
access_allowed=bool(auth_result.access_allowed),
).model_dump(exclude_none=True)
async def _resolve_gateway_execute_auth_context(
*,
payload: GatewayExecuteRequest,
decision: GatewayRouteDecision,
) -> GatewayAuthContext | None:
gateway_module = _gateway_module()
if payload.auth_context is not None:
return payload.auth_context
resolved = await gateway_module._resolve_auth_context_signature(
headers=payload.headers,
query_string=payload.query_string,
auth_endpoint_signature=str(decision.auth_endpoint_signature or ""),
)
if not resolved:
return None
return GatewayAuthContext.model_validate(resolved)
def _resolve_gateway_sync_adapter(
decision: GatewayRouteDecision,
path: str,
) -> tuple[Any | None, dict[str, Any]]:
gateway_module = _gateway_module()
if decision.route_class != "ai_public":
return None, {}
family = str(decision.route_family or "").strip().lower()
kind = str(decision.route_kind or "").strip().lower()
if family == "openai" and kind == "chat":
from src.api.handlers.openai import OpenAIChatAdapter
return OpenAIChatAdapter(), {}
if family == "openai" and kind == "cli":
from src.api.handlers.openai_cli import OpenAICliAdapter
return OpenAICliAdapter(), {}
if family == "openai" and kind == "compact":
from src.api.handlers.openai_cli import OpenAICompactAdapter
return OpenAICompactAdapter(), {}
if family == "claude" and kind == "chat":
from src.api.handlers.claude.adapter import ClaudeChatAdapter
return ClaudeChatAdapter(), {}
if family == "claude" and kind == "cli":
from src.api.handlers.claude_cli import ClaudeCliAdapter
return ClaudeCliAdapter(), {}
if family == "gemini" and kind == "chat":
from src.api.handlers.gemini.adapter import GeminiChatAdapter
return GeminiChatAdapter(), gateway_module._extract_gemini_path_params(path)
if family == "gemini" and kind == "cli":
from src.api.handlers.gemini_cli import GeminiCliAdapter
return GeminiCliAdapter(), gateway_module._extract_gemini_path_params(path)
if family == "openai" and kind == "video":
from src.api.handlers.openai.video_adapter import OpenAIVideoAdapter
return OpenAIVideoAdapter(), gateway_module._extract_openai_video_path_params(path)
if family == "gemini" and kind == "video":
from src.api.handlers.gemini.video_adapter import GeminiVeoAdapter
return GeminiVeoAdapter(), gateway_module._extract_gemini_video_path_params(path)
return None, {}
async def _build_gateway_stream_plan_response(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
) -> GatewayExecutionPlanResponse | None:
gateway_module = _gateway_module()
decision = classify_gateway_route(payload.method, payload.path, payload.headers)
auth_context = await gateway_module._resolve_gateway_execute_auth_context(
payload=payload,
decision=decision,
)
if auth_context is None or not auth_context.access_allowed:
return None
if decision.route_family == "openai" and decision.route_kind == "chat":
planned = await gateway_module._build_openai_chat_stream_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
if planned is not None:
planned.auth_context = auth_context
return planned
if decision.route_family == "claude" and decision.route_kind == "chat":
planned = await gateway_module._build_claude_chat_stream_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
if planned is not None:
planned.auth_context = auth_context
return planned
if decision.route_family == "gemini" and decision.route_kind == "chat":
planned = await gateway_module._build_gemini_chat_stream_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
if planned is not None:
planned.auth_context = auth_context
return planned
if decision.route_family == "openai" and decision.route_kind in {"cli", "compact"}:
planned = await gateway_module._build_openai_cli_stream_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
if planned is not None:
planned.auth_context = auth_context
return planned
if decision.route_family == "claude" and decision.route_kind == "cli":
planned = await gateway_module._build_claude_cli_stream_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
if planned is not None:
planned.auth_context = auth_context
return planned
if decision.route_family == "gemini" and decision.route_kind == "cli":
planned = await gateway_module._build_gemini_cli_stream_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
if planned is not None:
planned.auth_context = auth_context
return planned
if decision.route_family == "gemini" and decision.route_kind == "files":
download_match = _GEMINI_FILES_DOWNLOAD_ROUTE_RE.match(str(payload.path or "").strip())
if not download_match or str(payload.method or "").strip().upper() != "GET":
return None
gateway_request = gateway_module._build_gateway_forward_request(
request=request, payload=payload
)
user, api_key = gateway_module._load_gateway_auth_models(db, auth_context)
pipeline = gateway_module.get_pipeline()
await pipeline._check_user_rate_limit(gateway_request, db, user, api_key)
plan = await gateway_module._build_gemini_files_download_stream_plan(
request=gateway_request,
db=db,
user=user,
user_api_key=api_key,
file_id=download_match.group("file_id"),
)
return GatewayExecutionPlanResponse(
action="executor_stream",
plan_kind="gemini_files_download",
plan=plan,
auth_context=auth_context,
)
if decision.route_family == "openai" and decision.route_kind == "video":
content_match = _OPENAI_VIDEO_CONTENT_ROUTE_RE.match(str(payload.path or "").strip())
if not content_match or str(payload.method or "").strip().upper() != "GET":
return None
gateway_request = gateway_module._build_gateway_forward_request(
request=request, payload=payload
)
user, api_key = gateway_module._load_gateway_auth_models(db, auth_context)
pipeline = gateway_module.get_pipeline()
await pipeline._check_user_rate_limit(gateway_request, db, user, api_key)
plan = await gateway_module._build_openai_video_content_stream_plan(
request=gateway_request,
db=db,
user=user,
user_api_key=api_key,
task_id=content_match.group("task_id"),
)
return GatewayExecutionPlanResponse(
action="executor_stream",
plan_kind="openai_video_content",
plan=plan,
auth_context=auth_context,
)
return None
async def _build_gateway_stream_decision_response(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionDecisionResponse | None:
gateway_module = _gateway_module()
if decision.route_family == "openai" and decision.route_kind == "chat":
return await gateway_module._build_openai_chat_stream_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
if decision.route_family == "claude" and decision.route_kind == "chat":
return await gateway_module._build_claude_chat_stream_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
if decision.route_family == "gemini" and decision.route_kind == "chat":
return await gateway_module._build_gemini_chat_stream_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
if decision.route_family == "openai" and decision.route_kind in {"cli", "compact"}:
return await gateway_module._build_openai_cli_stream_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
if decision.route_family == "claude" and decision.route_kind == "cli":
return await gateway_module._build_claude_cli_stream_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
if decision.route_family == "gemini" and decision.route_kind == "cli":
return await gateway_module._build_gemini_cli_stream_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
if decision.route_family == "gemini" and decision.route_kind == "files":
normalized_path = str(payload.path or "").strip()
download_match = _GEMINI_FILES_DOWNLOAD_ROUTE_RE.match(normalized_path)
if str(payload.method or "").strip().upper() == "GET" and download_match:
gateway_request = gateway_module._build_gateway_forward_request(
request=request, payload=payload
)
user, api_key = gateway_module._load_gateway_auth_models(db, auth_context)
return await gateway_module._build_gemini_files_download_stream_decision(
request=gateway_request,
db=db,
user=user,
user_api_key=api_key,
file_id=str(download_match.group("file_id") or "").strip(),
)
if decision.route_family == "openai" and decision.route_kind == "video":
normalized_path = str(payload.path or "").strip()
content_match = _OPENAI_VIDEO_CONTENT_ROUTE_RE.match(normalized_path)
if str(payload.method or "").strip().upper() == "GET" and content_match:
gateway_request = gateway_module._build_gateway_forward_request(
request=request, payload=payload
)
user, api_key = gateway_module._load_gateway_auth_models(db, auth_context)
return await gateway_module._build_openai_video_content_stream_decision(
request=gateway_request,
db=db,
user=user,
user_api_key=api_key,
task_id=str(content_match.group("task_id") or "").strip(),
)
return None
async def _build_gateway_sync_plan_response(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
) -> GatewayExecutionPlanResponse | None:
gateway_module = _gateway_module()
decision = classify_gateway_route(payload.method, payload.path, payload.headers)
auth_context = await gateway_module._resolve_gateway_execute_auth_context(
payload=payload,
decision=decision,
)
if auth_context is None or not auth_context.access_allowed:
return None
if decision.route_family == "openai" and decision.route_kind == "chat":
planned = await gateway_module._build_openai_chat_sync_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
if planned is not None:
planned.auth_context = auth_context
return planned
if decision.route_family == "claude" and decision.route_kind == "chat":
planned = await gateway_module._build_claude_chat_sync_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
if planned is not None:
planned.auth_context = auth_context
return planned
if decision.route_family == "gemini" and decision.route_kind == "chat":
planned = await gateway_module._build_gemini_chat_sync_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
if planned is not None:
planned.auth_context = auth_context
return planned
if decision.route_family == "openai" and decision.route_kind in {"cli", "compact"}:
planned = await gateway_module._build_openai_cli_sync_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
if planned is not None:
planned.auth_context = auth_context
return planned
if decision.route_family == "claude" and decision.route_kind == "cli":
planned = await gateway_module._build_claude_cli_sync_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
if planned is not None:
planned.auth_context = auth_context
return planned
if decision.route_family == "gemini" and decision.route_kind == "cli":
planned = await gateway_module._build_gemini_cli_sync_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
if planned is not None:
planned.auth_context = auth_context
return planned
if decision.route_family == "openai" and decision.route_kind == "video":
normalized_path = str(payload.path or "").strip()
normalized_method = str(payload.method or "").strip().upper()
if normalized_method == "POST" and normalized_path == "/v1/videos":
planned = await gateway_module._build_openai_video_create_sync_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
if planned is not None:
planned.auth_context = auth_context
return planned
remix_match = _OPENAI_VIDEO_REMIX_ROUTE_RE.match(normalized_path)
if normalized_method == "POST" and remix_match:
planned = await gateway_module._build_openai_video_remix_sync_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
task_id=str(remix_match.group("task_id") or "").strip(),
)
if planned is not None:
planned.auth_context = auth_context
return planned
cancel_match = _OPENAI_VIDEO_CANCEL_ROUTE_RE.match(normalized_path)
if normalized_method == "POST" and cancel_match:
planned = await gateway_module._build_openai_video_cancel_sync_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
task_id=str(cancel_match.group("task_id") or "").strip(),
)
if planned is not None:
planned.auth_context = auth_context
return planned
task_match = _OPENAI_VIDEO_TASK_ROUTE_RE.match(normalized_path)
if normalized_method == "DELETE" and task_match:
planned = await gateway_module._build_openai_video_delete_sync_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
task_id=str(task_match.group("task_id") or "").strip(),
)
if planned is not None:
planned.auth_context = auth_context
return planned
if decision.route_family == "gemini" and decision.route_kind == "video":
normalized_path = str(payload.path or "").strip()
normalized_method = str(payload.method or "").strip().upper()
create_match = _GEMINI_VIDEO_CREATE_ROUTE_RE.match(normalized_path)
if normalized_method == "POST" and create_match:
planned = await gateway_module._build_gemini_video_create_sync_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
model=str(create_match.group("model") or "").strip(),
)
if planned is not None:
planned.auth_context = auth_context
return planned
if normalized_method == "POST" and (
_GEMINI_MODEL_OPERATION_CANCEL_RE.match(normalized_path)
or _GEMINI_OPERATION_CANCEL_RE.match(normalized_path)
):
task_id = str(
gateway_module._extract_gemini_video_path_params(normalized_path).get("task_id")
or ""
)
planned = await gateway_module._build_gemini_video_cancel_sync_plan(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
task_id=task_id,
)
if planned is not None:
planned.auth_context = auth_context
return planned
if decision.route_family == "gemini" and decision.route_kind == "files":
resource_match = _GEMINI_FILES_RESOURCE_ROUTE_RE.match(str(payload.path or "").strip())
download_match = _GEMINI_FILES_DOWNLOAD_ROUTE_RE.match(str(payload.path or "").strip())
method = str(payload.method or "").strip().upper()
if method == "POST" and str(payload.path or "").strip() == "/upload/v1beta/files":
gateway_request = gateway_module._build_gateway_forward_request(
request=request, payload=payload
)
user, api_key = gateway_module._load_gateway_auth_models(db, auth_context)
pipeline = gateway_module.get_pipeline()
await pipeline._check_user_rate_limit(gateway_request, db, user, api_key)
plan, report_context = await gateway_module._build_gemini_files_proxy_sync_plan(
request=gateway_request,
db=db,
method="POST",
upstream_path="/v1beta/files",
is_upload=True,
)
return GatewayExecutionPlanResponse(
action="executor_sync",
plan_kind="gemini_files_upload",
plan=plan,
report_kind="gemini_files_store_mapping",
report_context=report_context,
auth_context=auth_context,
)
if method == "GET" and str(payload.path or "").strip() == "/v1beta/files":
gateway_request = gateway_module._build_gateway_forward_request(
request=request, payload=payload
)
user, api_key = gateway_module._load_gateway_auth_models(db, auth_context)
pipeline = gateway_module.get_pipeline()
await pipeline._check_user_rate_limit(gateway_request, db, user, api_key)
plan, report_context = await gateway_module._build_gemini_files_proxy_sync_plan(
request=gateway_request,
db=db,
method="GET",
upstream_path="/v1beta/files",
)
return GatewayExecutionPlanResponse(
action="executor_sync",
plan_kind="gemini_files_list",
plan=plan,
report_kind="gemini_files_store_mapping",
report_context=report_context,
auth_context=auth_context,
)
if (
not resource_match
or download_match
or str(payload.path or "").strip() == "/v1beta/files"
):
return None
gateway_request = gateway_module._build_gateway_forward_request(
request=request, payload=payload
)
user, api_key = gateway_module._load_gateway_auth_models(db, auth_context)
pipeline = gateway_module.get_pipeline()
await pipeline._check_user_rate_limit(gateway_request, db, user, api_key)
file_name = resource_match.group("file_name")
if method == "GET":
plan, report_context = await gateway_module._build_gemini_files_get_sync_plan(
request=gateway_request,
db=db,
file_name=file_name,
)
return GatewayExecutionPlanResponse(
action="executor_sync",
plan_kind="gemini_files_get",
plan=plan,
report_kind="gemini_files_store_mapping",
report_context=report_context,
auth_context=auth_context,
)
if method == "DELETE":
normalized_file_name = (
file_name if str(file_name or "").startswith("files/") else f"files/{file_name}"
)
plan, _report_context = await gateway_module._build_gemini_files_proxy_sync_plan(
request=gateway_request,
db=db,
method="DELETE",
upstream_path=f"/v1beta/{normalized_file_name}",
)
return GatewayExecutionPlanResponse(
action="executor_sync",
plan_kind="gemini_files_delete",
plan=plan,
report_kind="gemini_files_delete_mapping",
report_context={"file_name": normalized_file_name},
auth_context=auth_context,
)
return None
async def _build_gateway_sync_decision_response(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
auth_context: GatewayAuthContext,
decision: GatewayRouteDecision,
) -> GatewayExecutionDecisionResponse | None:
gateway_module = _gateway_module()
if decision.route_family == "openai" and decision.route_kind == "chat":
return await gateway_module._build_openai_chat_sync_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
if decision.route_family == "openai" and decision.route_kind in {"cli", "compact"}:
return await gateway_module._build_openai_cli_sync_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
if decision.route_family == "claude" and decision.route_kind == "chat":
return await gateway_module._build_claude_chat_sync_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
if decision.route_family == "claude" and decision.route_kind == "cli":
return await gateway_module._build_claude_cli_sync_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
if decision.route_family == "gemini" and decision.route_kind == "chat":
return await gateway_module._build_gemini_chat_sync_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
if decision.route_family == "gemini" and decision.route_kind == "cli":
return await gateway_module._build_gemini_cli_sync_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
if decision.route_family == "openai" and decision.route_kind == "video":
normalized_path = str(payload.path or "").strip()
normalized_method = str(payload.method or "").strip().upper()
if normalized_method == "POST" and normalized_path == "/v1/videos":
return await gateway_module._build_openai_video_create_sync_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
remix_match = _OPENAI_VIDEO_REMIX_ROUTE_RE.match(normalized_path)
if normalized_method == "POST" and remix_match:
return await gateway_module._build_openai_video_remix_sync_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
task_id=str(remix_match.group("task_id") or "").strip(),
)
cancel_match = _OPENAI_VIDEO_CANCEL_ROUTE_RE.match(normalized_path)
if normalized_method == "POST" and cancel_match:
return await gateway_module._build_openai_video_cancel_sync_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
task_id=str(cancel_match.group("task_id") or "").strip(),
)
task_match = _OPENAI_VIDEO_TASK_ROUTE_RE.match(normalized_path)
if normalized_method == "DELETE" and task_match:
return await gateway_module._build_openai_video_delete_sync_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
task_id=str(task_match.group("task_id") or "").strip(),
)
if decision.route_family == "gemini" and decision.route_kind == "video":
normalized_path = str(payload.path or "").strip()
normalized_method = str(payload.method or "").strip().upper()
create_match = _GEMINI_VIDEO_CREATE_ROUTE_RE.match(normalized_path)
if normalized_method == "POST" and create_match:
return await gateway_module._build_gemini_video_create_sync_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
model=str(create_match.group("model") or "").strip(),
)
if normalized_method == "POST" and (
_GEMINI_MODEL_OPERATION_CANCEL_RE.match(normalized_path)
or _GEMINI_OPERATION_CANCEL_RE.match(normalized_path)
):
task_id = str(
gateway_module._extract_gemini_video_path_params(normalized_path).get("task_id")
or ""
)
return await gateway_module._build_gemini_video_cancel_sync_decision(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
task_id=task_id,
)
if decision.route_family == "gemini" and decision.route_kind == "files":
normalized_path = str(payload.path or "").strip()
method = str(payload.method or "").strip().upper()
resource_match = _GEMINI_FILES_RESOURCE_ROUTE_RE.match(normalized_path)
download_match = _GEMINI_FILES_DOWNLOAD_ROUTE_RE.match(normalized_path)
gateway_request = gateway_module._build_gateway_forward_request(
request=request, payload=payload
)
user, api_key = gateway_module._load_gateway_auth_models(db, auth_context)
pipeline = gateway_module.get_pipeline()
await pipeline._check_user_rate_limit(gateway_request, db, user, api_key)
if method == "GET" and normalized_path == "/v1beta/files":
return await gateway_module._build_gemini_files_proxy_sync_decision(
request=gateway_request,
db=db,
method="GET",
upstream_path="/v1beta/files",
decision_kind="gemini_files_list",
report_kind="gemini_files_store_mapping",
)
if (
method == "GET"
and resource_match
and not download_match
and normalized_path != "/v1beta/files"
):
file_name = str(resource_match.group("file_name") or "").strip()
return await gateway_module._build_gemini_files_proxy_sync_decision(
request=gateway_request,
db=db,
method="GET",
upstream_path=f"/v1beta/{file_name if file_name.startswith('files/') else f'files/{file_name}'}",
decision_kind="gemini_files_get",
report_kind="gemini_files_store_mapping",
)
if (
method == "DELETE"
and resource_match
and not download_match
and normalized_path != "/v1beta/files"
):
file_name = str(resource_match.group("file_name") or "").strip()
normalized_file_name = (
file_name if file_name.startswith("files/") else f"files/{file_name}"
)
return await gateway_module._build_gemini_files_proxy_sync_decision(
request=gateway_request,
db=db,
method="DELETE",
upstream_path=f"/v1beta/{normalized_file_name}",
decision_kind="gemini_files_delete",
report_kind="gemini_files_delete_mapping",
report_context={"file_name": normalized_file_name},
)
return None
@@ -1,546 +0,0 @@
from __future__ import annotations
import base64
import time
import uuid
from dataclasses import asdict
from datetime import datetime, timezone
from typing import Any
from fastapi import BackgroundTasks, HTTPException, Request
from fastapi.responses import JSONResponse, Response
from sqlalchemy.orm import Session
from src.api.base.context import ApiRequestContext
from src.api.base.pipeline import get_pipeline
from src.core.api_format.headers import extract_client_api_key_for_endpoint_with_query
from src.core.api_format.metadata import get_auth_config_for_endpoint
from src.core.crypto import crypto_service
from src.core.http_compression import normalize_content_encoding
from src.core.logger import logger
from src.database import create_session, get_db
from src.models.database import ApiKey, RequestCandidate, User
from src.services.auth.service import AuthService
from src.utils.async_utils import safe_create_task
from .common import ensure_loopback
from .gateway_contract import (
_GEMINI_FILES_DOWNLOAD_ROUTE_RE,
_GEMINI_FILES_RESOURCE_ROUTE_RE,
_GEMINI_MODEL_OPERATION_CANCEL_RE,
_GEMINI_OPERATION_CANCEL_RE,
_GEMINI_SYNC_ROUTE_RE,
_GEMINI_VIDEO_CREATE_ROUTE_RE,
_GEMINI_VIDEO_MODEL_OPERATION_ANY_RE,
_OPENAI_VIDEO_CANCEL_ROUTE_RE,
_OPENAI_VIDEO_CONTENT_ROUTE_RE,
_OPENAI_VIDEO_REMIX_ROUTE_RE,
_OPENAI_VIDEO_TASK_ROUTE_RE,
CONTROL_ACTION_HEADER,
CONTROL_ACTION_PROXY_PUBLIC,
CONTROL_EXECUTED_HEADER,
GatewayAuthContext,
GatewayExecuteRequest,
GatewayExecutionDecisionResponse,
GatewayExecutionPlanResponse,
GatewayResolveRequest,
GatewayRouteDecision,
GatewayStreamReportRequest,
GatewaySyncReportRequest,
classify_gateway_route,
)
class _GatewayProxy:
def __getattr__(self, name: str) -> Any:
from . import gateway as gateway_module
return getattr(gateway_module, name)
gateway_module = _GatewayProxy()
async def _build_gemini_files_download_stream_plan(
*,
request: Request,
db: Session,
user: User,
user_api_key: ApiKey,
file_id: str,
) -> dict[str, Any]:
from src.api.public.gemini_files import (
GEMINI_FILES_BASE_URL,
UpstreamContext,
_build_upstream_headers,
_build_upstream_url,
_enrich_upstream_context_proxy,
_find_video_task_by_id,
_resolve_files_model_name,
_select_provider_candidate,
resolve_provider_proxy,
)
from src.services.request.execution_runtime_plan import (
ExecutionPlan,
ExecutionPlanBody,
ExecutionPlanTimeouts,
ExecutionProxySnapshot,
)
proxy_snapshot = None
provider_id = ""
endpoint_id = ""
file_key_id = ""
if file_id.startswith("aev_"):
short_id = file_id[4:]
upstream_key, video_url = await _find_video_task_by_id(db, short_id, user.id)
if not upstream_key or not video_url:
raise HTTPException(
status_code=404,
detail={
"error": {
"code": 404,
"message": f"Video not found or not ready: {file_id}",
"status": "NOT_FOUND",
}
},
)
upstream_url = video_url
try:
from src.services.proxy_node.resolver import (
build_proxy_url_async,
get_system_proxy_config_async,
resolve_delegate_config_async,
resolve_proxy_info_async,
)
system_proxy = await get_system_proxy_config_async()
delegate_cfg = await resolve_delegate_config_async(system_proxy)
proxy_url: str | None = None
if system_proxy and not (delegate_cfg and delegate_cfg.get("tunnel")):
proxy_url = await build_proxy_url_async(system_proxy)
proxy_info = await resolve_proxy_info_async(system_proxy)
proxy_snapshot = ExecutionProxySnapshot.from_proxy_info(
proxy_info,
proxy_url=proxy_url,
mode_override="tunnel" if delegate_cfg and delegate_cfg.get("tunnel") else None,
node_id_override=(
str(delegate_cfg.get("node_id") or "").strip() or None
if delegate_cfg and delegate_cfg.get("tunnel")
else None
),
)
except Exception:
proxy_snapshot = None
else:
model_name = _resolve_files_model_name(db, user_api_key, user)
if not model_name:
raise HTTPException(
status_code=503,
detail={
"error": {
"code": 503,
"message": "No available model for Gemini Files API routing",
"status": "UNAVAILABLE",
}
},
)
candidate = await _select_provider_candidate(
db,
user_api_key,
model_name,
require_files_capability=True,
)
if not candidate:
raise HTTPException(
status_code=503,
detail={
"error": {
"code": 503,
"message": "No available Gemini key with 'gemini_files' capability enabled",
"status": "UNAVAILABLE",
}
},
)
try:
upstream_key = crypto_service.decrypt(candidate.key.api_key)
except Exception as exc:
raise HTTPException(
status_code=500,
detail={
"error": {
"code": 500,
"message": "Failed to decrypt provider key",
"status": "INTERNAL",
}
},
) from exc
regular_ctx = UpstreamContext(
upstream_key=upstream_key,
base_url=candidate.endpoint.base_url or GEMINI_FILES_BASE_URL,
file_key_id=str(candidate.key.id),
user_id=str(user.id),
provider_id=str(candidate.provider.id),
endpoint_id=str(candidate.endpoint.id),
provider_proxy=resolve_provider_proxy(endpoint=candidate.endpoint, key=candidate.key),
key_proxy=(
candidate.key.proxy
if isinstance(getattr(candidate.key, "proxy", None), dict)
else None
),
)
ctx = await _enrich_upstream_context_proxy(regular_ctx)
file_key_id = ctx.file_key_id
provider_id = ctx.provider_id
endpoint_id = ctx.endpoint_id
proxy_snapshot = ctx.proxy_snapshot
file_name = f"files/{file_id}" if not file_id.startswith("files/") else file_id
upstream_url = _build_upstream_url(
ctx.base_url,
f"/v1beta/{file_name}:download",
dict(request.query_params),
)
headers = _build_upstream_headers(dict(request.headers), upstream_key)
plan = ExecutionPlan(
request_id=str(getattr(request.state, "request_id", "") or uuid.uuid4().hex),
candidate_id=None,
provider_name="gemini",
provider_id=provider_id,
endpoint_id=endpoint_id,
key_id=str(file_key_id or ""),
method="GET",
url=upstream_url,
headers=headers,
body=ExecutionPlanBody(),
stream=True,
provider_api_format="gemini:files",
client_api_format="gemini:files",
model_name="gemini-files",
proxy=proxy_snapshot,
timeouts=ExecutionPlanTimeouts(
connect_ms=30_000,
read_ms=300_000,
write_ms=300_000,
pool_ms=30_000,
total_ms=None,
),
)
return plan.to_payload()
async def _build_gemini_files_download_stream_decision(
*,
request: Request,
db: Session,
user: User,
user_api_key: ApiKey,
file_id: str,
) -> GatewayExecutionDecisionResponse:
from src.api.public.gemini_files import (
GEMINI_FILES_BASE_URL,
UpstreamContext,
_build_upstream_headers,
_build_upstream_url,
_enrich_upstream_context_proxy,
_find_video_task_by_id,
_resolve_files_model_name,
_select_provider_candidate,
resolve_provider_proxy,
)
proxy_snapshot = None
provider_id = ""
endpoint_id = ""
file_key_id = ""
if file_id.startswith("aev_"):
short_id = file_id[4:]
upstream_key, video_url = await _find_video_task_by_id(db, short_id, user.id)
if not upstream_key or not video_url:
raise HTTPException(
status_code=404,
detail={
"error": {
"code": 404,
"message": f"Video not found or not ready: {file_id}",
"status": "NOT_FOUND",
}
},
)
upstream_url = video_url
try:
from src.services.proxy_node.resolver import (
build_proxy_url_async,
get_system_proxy_config_async,
resolve_delegate_config_async,
resolve_proxy_info_async,
)
from src.services.request.execution_runtime_plan import ExecutionProxySnapshot
system_proxy = await get_system_proxy_config_async()
delegate_cfg = await resolve_delegate_config_async(system_proxy)
proxy_url: str | None = None
if system_proxy and not (delegate_cfg and delegate_cfg.get("tunnel")):
proxy_url = await build_proxy_url_async(system_proxy)
proxy_info = await resolve_proxy_info_async(system_proxy)
proxy_snapshot = ExecutionProxySnapshot.from_proxy_info(
proxy_info,
proxy_url=proxy_url,
mode_override="tunnel" if delegate_cfg and delegate_cfg.get("tunnel") else None,
node_id_override=(
str(delegate_cfg.get("node_id") or "").strip() or None
if delegate_cfg and delegate_cfg.get("tunnel")
else None
),
)
except Exception:
proxy_snapshot = None
else:
model_name = _resolve_files_model_name(db, user_api_key, user)
if not model_name:
raise HTTPException(
status_code=503,
detail={
"error": {
"code": 503,
"message": "No available model for Gemini Files API routing",
"status": "UNAVAILABLE",
}
},
)
candidate = await _select_provider_candidate(
db,
user_api_key,
model_name,
require_files_capability=True,
)
if not candidate:
raise HTTPException(
status_code=503,
detail={
"error": {
"code": 503,
"message": "No available Gemini key with 'gemini_files' capability enabled",
"status": "UNAVAILABLE",
}
},
)
try:
upstream_key = crypto_service.decrypt(candidate.key.api_key)
except Exception as exc:
raise HTTPException(
status_code=500,
detail={
"error": {
"code": 500,
"message": "Failed to decrypt provider key",
"status": "INTERNAL",
}
},
) from exc
regular_ctx = UpstreamContext(
upstream_key=upstream_key,
base_url=candidate.endpoint.base_url or GEMINI_FILES_BASE_URL,
file_key_id=str(candidate.key.id),
user_id=str(user.id),
provider_id=str(candidate.provider.id),
endpoint_id=str(candidate.endpoint.id),
provider_proxy=resolve_provider_proxy(endpoint=candidate.endpoint, key=candidate.key),
key_proxy=(
candidate.key.proxy
if isinstance(getattr(candidate.key, "proxy", None), dict)
else None
),
)
ctx = await _enrich_upstream_context_proxy(regular_ctx)
file_key_id = ctx.file_key_id
provider_id = ctx.provider_id
endpoint_id = ctx.endpoint_id
proxy_snapshot = ctx.proxy_snapshot
file_name = f"files/{file_id}" if not file_id.startswith("files/") else file_id
upstream_url = _build_upstream_url(
ctx.base_url,
f"/v1beta/{file_name}:download",
dict(request.query_params),
)
headers = _build_upstream_headers(dict(request.headers), upstream_key)
return GatewayExecutionDecisionResponse(
action="executor_stream_decision",
decision_kind="gemini_files_download",
request_id=str(getattr(request.state, "request_id", "") or uuid.uuid4().hex),
provider_name="gemini",
provider_id=provider_id,
endpoint_id=endpoint_id,
key_id=str(file_key_id or ""),
upstream_base_url="",
upstream_url=upstream_url,
auth_header="",
auth_value="",
provider_api_format="gemini:files",
client_api_format="gemini:files",
model_name="gemini-files",
provider_request_headers=headers,
proxy=asdict(proxy_snapshot) if proxy_snapshot else None,
timeouts={
"connect_ms": 30_000,
"read_ms": 300_000,
"write_ms": 300_000,
"pool_ms": 30_000,
},
)
async def _build_gemini_files_get_sync_plan(
*,
request: Request,
db: Session,
file_name: str,
) -> tuple[dict[str, Any], dict[str, Any]]:
normalized_file_name = (
file_name if str(file_name or "").startswith("files/") else f"files/{file_name}"
)
return await _build_gemini_files_proxy_sync_plan(
request=request,
db=db,
method="GET",
upstream_path=f"/v1beta/{normalized_file_name}",
)
async def _build_gemini_files_proxy_sync_decision(
*,
request: Request,
db: Session,
method: str,
upstream_path: str,
decision_kind: str,
report_kind: str,
report_context: dict[str, Any] | None = None,
) -> GatewayExecutionDecisionResponse:
from src.api.public.gemini_files import (
_build_upstream_headers,
_build_upstream_url,
_enrich_upstream_context_proxy,
_resolve_upstream_context,
)
ctx = await _resolve_upstream_context(request, db)
ctx = await _enrich_upstream_context_proxy(ctx)
upstream_url = _build_upstream_url(
ctx.base_url,
upstream_path,
dict(request.query_params),
)
headers = _build_upstream_headers(dict(request.headers), ctx.upstream_key)
merged_report_context = {
"file_key_id": str(ctx.file_key_id or ""),
"user_id": str(ctx.user_id or ""),
}
if report_context:
merged_report_context.update(report_context)
return GatewayExecutionDecisionResponse(
action="executor_sync_decision",
decision_kind=decision_kind,
request_id=str(getattr(request.state, "request_id", "") or uuid.uuid4().hex),
provider_name="gemini",
provider_id=str(ctx.provider_id or ""),
endpoint_id=str(ctx.endpoint_id or ""),
key_id=str(ctx.file_key_id or ""),
upstream_base_url=str(ctx.base_url or ""),
upstream_url=upstream_url,
auth_header="",
auth_value="",
provider_api_format="gemini:files",
client_api_format="gemini:files",
model_name="gemini-files",
provider_request_headers=headers,
proxy=asdict(ctx.proxy_snapshot) if ctx.proxy_snapshot else None,
timeouts={
"connect_ms": 30_000,
"read_ms": 300_000,
"write_ms": 300_000,
"pool_ms": 30_000,
"total_ms": 300_000,
},
report_kind=report_kind,
report_context=merged_report_context,
)
async def _build_gemini_files_proxy_sync_plan(
*,
request: Request,
db: Session,
method: str,
upstream_path: str,
is_upload: bool = False,
) -> tuple[dict[str, Any], dict[str, Any]]:
from src.api.public.gemini_files import (
_build_upstream_headers,
_build_upstream_url,
_enrich_upstream_context_proxy,
_resolve_upstream_context,
)
from src.services.request.execution_runtime_plan import (
ExecutionPlan,
ExecutionPlanTimeouts,
build_execution_plan_body,
)
ctx = await _resolve_upstream_context(request, db)
ctx = await _enrich_upstream_context_proxy(ctx)
upstream_url = _build_upstream_url(
ctx.base_url,
upstream_path,
dict(request.query_params),
is_upload=is_upload,
)
headers = _build_upstream_headers(dict(request.headers), ctx.upstream_key)
body_bytes = await request.body()
content_type = str(headers.get("content-type") or "").strip() or None
plan = ExecutionPlan(
request_id=str(getattr(request.state, "request_id", "") or uuid.uuid4().hex),
candidate_id=None,
provider_name="gemini",
provider_id=str(ctx.provider_id or ""),
endpoint_id=str(ctx.endpoint_id or ""),
key_id=str(ctx.file_key_id or ""),
method=method.upper(),
url=upstream_url,
headers=headers,
body=(
build_execution_plan_body(body_bytes, content_type=content_type)
if body_bytes
else build_execution_plan_body(None)
),
stream=False,
provider_api_format="gemini:files",
client_api_format="gemini:files",
model_name="gemini-files",
proxy=ctx.proxy_snapshot,
timeouts=ExecutionPlanTimeouts(
connect_ms=30_000,
read_ms=300_000,
write_ms=300_000,
pool_ms=30_000,
total_ms=300_000,
),
)
return plan.to_payload(), {
"file_key_id": str(ctx.file_key_id or ""),
"user_id": str(ctx.user_id or ""),
}
@@ -1,41 +0,0 @@
"""Compatibility re-export layer for gateway finalize helpers."""
from __future__ import annotations
from .gateway_finalize_chat import (
_finalize_gateway_chat_sync,
_run_gateway_chat_sync_finalize_background,
_run_gateway_chat_sync_finalize_background_with_session,
)
from .gateway_finalize_cli import (
_finalize_gateway_cli_sync,
_run_gateway_cli_sync_finalize_background,
_run_gateway_cli_sync_finalize_background_with_session,
)
from .gateway_finalize_common import (
_build_gateway_embedded_error_payload,
_build_gateway_sync_error_payload,
_extract_gateway_report_body_bytes,
_extract_gateway_sync_error_message,
_finalize_gateway_sync_response,
_gateway_module,
_resolve_gateway_finalize_db,
_resolve_gateway_sync_error_status_code,
)
__all__ = [
"_build_gateway_embedded_error_payload",
"_build_gateway_sync_error_payload",
"_extract_gateway_report_body_bytes",
"_extract_gateway_sync_error_message",
"_finalize_gateway_chat_sync",
"_finalize_gateway_cli_sync",
"_finalize_gateway_sync_response",
"_gateway_module",
"_resolve_gateway_finalize_db",
"_resolve_gateway_sync_error_status_code",
"_run_gateway_chat_sync_finalize_background",
"_run_gateway_chat_sync_finalize_background_with_session",
"_run_gateway_cli_sync_finalize_background",
"_run_gateway_cli_sync_finalize_background_with_session",
]
@@ -1,360 +0,0 @@
from __future__ import annotations
import base64
import inspect
import time
import uuid
from typing import Any
from fastapi import BackgroundTasks, HTTPException, Request
from fastapi.responses import JSONResponse, Response
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.database import create_session, get_db
from src.models.database import ApiKey, User
from .gateway_contract import GatewaySyncReportRequest
from .gateway_finalize_common import _gateway_module
async def _run_gateway_chat_sync_finalize_background(
payload: GatewaySyncReportRequest,
db: Session,
) -> None:
gateway_module = _gateway_module()
try:
await gateway_module._finalize_gateway_chat_sync(
payload,
db=db,
background_tasks=None,
allow_fast_path=False,
)
except Exception as exc:
logger.warning("gateway background chat finalize failed: {}", exc)
async def _run_gateway_chat_sync_finalize_background_with_session(
payload: GatewaySyncReportRequest,
) -> None:
gateway_module = _gateway_module()
db = create_session()
try:
await gateway_module._run_gateway_chat_sync_finalize_background(payload, db)
finally:
gateway_module._close_gateway_session(db)
async def _finalize_gateway_chat_sync(
payload: GatewaySyncReportRequest,
*,
db: Session,
background_tasks: BackgroundTasks | None = None,
allow_fast_path: bool = True,
) -> Response:
from src.api.handlers.base.chat_sync_executor import ChatSyncExecutor
from src.api.handlers.base.parsers import get_parser_for_format
from src.api.handlers.base.stream_context import is_format_converted
from src.api.handlers.base.utils import (
build_json_response_for_client,
filter_proxy_response_headers,
get_format_converter_registry,
resolve_client_accept_encoding,
)
from src.api.handlers.claude import ClaudeChatAdapter
from src.api.handlers.gemini import GeminiChatAdapter
from src.api.handlers.openai import OpenAIChatAdapter
from src.core.exceptions import EmbeddedErrorException
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
from src.services.scheduling.schemas import ProviderCandidate
gateway_module = _gateway_module()
context = dict(payload.report_context or {})
user_id = str(context.get("user_id") or "").strip()
api_key_id = str(context.get("api_key_id") or "").strip()
provider_id = str(context.get("provider_id") or "").strip()
endpoint_id = str(context.get("endpoint_id") or "").strip()
key_id = str(context.get("key_id") or "").strip()
request_id = str(context.get("request_id") or payload.trace_id or uuid.uuid4().hex[:8]).strip()
model = str(context.get("model") or "unknown").strip() or "unknown"
provider_api_format = str(context.get("provider_api_format") or "").strip().lower()
client_api_format = str(context.get("client_api_format") or "").strip().lower()
if not all([user_id, api_key_id, provider_id, endpoint_id, key_id, client_api_format]):
raise HTTPException(status_code=400, detail="Missing gateway chat finalize context")
if allow_fast_path and background_tasks is not None:
fast_response = await gateway_module._maybe_build_gateway_core_sync_fast_success_response(
payload
)
if fast_response is not None:
background_tasks.add_task(
gateway_module._run_gateway_chat_sync_finalize_background,
payload.model_copy(deep=True),
db,
)
return fast_response
user = db.query(User).filter(User.id == user_id).first()
api_key = db.query(ApiKey).filter(ApiKey.id == api_key_id).first()
provider = db.query(Provider).filter(Provider.id == provider_id).first()
endpoint = db.query(ProviderEndpoint).filter(ProviderEndpoint.id == endpoint_id).first()
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
if not user or not api_key or not provider or not endpoint or not key:
raise HTTPException(status_code=400, detail="Invalid gateway chat finalize context")
if client_api_format == "claude:chat":
adapter = ClaudeChatAdapter()
elif client_api_format == "gemini:chat":
adapter = GeminiChatAdapter()
else:
adapter = OpenAIChatAdapter()
handler = adapter._create_handler(
db=db,
user=user,
api_key=api_key,
request_id=request_id,
client_ip="127.0.0.1",
user_agent=str(
(context.get("original_headers") or {}).get("user-agent") or "aether-gateway"
),
start_time=time.time(),
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
)
original_headers = dict(context.get("original_headers") or {})
original_request_body = dict(context.get("original_request_body") or {})
mapped_model = str(context.get("mapped_model") or "").strip() or None
candidate = ProviderCandidate(
provider=provider,
endpoint=endpoint,
key=key,
mapping_matched_model=mapped_model,
needs_conversion=is_format_converted(provider_api_format, client_api_format),
provider_api_format=provider_api_format,
)
prep = await handler._prepare_provider_request(
model=model,
provider=provider,
endpoint=endpoint,
key=key,
working_request_body=dict(original_request_body),
original_headers=original_headers,
client_api_format=client_api_format,
provider_api_format=provider_api_format,
candidate=candidate,
client_is_stream=False,
)
response_json: dict[str, Any] | None = None
provider_response_json: dict[str, Any] | None = None
if prep.upstream_is_stream and payload.body_base64 and payload.status_code < 400:
sync_executor = ChatSyncExecutor(handler)
sync_executor._ctx.provider_api_format_for_error = provider_api_format
sync_executor._ctx.client_api_format_for_error = client_api_format
sync_executor._ctx.needs_conversion_for_error = bool(prep.needs_conversion)
sync_executor._ctx.mapped_model_result = mapped_model
try:
response_json = await sync_executor._finalize_rust_stream_sync_result(
prepared_plan=prep,
provider=provider,
model=model,
response_body_bytes=gateway_module._extract_gateway_report_body_bytes(payload),
)
if isinstance(sync_executor._ctx.provider_response_json, dict):
provider_response_json = dict(sync_executor._ctx.provider_response_json)
except EmbeddedErrorException as exc:
payload = payload.model_copy(
update={
"status_code": int(exc.error_code or 400),
"body_json": gateway_module._build_gateway_embedded_error_payload(exc),
"body_base64": None,
}
)
except Exception as exc:
raise HTTPException(
status_code=502,
detail="Invalid upstream chat stream response",
) from exc
provider_error_parser = get_parser_for_format(provider_api_format or client_api_format or "")
is_error_response = payload.status_code >= 400
if isinstance(payload.body_json, dict):
try:
is_error_response = is_error_response or provider_error_parser.is_error_response(
dict(payload.body_json)
)
except Exception:
is_error_response = is_error_response or payload.body_json.get("error") is not None
client_accept_encoding = resolve_client_accept_encoding(original_headers, None)
if is_error_response:
error_status_code = gateway_module._resolve_gateway_sync_error_status_code(
payload,
provider_parser=provider_error_parser,
)
error_payload = gateway_module._build_gateway_sync_error_payload(
payload,
client_api_format=client_api_format,
provider_api_format=provider_api_format or client_api_format,
needs_conversion=bool(prep.needs_conversion),
)
client_response_headers = filter_proxy_response_headers(dict(payload.headers or {}))
client_response_headers["content-type"] = "application/json"
client_response = build_json_response_for_client(
status_code=error_status_code,
content=error_payload,
headers=client_response_headers,
client_accept_encoding=client_accept_encoding,
)
request_metadata: dict[str, Any] = {
"gateway_direct_executor": True,
"phase": "3c_trial",
}
proxy_info = context.get("proxy_info")
if isinstance(proxy_info, dict):
request_metadata["proxy"] = proxy_info
telemetry_writer = gateway_module._build_gateway_sync_telemetry_writer(
db=db,
request_id=request_id,
user_id=user_id,
api_key_id=api_key_id,
fallback_telemetry=handler.telemetry,
)
response_time_ms = 0
if isinstance(payload.telemetry, dict):
raw_elapsed_ms = payload.telemetry.get("elapsed_ms")
try:
response_time_ms = max(int(raw_elapsed_ms or 0), 0)
except (TypeError, ValueError):
response_time_ms = 0
await gateway_module._schedule_gateway_sync_telemetry(
background_tasks=background_tasks,
telemetry_writer=telemetry_writer,
operation="record_failure",
provider=str(context.get("provider_name") or provider.name or "unknown"),
model=model,
response_time_ms=response_time_ms,
status_code=error_status_code,
error_message=gateway_module._extract_gateway_sync_error_message(payload),
request_headers=original_headers,
request_body=original_request_body,
provider_request_body=context.get("provider_request_body"),
is_stream=False,
api_format=client_api_format,
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
provider_request_headers=dict(context.get("provider_request_headers") or {}),
response_headers=dict(payload.headers or {}),
client_response_headers=dict(client_response.headers),
provider_id=str(context.get("provider_id") or "") or None,
provider_endpoint_id=str(context.get("endpoint_id") or "") or None,
provider_api_key_id=str(context.get("key_id") or "") or None,
endpoint_api_format=provider_api_format or None,
has_format_conversion=is_format_converted(provider_api_format, client_api_format),
target_model=mapped_model,
metadata=request_metadata,
)
return client_response
if response_json is None:
if not isinstance(payload.body_json, dict):
raise HTTPException(status_code=502, detail="Invalid upstream chat response")
response_json = dict(payload.body_json)
if prep.envelope:
response_json = prep.envelope.unwrap_response(response_json)
prep.envelope.postprocess_unwrapped_response(model=model, data=response_json)
if prep.needs_conversion:
provider_response_json = dict(response_json)
registry = get_format_converter_registry()
response_json = registry.convert_response(
response_json,
provider_api_format,
client_api_format,
requested_model=model,
)
response_json = handler._normalize_response(response_json)
extract_usage = getattr(handler, "_extract_usage", None)
if callable(extract_usage):
usage_info = extract_usage(response_json)
else:
parser = getattr(handler, "parser", None)
parser_extract_usage = getattr(parser, "extract_usage_from_response", None)
usage_info = parser_extract_usage(response_json) if callable(parser_extract_usage) else {}
extract_response_metadata = getattr(handler, "_extract_response_metadata", None)
response_metadata = (
extract_response_metadata(response_json) if callable(extract_response_metadata) else None
)
client_response_headers = filter_proxy_response_headers(dict(payload.headers or {}))
client_response_headers["content-type"] = "application/json"
client_response = build_json_response_for_client(
status_code=payload.status_code,
content=response_json,
headers=client_response_headers,
client_accept_encoding=client_accept_encoding,
)
request_metadata: dict[str, Any] = {
"gateway_direct_executor": True,
"phase": "3c_trial",
}
proxy_info = context.get("proxy_info")
if isinstance(proxy_info, dict):
request_metadata["proxy"] = proxy_info
telemetry_writer = gateway_module._build_gateway_sync_telemetry_writer(
db=db,
request_id=request_id,
user_id=user_id,
api_key_id=api_key_id,
fallback_telemetry=handler.telemetry,
)
response_time_ms = 0
if isinstance(payload.telemetry, dict):
raw_elapsed_ms = payload.telemetry.get("elapsed_ms")
try:
response_time_ms = max(int(raw_elapsed_ms or 0), 0)
except (TypeError, ValueError):
response_time_ms = 0
await gateway_module._schedule_gateway_sync_telemetry(
background_tasks=background_tasks,
telemetry_writer=telemetry_writer,
operation="record_success",
provider=str(context.get("provider_name") or provider.name or "unknown"),
model=model,
input_tokens=int(usage_info.get("input_tokens", 0) or 0),
output_tokens=int(usage_info.get("output_tokens", 0) or 0),
response_time_ms=response_time_ms,
status_code=payload.status_code,
request_headers=original_headers,
request_body=original_request_body,
response_headers=dict(payload.headers or {}),
client_response_headers=dict(client_response.headers),
response_body=provider_response_json or response_json,
client_response_body=response_json if provider_response_json else None,
provider_request_headers=dict(context.get("provider_request_headers") or {}),
provider_request_body=context.get("provider_request_body"),
is_stream=False,
provider_id=str(context.get("provider_id") or "") or None,
provider_endpoint_id=str(context.get("endpoint_id") or "") or None,
provider_api_key_id=str(context.get("key_id") or "") or None,
api_format=client_api_format,
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
endpoint_api_format=provider_api_format or None,
has_format_conversion=is_format_converted(provider_api_format, client_api_format),
target_model=mapped_model,
metadata=gateway_module._build_gateway_usage_metadata(
request_metadata=request_metadata,
response_metadata=response_metadata if response_metadata else None,
),
)
return client_response
@@ -1,366 +0,0 @@
from __future__ import annotations
import base64
import inspect
import time
import uuid
from typing import Any
from fastapi import BackgroundTasks, HTTPException, Request
from fastapi.responses import JSONResponse, Response
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.database import create_session, get_db
from src.models.database import ApiKey, User
from .gateway_contract import GatewaySyncReportRequest
from .gateway_finalize_common import _gateway_module
async def _run_gateway_cli_sync_finalize_background(
payload: GatewaySyncReportRequest,
db: Session,
) -> None:
gateway_module = _gateway_module()
try:
await gateway_module._finalize_gateway_cli_sync(
payload,
db=db,
background_tasks=None,
allow_fast_path=False,
)
except Exception as exc:
logger.warning("gateway background cli finalize failed: {}", exc)
async def _run_gateway_cli_sync_finalize_background_with_session(
payload: GatewaySyncReportRequest,
) -> None:
gateway_module = _gateway_module()
db = create_session()
try:
await gateway_module._run_gateway_cli_sync_finalize_background(payload, db)
finally:
gateway_module._close_gateway_session(db)
async def _finalize_gateway_cli_sync(
payload: GatewaySyncReportRequest,
*,
db: Session,
background_tasks: BackgroundTasks | None = None,
allow_fast_path: bool = True,
) -> Response:
from src.api.handlers.base.parsers import get_parser_for_format
from src.api.handlers.base.stream_context import is_format_converted
from src.api.handlers.base.utils import (
build_json_response_for_client,
filter_proxy_response_headers,
get_format_converter_registry,
resolve_client_accept_encoding,
)
from src.api.handlers.claude_cli import ClaudeCliAdapter
from src.api.handlers.gemini_cli import GeminiCliAdapter
from src.api.handlers.openai_cli import OpenAICliAdapter, OpenAICompactAdapter
from src.core.exceptions import EmbeddedErrorException
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
from src.services.scheduling.schemas import ProviderCandidate
gateway_module = _gateway_module()
context = dict(payload.report_context or {})
user_id = str(context.get("user_id") or "").strip()
api_key_id = str(context.get("api_key_id") or "").strip()
provider_id = str(context.get("provider_id") or "").strip()
endpoint_id = str(context.get("endpoint_id") or "").strip()
key_id = str(context.get("key_id") or "").strip()
request_id = str(context.get("request_id") or payload.trace_id or uuid.uuid4().hex[:8]).strip()
model = str(context.get("model") or "unknown").strip() or "unknown"
provider_api_format = str(context.get("provider_api_format") or "").strip().lower()
client_api_format = str(context.get("client_api_format") or "").strip().lower()
if not all([user_id, api_key_id, provider_id, endpoint_id, key_id, client_api_format]):
raise HTTPException(status_code=400, detail="Missing gateway CLI finalize context")
if allow_fast_path and background_tasks is not None:
fast_response = await gateway_module._maybe_build_gateway_core_sync_fast_success_response(
payload
)
if fast_response is not None:
background_tasks.add_task(
gateway_module._run_gateway_cli_sync_finalize_background,
payload.model_copy(deep=True),
db,
)
return fast_response
user = db.query(User).filter(User.id == user_id).first()
api_key = db.query(ApiKey).filter(ApiKey.id == api_key_id).first()
provider = db.query(Provider).filter(Provider.id == provider_id).first()
endpoint = db.query(ProviderEndpoint).filter(ProviderEndpoint.id == endpoint_id).first()
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
if not user or not api_key or not provider or not endpoint or not key:
raise HTTPException(status_code=400, detail="Invalid gateway CLI finalize context")
if client_api_format == "openai:compact":
adapter = OpenAICompactAdapter()
elif client_api_format == "claude:cli":
adapter = ClaudeCliAdapter()
elif client_api_format == "gemini:cli":
adapter = GeminiCliAdapter()
else:
adapter = OpenAICliAdapter()
handler = adapter.HANDLER_CLASS(
db=db,
user=user,
api_key=api_key,
request_id=request_id,
client_ip="127.0.0.1",
user_agent=str(
(context.get("original_headers") or {}).get("user-agent") or "aether-gateway"
),
start_time=time.time(),
allowed_api_formats=adapter.allowed_api_formats,
adapter_detector=adapter.detect_capability_requirements,
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
)
original_headers = dict(context.get("original_headers") or {})
original_request_body = dict(context.get("original_request_body") or {})
mapped_model = str(context.get("mapped_model") or "").strip() or None
candidate = ProviderCandidate(
provider=provider,
endpoint=endpoint,
key=key,
mapping_matched_model=mapped_model,
needs_conversion=is_format_converted(provider_api_format, client_api_format),
provider_api_format=provider_api_format,
)
upstream_request = await handler._build_upstream_request(
provider=provider,
endpoint=endpoint,
key=key,
request_body=dict(original_request_body),
original_headers=original_headers,
query_params=None,
client_api_format=client_api_format,
provider_api_format=provider_api_format,
fallback_model=model,
mapped_model=mapped_model,
client_is_stream=False,
needs_conversion=bool(candidate.needs_conversion),
output_limit=None,
)
response_json: dict[str, Any] | None = None
provider_response_json: dict[str, Any] | None = None
if upstream_request.upstream_is_stream and payload.body_base64 and payload.status_code < 400:
try:
response_json = await handler._aggregate_upstream_stream_sync_response(
body_bytes=gateway_module._extract_gateway_report_body_bytes(payload),
provider_api_format=provider_api_format,
client_api_format=client_api_format,
provider_name=str(provider.name),
provider_type=str(getattr(provider, "provider_type", "") or "").lower(),
model=model,
request_id=request_id,
envelope=upstream_request.envelope,
)
except EmbeddedErrorException as exc:
payload = payload.model_copy(
update={
"status_code": int(exc.error_code or 400),
"body_json": gateway_module._build_gateway_embedded_error_payload(exc),
"body_base64": None,
}
)
except Exception as exc:
raise HTTPException(
status_code=502,
detail="Invalid upstream CLI stream response",
) from exc
provider_error_parser = get_parser_for_format(provider_api_format or client_api_format or "")
is_error_response = payload.status_code >= 400
if isinstance(payload.body_json, dict):
try:
is_error_response = is_error_response or provider_error_parser.is_error_response(
dict(payload.body_json)
)
except Exception:
is_error_response = is_error_response or payload.body_json.get("error") is not None
client_accept_encoding = resolve_client_accept_encoding(original_headers, None)
if is_error_response:
error_status_code = gateway_module._resolve_gateway_sync_error_status_code(
payload,
provider_parser=provider_error_parser,
)
error_payload = gateway_module._build_gateway_sync_error_payload(
payload,
client_api_format=client_api_format,
provider_api_format=provider_api_format or client_api_format,
needs_conversion=bool(candidate.needs_conversion),
)
client_response_headers = filter_proxy_response_headers(dict(payload.headers or {}))
client_response_headers["content-type"] = "application/json"
client_response = build_json_response_for_client(
status_code=error_status_code,
content=error_payload,
headers=client_response_headers,
client_accept_encoding=client_accept_encoding,
)
request_metadata: dict[str, Any] = {
"gateway_direct_executor": True,
"phase": "3c_trial",
}
proxy_info = context.get("proxy_info")
if isinstance(proxy_info, dict):
request_metadata["proxy"] = proxy_info
telemetry_writer = gateway_module._build_gateway_sync_telemetry_writer(
db=db,
request_id=request_id,
user_id=user_id,
api_key_id=api_key_id,
fallback_telemetry=handler.telemetry,
)
response_time_ms = 0
if isinstance(payload.telemetry, dict):
raw_elapsed_ms = payload.telemetry.get("elapsed_ms")
try:
response_time_ms = max(int(raw_elapsed_ms or 0), 0)
except (TypeError, ValueError):
response_time_ms = 0
await gateway_module._schedule_gateway_sync_telemetry(
background_tasks=background_tasks,
telemetry_writer=telemetry_writer,
operation="record_failure",
provider=str(context.get("provider_name") or provider.name or "unknown"),
model=model,
response_time_ms=response_time_ms,
status_code=error_status_code,
error_message=gateway_module._extract_gateway_sync_error_message(payload),
request_headers=original_headers,
request_body=original_request_body,
provider_request_body=context.get("provider_request_body"),
is_stream=False,
api_format=client_api_format,
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
provider_request_headers=dict(context.get("provider_request_headers") or {}),
response_headers=dict(payload.headers or {}),
client_response_headers=dict(client_response.headers),
provider_id=str(context.get("provider_id") or "") or None,
provider_endpoint_id=str(context.get("endpoint_id") or "") or None,
provider_api_key_id=str(context.get("key_id") or "") or None,
endpoint_api_format=provider_api_format or None,
has_format_conversion=is_format_converted(provider_api_format, client_api_format),
target_model=mapped_model,
metadata=request_metadata,
)
return client_response
if response_json is None:
if not isinstance(payload.body_json, dict):
raise HTTPException(status_code=502, detail="Invalid upstream CLI response")
response_json = dict(payload.body_json)
if upstream_request.envelope:
response_json = upstream_request.envelope.unwrap_response(response_json)
upstream_request.envelope.postprocess_unwrapped_response(
model=model, data=response_json
)
if candidate.needs_conversion:
provider_response_json = dict(response_json)
registry = get_format_converter_registry()
response_json = registry.convert_response(
response_json,
provider_api_format,
client_api_format,
requested_model=model,
)
response_json = handler._normalize_response(response_json)
extract_usage = getattr(handler, "_extract_usage", None)
if callable(extract_usage):
usage_info = extract_usage(response_json)
else:
parser = getattr(handler, "parser", None)
parser_extract_usage = getattr(parser, "extract_usage_from_response", None)
usage_info = parser_extract_usage(response_json) if callable(parser_extract_usage) else {}
extract_response_metadata = getattr(handler, "_extract_response_metadata", None)
response_metadata = (
extract_response_metadata(response_json) if callable(extract_response_metadata) else None
)
client_response_headers = filter_proxy_response_headers(dict(payload.headers or {}))
client_response_headers["content-type"] = "application/json"
client_response = build_json_response_for_client(
status_code=payload.status_code,
content=response_json,
headers=client_response_headers,
client_accept_encoding=client_accept_encoding,
)
request_metadata: dict[str, Any] = {
"gateway_direct_executor": True,
"phase": "3c_trial",
}
proxy_info = context.get("proxy_info")
if isinstance(proxy_info, dict):
request_metadata["proxy"] = proxy_info
telemetry_writer = gateway_module._build_gateway_sync_telemetry_writer(
db=db,
request_id=request_id,
user_id=user_id,
api_key_id=api_key_id,
fallback_telemetry=handler.telemetry,
)
response_time_ms = 0
if isinstance(payload.telemetry, dict):
raw_elapsed_ms = payload.telemetry.get("elapsed_ms")
try:
response_time_ms = max(int(raw_elapsed_ms or 0), 0)
except (TypeError, ValueError):
response_time_ms = 0
await gateway_module._schedule_gateway_sync_telemetry(
background_tasks=background_tasks,
telemetry_writer=telemetry_writer,
operation="record_success",
provider=str(context.get("provider_name") or provider.name or "unknown"),
model=model,
input_tokens=int(usage_info.get("input_tokens", 0) or 0),
output_tokens=int(usage_info.get("output_tokens", 0) or 0),
response_time_ms=response_time_ms,
status_code=payload.status_code,
request_headers=original_headers,
request_body=original_request_body,
response_headers=dict(payload.headers or {}),
client_response_headers=dict(client_response.headers),
response_body=provider_response_json or response_json,
client_response_body=response_json if provider_response_json else None,
provider_request_headers=dict(context.get("provider_request_headers") or {}),
provider_request_body=context.get("provider_request_body"),
is_stream=False,
provider_id=str(context.get("provider_id") or "") or None,
provider_endpoint_id=str(context.get("endpoint_id") or "") or None,
provider_api_key_id=str(context.get("key_id") or "") or None,
api_format=client_api_format,
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
endpoint_api_format=provider_api_format or None,
has_format_conversion=is_format_converted(provider_api_format, client_api_format),
target_model=mapped_model,
metadata=gateway_module._build_gateway_usage_metadata(
request_metadata=request_metadata,
response_metadata=response_metadata if response_metadata else None,
),
)
return client_response
@@ -1,248 +0,0 @@
from __future__ import annotations
import base64
import inspect
import time
import uuid
from typing import Any
from fastapi import BackgroundTasks, HTTPException, Request
from fastapi.responses import JSONResponse, Response
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.database import create_session, get_db
from src.models.database import ApiKey, User
from .gateway_contract import GatewaySyncReportRequest
def _gateway_module() -> Any:
from . import gateway as gateway_module
return gateway_module
async def _finalize_gateway_sync_response(
payload: GatewaySyncReportRequest,
*,
db: Session,
background_tasks: BackgroundTasks | None = None,
) -> Response:
gateway_module = _gateway_module()
if payload.report_kind == "openai_chat_sync_finalize":
return await gateway_module._finalize_gateway_chat_sync(
payload,
db=db,
background_tasks=background_tasks,
)
if payload.report_kind == "claude_chat_sync_finalize":
return await gateway_module._finalize_gateway_chat_sync(
payload,
db=db,
background_tasks=background_tasks,
)
if payload.report_kind == "gemini_chat_sync_finalize":
return await gateway_module._finalize_gateway_chat_sync(
payload,
db=db,
background_tasks=background_tasks,
)
if payload.report_kind == "openai_cli_sync_finalize":
return await gateway_module._finalize_gateway_cli_sync(
payload,
db=db,
background_tasks=background_tasks,
)
if payload.report_kind == "openai_compact_sync_finalize":
return await gateway_module._finalize_gateway_cli_sync(
payload,
db=db,
background_tasks=background_tasks,
)
if payload.report_kind == "claude_cli_sync_finalize":
return await gateway_module._finalize_gateway_cli_sync(
payload,
db=db,
background_tasks=background_tasks,
)
if payload.report_kind == "gemini_cli_sync_finalize":
return await gateway_module._finalize_gateway_cli_sync(
payload,
db=db,
background_tasks=background_tasks,
)
if payload.report_kind == "openai_video_create_sync_finalize":
return await gateway_module._finalize_gateway_openai_video_create_sync(
payload,
db=db,
background_tasks=background_tasks,
)
if payload.report_kind == "openai_video_remix_sync_finalize":
return await gateway_module._finalize_gateway_openai_video_remix_sync(payload, db=db)
if payload.report_kind == "gemini_video_create_sync_finalize":
return await gateway_module._finalize_gateway_gemini_video_create_sync(
payload,
db=db,
background_tasks=background_tasks,
)
if payload.report_kind == "openai_video_cancel_sync_finalize":
return await gateway_module._finalize_gateway_openai_video_cancel_sync(payload, db=db)
if payload.report_kind == "openai_video_delete_sync_finalize":
return await gateway_module._finalize_gateway_openai_video_delete_sync(payload, db=db)
if payload.report_kind == "gemini_video_cancel_sync_finalize":
return await gateway_module._finalize_gateway_gemini_video_cancel_sync(payload, db=db)
raise HTTPException(status_code=400, detail="Unsupported gateway sync finalize kind")
def _resolve_gateway_finalize_db(
request: Request,
) -> tuple[Any, Any | None]:
gateway_module = _gateway_module()
overrides = getattr(getattr(request, "app", None), "dependency_overrides", None)
if isinstance(overrides, dict):
override = overrides.get(get_db)
if callable(override):
override_value = override()
if inspect.isgenerator(override_value):
generator = override_value
db = next(generator)
def _cleanup_override_generator() -> None:
try:
next(generator)
except StopIteration:
return
except Exception as exc:
logger.warning(
"gateway finalize override generator cleanup failed: {}",
exc,
)
return db, _cleanup_override_generator
return override_value, None
db = create_session()
return db, lambda: gateway_module._close_gateway_session(db)
def _extract_gateway_report_body_bytes(payload: Any) -> bytes:
if payload.body_base64:
try:
return base64.b64decode(payload.body_base64, validate=True)
except Exception as exc:
raise HTTPException(
status_code=400, detail="Invalid gateway report body payload"
) from exc
if hasattr(payload, "body_json") and payload.body_json is not None:
return JSONResponse(content=payload.body_json).body
return b""
def _build_gateway_embedded_error_payload(exc: Exception) -> dict[str, Any]:
from src.core.error_utils import extract_client_error_message
from src.core.exceptions import EmbeddedErrorException
message = extract_client_error_message(exc)
payload: dict[str, Any] = {
"error": {
"message": message,
}
}
if isinstance(exc, EmbeddedErrorException):
if exc.error_message and str(exc.error_message).strip():
payload["error"]["message"] = str(exc.error_message).strip()
if exc.error_code is not None:
payload["error"]["code"] = int(exc.error_code)
if exc.error_status:
payload["error"]["status"] = str(exc.error_status)
return payload
def _extract_gateway_sync_error_message(payload: GatewaySyncReportRequest) -> str:
if isinstance(payload.body_json, dict):
error_obj = payload.body_json.get("error")
if isinstance(error_obj, dict):
for key in ("message", "detail", "status", "type", "code"):
value = error_obj.get(key)
if isinstance(value, str) and value.strip():
return value.strip()
elif isinstance(error_obj, str) and error_obj.strip():
return error_obj.strip()
for key in ("message", "detail", "status", "type"):
value = payload.body_json.get(key)
if isinstance(value, str) and value.strip():
return value.strip()
if payload.body_base64:
try:
body_text = base64.b64decode(payload.body_base64, validate=True).decode(
"utf-8", errors="replace"
)
if body_text.strip():
return body_text[:4000]
except Exception:
pass
return f"HTTP {payload.status_code}"
def _resolve_gateway_sync_error_status_code(
payload: GatewaySyncReportRequest,
*,
provider_parser: Any | None = None,
) -> int:
status_code = int(getattr(payload, "status_code", 0) or 0)
if 400 <= status_code < 600:
return status_code
if isinstance(payload.body_json, dict):
if provider_parser is not None:
try:
parsed = provider_parser.parse_response(dict(payload.body_json), status_code or 200)
embedded_status = getattr(parsed, "embedded_status_code", None)
if isinstance(embedded_status, int) and 100 <= embedded_status < 600:
return embedded_status
except Exception:
pass
error_obj = payload.body_json.get("error")
if isinstance(error_obj, dict):
for key in ("code", "status"):
value = error_obj.get(key)
if isinstance(value, int) and 100 <= value < 600:
return value
if isinstance(value, str) and value.isdigit():
parsed_value = int(value)
if 100 <= parsed_value < 600:
return parsed_value
return status_code if 400 <= status_code < 600 else 400
def _build_gateway_sync_error_payload(
payload: GatewaySyncReportRequest,
*,
client_api_format: str,
provider_api_format: str,
needs_conversion: bool,
) -> dict[str, Any]:
from src.api.handlers.base.chat_error_utils import (
_build_client_error_response_best_effort,
_convert_error_response_best_effort,
)
if isinstance(payload.body_json, dict):
if needs_conversion and provider_api_format and client_api_format:
return _convert_error_response_best_effort(
dict(payload.body_json),
provider_api_format,
client_api_format,
)
return dict(payload.body_json)
return _build_client_error_response_best_effort(
_extract_gateway_sync_error_message(payload),
client_api_format or provider_api_format or "openai:chat",
)
@@ -1,95 +0,0 @@
"""Compatibility re-export layer for gateway reporting helpers."""
from __future__ import annotations
from .gateway_reporting_background import (
_apply_gateway_stream_report,
_apply_gateway_sync_report,
_close_gateway_session,
_gateway_sync_report_requires_inline,
_resolve_gateway_background_db,
_run_gateway_stream_report_background,
_run_gateway_stream_report_background_with_session,
_run_gateway_sync_report_background,
_run_gateway_sync_report_background_with_session,
_run_gateway_video_finalize_submitted_background,
)
from .gateway_reporting_candidates import (
_ensure_gateway_request_candidate,
_mark_gateway_sync_candidate_terminal_state,
_record_gateway_direct_candidate_graph,
)
from .gateway_reporting_common import _gateway_module
from .gateway_reporting_failures import (
_record_gateway_chat_sync_failure,
_record_gateway_cli_sync_failure,
_record_gateway_sync_failure,
_record_gateway_video_sync_failure,
_resolve_gateway_failure_adapter,
)
from .gateway_reporting_success_stream import (
_GatewayReportStreamContext,
_iter_gateway_report_body_chunks,
_record_gateway_openai_chat_stream_success,
_record_gateway_passthrough_chat_stream_success,
_record_gateway_passthrough_cli_stream_success,
)
from .gateway_reporting_success_sync import (
_postprocess_gateway_report_provider_response,
_record_gateway_gemini_video_cancel_sync_success,
_record_gateway_gemini_video_create_sync_success,
_record_gateway_openai_chat_sync_success,
_record_gateway_openai_video_cancel_sync_success,
_record_gateway_openai_video_create_sync_success,
_record_gateway_openai_video_delete_sync_success,
_record_gateway_openai_video_remix_sync_success,
_record_gateway_passthrough_chat_sync_success,
_record_gateway_passthrough_cli_sync_success,
)
from .gateway_reporting_telemetry import (
_build_gateway_sync_telemetry_writer,
_build_gateway_usage_metadata,
_dispatch_gateway_sync_telemetry,
_schedule_gateway_sync_telemetry,
)
__all__ = [
"_GatewayReportStreamContext",
"_apply_gateway_stream_report",
"_apply_gateway_sync_report",
"_build_gateway_sync_telemetry_writer",
"_build_gateway_usage_metadata",
"_close_gateway_session",
"_dispatch_gateway_sync_telemetry",
"_ensure_gateway_request_candidate",
"_gateway_module",
"_gateway_sync_report_requires_inline",
"_iter_gateway_report_body_chunks",
"_record_gateway_gemini_video_create_sync_success",
"_mark_gateway_sync_candidate_terminal_state",
"_postprocess_gateway_report_provider_response",
"_record_gateway_gemini_video_cancel_sync_success",
"_record_gateway_chat_sync_failure",
"_record_gateway_cli_sync_failure",
"_record_gateway_direct_candidate_graph",
"_record_gateway_openai_chat_stream_success",
"_record_gateway_openai_chat_sync_success",
"_record_gateway_openai_video_create_sync_success",
"_record_gateway_openai_video_cancel_sync_success",
"_record_gateway_openai_video_delete_sync_success",
"_record_gateway_openai_video_remix_sync_success",
"_record_gateway_passthrough_chat_stream_success",
"_record_gateway_passthrough_chat_sync_success",
"_record_gateway_passthrough_cli_stream_success",
"_record_gateway_passthrough_cli_sync_success",
"_record_gateway_sync_failure",
"_record_gateway_video_sync_failure",
"_resolve_gateway_background_db",
"_resolve_gateway_failure_adapter",
"_run_gateway_stream_report_background",
"_run_gateway_stream_report_background_with_session",
"_run_gateway_sync_report_background",
"_run_gateway_sync_report_background_with_session",
"_run_gateway_video_finalize_submitted_background",
"_schedule_gateway_sync_telemetry",
]
@@ -1,265 +0,0 @@
from __future__ import annotations
import inspect
import time
import uuid
from datetime import datetime, timezone
from typing import Any
from fastapi import BackgroundTasks
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.models.database import ApiKey, RequestCandidate, User
from .gateway_contract import GatewayStreamReportRequest, GatewaySyncReportRequest
from .gateway_reporting_common import _gateway_module
_INLINE_SYNC_REPORT_KINDS = {
"openai_video_create_sync_success",
"openai_video_remix_sync_success",
"gemini_video_create_sync_success",
}
def _gateway_sync_report_requires_inline(payload: GatewaySyncReportRequest) -> bool:
return payload.report_kind in _INLINE_SYNC_REPORT_KINDS
async def _apply_gateway_sync_report(
payload: GatewaySyncReportRequest,
*,
db: Session,
) -> None:
gateway_module = _gateway_module()
if payload.report_kind == "gemini_files_store_mapping":
from src.api.public.gemini_files import _maybe_store_file_mapping_from_payload
await _maybe_store_file_mapping_from_payload(
status_code=payload.status_code,
headers=dict(payload.headers or {}),
content_bytes=gateway_module._extract_gateway_report_body_bytes(payload),
file_key_id=str(payload.report_context.get("file_key_id") or "") or None,
user_id=str(payload.report_context.get("user_id") or "") or None,
)
elif payload.report_kind == "gemini_files_delete_mapping" and payload.status_code < 300:
from src.services.gemini_files_mapping import delete_file_key_mapping
file_name = str(payload.report_context.get("file_name") or "").strip()
if file_name:
await delete_file_key_mapping(file_name)
elif payload.report_kind == "openai_chat_sync_success":
await gateway_module._record_gateway_openai_chat_sync_success(payload, db=db)
elif payload.report_kind == "claude_chat_sync_success":
await gateway_module._record_gateway_passthrough_chat_sync_success(payload, db=db)
elif payload.report_kind == "gemini_chat_sync_success":
await gateway_module._record_gateway_passthrough_chat_sync_success(payload, db=db)
elif payload.report_kind == "openai_chat_sync_error":
await gateway_module._record_gateway_chat_sync_failure(payload, db=db)
elif payload.report_kind == "claude_chat_sync_error":
await gateway_module._record_gateway_chat_sync_failure(payload, db=db)
elif payload.report_kind == "gemini_chat_sync_error":
await gateway_module._record_gateway_chat_sync_failure(payload, db=db)
elif payload.report_kind == "openai_cli_sync_success":
await gateway_module._record_gateway_passthrough_cli_sync_success(payload, db=db)
elif payload.report_kind == "claude_cli_sync_success":
await gateway_module._record_gateway_passthrough_cli_sync_success(payload, db=db)
elif payload.report_kind == "gemini_cli_sync_success":
await gateway_module._record_gateway_passthrough_cli_sync_success(payload, db=db)
elif payload.report_kind == "openai_video_delete_sync_success":
await gateway_module._record_gateway_openai_video_delete_sync_success(payload, db=db)
elif payload.report_kind == "openai_video_cancel_sync_success":
await gateway_module._record_gateway_openai_video_cancel_sync_success(payload, db=db)
elif payload.report_kind == "gemini_video_cancel_sync_success":
await gateway_module._record_gateway_gemini_video_cancel_sync_success(payload, db=db)
elif payload.report_kind == "openai_video_create_sync_success":
await gateway_module._record_gateway_openai_video_create_sync_success(payload, db=db)
elif payload.report_kind == "openai_video_remix_sync_success":
await gateway_module._record_gateway_openai_video_remix_sync_success(payload, db=db)
elif payload.report_kind == "gemini_video_create_sync_success":
await gateway_module._record_gateway_gemini_video_create_sync_success(payload, db=db)
elif payload.report_kind == "openai_cli_sync_error":
await gateway_module._record_gateway_cli_sync_failure(payload, db=db)
elif payload.report_kind == "openai_compact_sync_error":
await gateway_module._record_gateway_cli_sync_failure(payload, db=db)
elif payload.report_kind == "claude_cli_sync_error":
await gateway_module._record_gateway_cli_sync_failure(payload, db=db)
elif payload.report_kind == "gemini_cli_sync_error":
await gateway_module._record_gateway_cli_sync_failure(payload, db=db)
elif payload.report_kind == "openai_video_create_sync_error":
await gateway_module._record_gateway_video_sync_failure(payload, db=db)
elif payload.report_kind == "openai_video_remix_sync_error":
await gateway_module._record_gateway_video_sync_failure(payload, db=db)
elif payload.report_kind == "gemini_video_create_sync_error":
await gateway_module._record_gateway_video_sync_failure(payload, db=db)
elif payload.report_kind == "openai_video_delete_sync_error":
await gateway_module._record_gateway_video_sync_failure(payload, db=db)
elif payload.report_kind == "openai_video_cancel_sync_error":
await gateway_module._record_gateway_video_sync_failure(payload, db=db)
elif payload.report_kind == "gemini_video_cancel_sync_error":
await gateway_module._record_gateway_video_sync_failure(payload, db=db)
async def _run_gateway_sync_report_background(
payload: GatewaySyncReportRequest,
db: Session,
) -> None:
gateway_module = _gateway_module()
context = dict(payload.report_context or {})
candidate = None
try:
candidate = gateway_module._ensure_gateway_request_candidate(
db=db,
report_context=context,
trace_id=payload.trace_id,
initial_status="pending",
)
await gateway_module._apply_gateway_sync_report(payload, db=db)
except Exception as exc:
logger.warning("gateway background sync report failed: {}", exc)
finally:
try:
gateway_module._mark_gateway_sync_candidate_terminal_state(
db=db,
candidate=candidate,
payload=payload,
)
except Exception as exc:
logger.warning("gateway sync candidate finalize failed: {}", exc)
def _close_gateway_session(db: Session) -> None:
try:
db.close()
except Exception as exc:
logger.warning("gateway finalize session close failed: {}", exc)
def _resolve_gateway_background_db(app: Any | None) -> tuple[Any, Any | None]:
gateway_module = _gateway_module()
overrides = getattr(app, "dependency_overrides", None)
if isinstance(overrides, dict):
override = overrides.get(gateway_module.get_db)
if callable(override):
override_value = override()
if inspect.isgenerator(override_value):
generator = override_value
db = next(generator)
def _cleanup_override_generator() -> None:
try:
next(generator)
except StopIteration:
return
except Exception as exc:
logger.warning(
"gateway background override generator cleanup failed: {}",
exc,
)
return db, _cleanup_override_generator
return override_value, None
db = gateway_module.create_session()
return db, lambda: gateway_module._close_gateway_session(db)
async def _run_gateway_sync_report_background_with_session(
payload: GatewaySyncReportRequest,
app: Any | None,
) -> None:
db, cleanup = _resolve_gateway_background_db(app)
try:
await _run_gateway_sync_report_background(payload, db)
finally:
if cleanup is not None:
cleanup()
async def _apply_gateway_stream_report(
payload: GatewayStreamReportRequest,
*,
db: Session,
) -> None:
gateway_module = _gateway_module()
if payload.report_kind == "openai_chat_stream_success":
await gateway_module._record_gateway_openai_chat_stream_success(payload, db=db)
elif payload.report_kind == "claude_chat_stream_success":
await gateway_module._record_gateway_passthrough_chat_stream_success(payload, db=db)
elif payload.report_kind == "gemini_chat_stream_success":
await gateway_module._record_gateway_passthrough_chat_stream_success(payload, db=db)
elif payload.report_kind == "openai_cli_stream_success":
await gateway_module._record_gateway_passthrough_cli_stream_success(payload, db=db)
elif payload.report_kind == "claude_cli_stream_success":
await gateway_module._record_gateway_passthrough_cli_stream_success(payload, db=db)
elif payload.report_kind == "gemini_cli_stream_success":
await gateway_module._record_gateway_passthrough_cli_stream_success(payload, db=db)
async def _run_gateway_stream_report_background(
payload: GatewayStreamReportRequest,
db: Session,
) -> None:
gateway_module = _gateway_module()
try:
gateway_module._ensure_gateway_request_candidate(
db=db,
report_context=dict(payload.report_context or {}),
trace_id=payload.trace_id,
initial_status="streaming",
)
await gateway_module._apply_gateway_stream_report(payload, db=db)
except Exception as exc:
logger.warning("gateway background stream report failed: {}", exc)
async def _run_gateway_stream_report_background_with_session(
payload: GatewayStreamReportRequest,
app: Any | None,
) -> None:
db, cleanup = _resolve_gateway_background_db(app)
try:
await _run_gateway_stream_report_background(payload, db)
finally:
if cleanup is not None:
cleanup()
async def _run_gateway_video_finalize_submitted_background(
*,
db: Session,
request_id: str,
provider_name: str,
provider_id: str | None,
provider_endpoint_id: str | None,
provider_api_key_id: str | None,
response_time_ms: int,
status_code: int,
endpoint_api_format: str | None,
provider_request_headers: dict[str, Any],
response_headers: dict[str, Any],
response_body: dict[str, Any],
) -> None:
from src.services.usage.service import UsageService
try:
UsageService.finalize_submitted(
db,
request_id=request_id,
provider_name=provider_name,
provider_id=provider_id,
provider_endpoint_id=provider_endpoint_id,
provider_api_key_id=provider_api_key_id,
response_time_ms=response_time_ms,
status_code=status_code,
endpoint_api_format=endpoint_api_format,
provider_request_headers=provider_request_headers,
response_headers=response_headers,
response_body=response_body,
)
db.commit()
except Exception as exc:
db.rollback()
logger.warning("gateway background video finalize_submitted failed: {}", exc)
@@ -1,273 +0,0 @@
from __future__ import annotations
import inspect
import time
import uuid
from datetime import datetime, timezone
from typing import Any
from fastapi import BackgroundTasks
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.models.database import ApiKey, RequestCandidate, User
from .gateway_contract import GatewayStreamReportRequest, GatewaySyncReportRequest
from .gateway_reporting_common import _gateway_module
def _load_existing_gateway_direct_candidate_record_map(
*,
db: Session,
request_id: str,
) -> dict[tuple[int, int], str]:
if not request_id or not hasattr(db, "query"):
return {}
rows = (
db.query(RequestCandidate)
.filter(RequestCandidate.request_id == request_id)
.order_by(RequestCandidate.candidate_index, RequestCandidate.retry_index)
.all()
)
return {
(int(row.candidate_index or 0), int(row.retry_index or 0)): str(row.id)
for row in rows
}
def _record_gateway_direct_candidate_graph(
*,
db: Session,
candidate_resolver: Any,
candidates: list[Any],
request_id: str,
user_api_key: ApiKey,
required_capabilities: dict[str, bool] | None,
selected_candidate_index: int,
) -> None:
from src.services.request.candidate import RequestCandidateService
from src.services.scheduling.schemas import PoolCandidate
if not candidates or not request_id or not hasattr(db, "query"):
return
selected_candidate_index = max(0, min(selected_candidate_index, len(candidates) - 1))
selected_candidate = candidates[selected_candidate_index]
try:
user = getattr(user_api_key, "user", None)
except Exception:
user = None
user_id = str(getattr(user_api_key, "user_id", "") or getattr(user, "id", "") or "") or None
try:
candidate_record_map = _load_existing_gateway_direct_candidate_record_map(
db=db,
request_id=request_id,
)
if not candidate_record_map:
candidate_record_map = candidate_resolver.create_candidate_records(
candidates,
request_id,
user_id,
user_api_key,
required_capabilities,
expand_retries=False,
)
except IntegrityError as exc:
db.rollback()
candidate_record_map = _load_existing_gateway_direct_candidate_record_map(
db=db,
request_id=request_id,
)
if not candidate_record_map:
logger.warning(
"[Gateway] failed to create direct candidate graph for request {}: {}",
request_id,
exc,
)
return
logger.debug(
"[Gateway] reused existing direct candidate graph for request {} after duplicate insert",
request_id,
)
except Exception as exc:
db.rollback()
logger.warning(
"[Gateway] failed to create direct candidate graph for request {}: {}",
request_id,
exc,
)
return
selected_retry_index = 0
if isinstance(selected_candidate, PoolCandidate):
selected_retry_index = int(getattr(selected_candidate, "_pool_key_index", 0) or 0)
selected_record_id = candidate_record_map.get((selected_candidate_index, selected_retry_index))
if not selected_record_id:
selected_record_id = candidate_record_map.get((selected_candidate_index, 0))
selected_retry_index = 0
if not selected_record_id:
return
setattr(selected_candidate, "request_candidate_id", selected_record_id)
RequestCandidateService.mark_candidate_started(db, selected_record_id)
unused_record_ids = [
record_id
for (candidate_index, retry_index), record_id in candidate_record_map.items()
if (candidate_index, retry_index) != (selected_candidate_index, selected_retry_index)
]
if not unused_record_ids:
return
now = datetime.now(timezone.utc)
unused_candidates = (
db.query(RequestCandidate)
.filter(
RequestCandidate.id.in_(unused_record_ids),
RequestCandidate.status == "available",
)
.all()
)
for candidate in unused_candidates:
candidate.status = "unused"
candidate.finished_at = now
if unused_candidates:
db.flush()
def _ensure_gateway_request_candidate(
*,
db: Session,
report_context: dict[str, Any],
trace_id: str,
initial_status: str,
) -> RequestCandidate | None:
from src.services.request.candidate import RequestCandidateService
if not hasattr(db, "query"):
return None
request_id = str(report_context.get("request_id") or trace_id or "").strip()
if not request_id:
return None
candidate_id = str(report_context.get("candidate_id") or "").strip() or None
provider_id = str(report_context.get("provider_id") or "").strip() or None
endpoint_id = str(report_context.get("endpoint_id") or "").strip() or None
key_id = str(report_context.get("key_id") or "").strip() or None
user_id = str(report_context.get("user_id") or "").strip() or None
api_key_id = str(report_context.get("api_key_id") or "").strip() or None
client_api_format = str(report_context.get("client_api_format") or "").strip() or None
if not any([candidate_id, provider_id, endpoint_id, key_id, client_api_format]):
return None
candidate: RequestCandidate | None = None
if candidate_id:
candidate = db.query(RequestCandidate).filter(RequestCandidate.id == candidate_id).first()
if candidate is None:
lookup = db.query(RequestCandidate).filter(RequestCandidate.request_id == request_id)
if provider_id:
lookup = lookup.filter(RequestCandidate.provider_id == provider_id)
if endpoint_id:
lookup = lookup.filter(RequestCandidate.endpoint_id == endpoint_id)
if key_id:
lookup = lookup.filter(RequestCandidate.key_id == key_id)
candidate = lookup.order_by(
RequestCandidate.retry_index.desc(),
RequestCandidate.candidate_index.desc(),
RequestCandidate.created_at.desc(),
).first()
if candidate is None:
user = db.query(User).filter(User.id == user_id).first() if user_id else None
api_key = db.query(ApiKey).filter(ApiKey.id == api_key_id).first() if api_key_id else None
candidate_index = (
db.query(RequestCandidate).filter(RequestCandidate.request_id == request_id).count()
)
candidate = RequestCandidateService.create_candidate(
db,
request_id=request_id,
candidate_index=candidate_index,
candidate_id=candidate_id,
user_id=user_id,
api_key_id=api_key_id,
username=str(getattr(user, "username", "") or "") or None,
api_key_name=str(getattr(api_key, "name", "") or "") or None,
provider_id=provider_id,
endpoint_id=endpoint_id,
key_id=key_id,
status=initial_status,
extra_data={
"gateway_direct_executor": True,
"phase": "3c_trial",
"client_api_format": client_api_format,
"provider_api_format": str(report_context.get("provider_api_format") or "") or None,
},
)
current_status = str(candidate.status or "").strip().lower()
if candidate.started_at is None:
candidate.started_at = datetime.now(timezone.utc)
if initial_status == "streaming" and current_status in {
"",
"available",
"unused",
"skipped",
"pending",
"streaming",
}:
candidate.status = "streaming"
elif current_status in {"", "available", "unused", "skipped"}:
candidate.status = "pending"
db.commit()
return candidate
def _mark_gateway_sync_candidate_terminal_state(
*,
db: Session,
candidate: RequestCandidate | None,
payload: GatewaySyncReportRequest,
) -> None:
from src.services.request.candidate import RequestCandidateService
if candidate is None:
return
gateway_module = _gateway_module()
elapsed_ms = 0
if isinstance(payload.telemetry, dict):
raw_elapsed_ms = payload.telemetry.get("elapsed_ms")
try:
elapsed_ms = max(int(raw_elapsed_ms or 0), 0)
except (TypeError, ValueError):
elapsed_ms = 0
has_error_payload = (
isinstance(payload.body_json, dict) and payload.body_json.get("error") is not None
)
if payload.status_code >= 400 or has_error_payload:
RequestCandidateService.mark_candidate_failed(
db=db,
candidate_id=candidate.id,
error_type="gateway_error",
error_message=gateway_module._extract_gateway_sync_error_message(payload),
status_code=payload.status_code,
latency_ms=elapsed_ms or None,
)
return
RequestCandidateService.mark_candidate_success(
db=db,
candidate_id=candidate.id,
status_code=payload.status_code,
latency_ms=elapsed_ms,
extra_data={"gateway_direct_executor": True, "phase": "3c_trial"},
)
@@ -1,21 +0,0 @@
from __future__ import annotations
import inspect
import time
import uuid
from datetime import datetime, timezone
from typing import Any
from fastapi import BackgroundTasks
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.models.database import ApiKey, RequestCandidate, User
from .gateway_contract import GatewayStreamReportRequest, GatewaySyncReportRequest
def _gateway_module() -> Any:
from . import gateway as gateway_module
return gateway_module
@@ -1,165 +0,0 @@
from __future__ import annotations
import inspect
import time
import uuid
from datetime import datetime, timezone
from typing import Any
from fastapi import BackgroundTasks
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.models.database import ApiKey, RequestCandidate, User
from .gateway_contract import GatewayStreamReportRequest, GatewaySyncReportRequest
from .gateway_reporting_common import _gateway_module
def _resolve_gateway_failure_adapter(client_api_format: str, *, cli: bool) -> Any | None:
if cli:
from src.api.handlers.claude_cli import ClaudeCliAdapter
from src.api.handlers.gemini_cli import GeminiCliAdapter
from src.api.handlers.openai_cli import OpenAICliAdapter, OpenAICompactAdapter
if client_api_format == "openai:compact":
return OpenAICompactAdapter()
if client_api_format == "claude:cli":
return ClaudeCliAdapter()
if client_api_format == "gemini:cli":
return GeminiCliAdapter()
if client_api_format == "openai:cli":
return OpenAICliAdapter()
return None
from src.api.handlers.claude import ClaudeChatAdapter
from src.api.handlers.gemini import GeminiChatAdapter
from src.api.handlers.openai import OpenAIChatAdapter
if client_api_format == "claude:chat":
return ClaudeChatAdapter()
if client_api_format == "gemini:chat":
return GeminiChatAdapter()
if client_api_format == "openai:chat":
return OpenAIChatAdapter()
return None
async def _record_gateway_sync_failure(
payload: GatewaySyncReportRequest,
*,
db: Session,
cli: bool,
) -> None:
from src.api.handlers.base.stream_context import is_format_converted
from src.api.handlers.base.utils import filter_proxy_response_headers
gateway_module = _gateway_module()
context = dict(payload.report_context or {})
user_id = str(context.get("user_id") or "").strip()
api_key_id = str(context.get("api_key_id") or "").strip()
if not user_id or not api_key_id:
return
client_api_format = str(context.get("client_api_format") or "").strip().lower()
provider_api_format = str(context.get("provider_api_format") or "").strip().lower()
adapter = _resolve_gateway_failure_adapter(client_api_format, cli=cli)
if adapter is None:
return
user = db.query(User).filter(User.id == user_id).first()
api_key = db.query(ApiKey).filter(ApiKey.id == api_key_id).first()
if not user or not api_key:
return
request_id = str(context.get("request_id") or payload.trace_id or uuid.uuid4().hex[:8]).strip()
model = str(context.get("model") or "unknown").strip() or "unknown"
handler = adapter._create_handler(
db=db,
user=user,
api_key=api_key,
request_id=request_id,
client_ip="127.0.0.1",
user_agent=str(
(context.get("original_headers") or {}).get("user-agent") or "aether-gateway"
),
start_time=time.time(),
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
)
request_metadata: dict[str, Any] = {
"gateway_direct_executor": True,
"phase": "3c_trial",
}
proxy_info = context.get("proxy_info")
if isinstance(proxy_info, dict):
request_metadata["proxy"] = proxy_info
response_time_ms = 0
if isinstance(payload.telemetry, dict):
raw_elapsed_ms = payload.telemetry.get("elapsed_ms")
try:
response_time_ms = max(int(raw_elapsed_ms or 0), 0)
except (TypeError, ValueError):
response_time_ms = 0
client_response_headers = filter_proxy_response_headers(dict(payload.headers or {}))
client_response_headers["content-type"] = "application/json"
telemetry_writer = gateway_module._build_gateway_sync_telemetry_writer(
db=db,
request_id=request_id,
user_id=user_id,
api_key_id=api_key_id,
fallback_telemetry=handler.telemetry,
)
await gateway_module._dispatch_gateway_sync_telemetry(
telemetry_writer=telemetry_writer,
operation="record_failure",
provider=str(context.get("provider_name") or "unknown"),
model=model,
response_time_ms=response_time_ms,
status_code=payload.status_code,
error_message=gateway_module._extract_gateway_sync_error_message(payload),
request_headers=dict(context.get("original_headers") or {}),
request_body=dict(context.get("original_request_body") or {}),
provider_request_body=context.get("provider_request_body"),
is_stream=False,
api_format=client_api_format,
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
provider_request_headers=dict(context.get("provider_request_headers") or {}),
response_headers=dict(payload.headers or {}),
client_response_headers=client_response_headers,
provider_id=str(context.get("provider_id") or "") or None,
provider_endpoint_id=str(context.get("endpoint_id") or "") or None,
provider_api_key_id=str(context.get("key_id") or "") or None,
endpoint_api_format=provider_api_format or None,
has_format_conversion=is_format_converted(provider_api_format, client_api_format),
target_model=str(context.get("mapped_model") or "") or None,
metadata=request_metadata,
)
async def _record_gateway_chat_sync_failure(
payload: GatewaySyncReportRequest,
*,
db: Session,
) -> None:
await _record_gateway_sync_failure(payload, db=db, cli=False)
async def _record_gateway_cli_sync_failure(
payload: GatewaySyncReportRequest,
*,
db: Session,
) -> None:
await _record_gateway_sync_failure(payload, db=db, cli=True)
async def _record_gateway_video_sync_failure(
payload: GatewaySyncReportRequest,
*,
db: Session,
) -> None:
_ = (payload, db)
@@ -1,407 +0,0 @@
from __future__ import annotations
import inspect
import time
import uuid
from datetime import datetime, timezone
from typing import Any
from fastapi import BackgroundTasks
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.models.database import ApiKey, RequestCandidate, User
from .gateway_contract import GatewayStreamReportRequest, GatewaySyncReportRequest
from .gateway_reporting_common import _gateway_module
class _GatewayReportStreamContext:
async def __aenter__(self) -> _GatewayReportStreamContext:
return self
async def __aexit__(self, exc_type: Any, exc: Any, tb: Any) -> None:
return None
async def _iter_gateway_report_body_chunks(body_bytes: bytes) -> Any:
yield body_bytes
async def _record_gateway_openai_chat_stream_success(
payload: GatewayStreamReportRequest,
*,
db: Session,
) -> None:
from src.api.handlers.base.stream_context import StreamContext
from src.api.handlers.base.stream_processor import StreamProcessor
from src.api.handlers.base.stream_telemetry import StreamTelemetryRecorder
from src.api.handlers.openai import OpenAIChatAdapter
from src.services.system.config import SystemConfigService
gateway_module = _gateway_module()
if payload.status_code >= 400:
return
context = dict(payload.report_context or {})
user_id = str(context.get("user_id") or "").strip()
api_key_id = str(context.get("api_key_id") or "").strip()
if not user_id or not api_key_id:
return
body_bytes = gateway_module._extract_gateway_report_body_bytes(payload)
if not body_bytes:
return
user = db.query(User).filter(User.id == user_id).first()
api_key = db.query(ApiKey).filter(ApiKey.id == api_key_id).first()
if not user or not api_key:
return
adapter = OpenAIChatAdapter()
handler = adapter._create_handler(
db=db,
user=user,
api_key=api_key,
request_id=str(context.get("request_id") or payload.trace_id or uuid.uuid4().hex[:8]),
client_ip="127.0.0.1",
user_agent=str(
(context.get("original_headers") or {}).get("user-agent") or "aether-gateway"
),
start_time=time.time(),
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
)
ctx = StreamContext(
model=str(context.get("model") or "unknown"),
api_format=handler.allowed_api_formats[0],
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
)
ctx.request_id = handler.request_id
ctx.client_api_format = str(context.get("client_api_format") or "openai:chat")
ctx.provider_type = str(context.get("provider_name") or "openai")
ctx.update_provider_info(
provider_name=str(context.get("provider_name") or "openai"),
provider_id=str(context.get("provider_id") or ""),
endpoint_id=str(context.get("endpoint_id") or ""),
key_id=str(context.get("key_id") or ""),
provider_api_format=str(context.get("provider_api_format") or "openai:chat"),
)
ctx.mapped_model = str(context.get("mapped_model") or "") or None
ctx.provider_request_headers = dict(context.get("provider_request_headers") or {})
ctx.provider_request_body = context.get("provider_request_body")
ctx.response_headers = dict(payload.headers or {})
ctx.status_code = payload.status_code
ctx.record_parsed_chunks = SystemConfigService.should_log_body(db)
if str(context.get("candidate_id") or "").strip():
ctx.attempt_id = str(context.get("candidate_id"))
proxy_info = context.get("proxy_info")
if isinstance(proxy_info, dict):
ctx.proxy_info = dict(proxy_info)
ctx.set_proxy_timing(ctx.response_headers)
telemetry = payload.telemetry if isinstance(payload.telemetry, dict) else {}
ttfb_ms = telemetry.get("ttfb_ms")
if ttfb_ms is not None:
try:
ctx.first_byte_time_ms = max(int(ttfb_ms), 0)
except (TypeError, ValueError):
ctx.first_byte_time_ms = None
if ctx.proxy_info is not None and ctx.first_byte_time_ms is not None:
ctx.set_ttfb_ms(ctx.first_byte_time_ms)
stream_processor = StreamProcessor(
request_id=ctx.request_id,
default_parser=handler.parser,
on_streaming_start=None,
)
stream = stream_processor.create_response_stream(
ctx,
_iter_gateway_report_body_chunks(body_bytes),
_GatewayReportStreamContext(),
start_time=time.time(),
)
async for _chunk in stream:
pass
elapsed_ms = telemetry.get("elapsed_ms")
try:
response_elapsed_ms = max(int(elapsed_ms or 0), 0)
except (TypeError, ValueError):
response_elapsed_ms = 0
request_start_time = (
time.time() - (response_elapsed_ms / 1000.0) if response_elapsed_ms > 0 else time.time()
)
telemetry_recorder = StreamTelemetryRecorder(
request_id=ctx.request_id,
user_id=str(user.id),
api_key_id=str(api_key.id),
client_ip="127.0.0.1",
format_id=handler.FORMAT_ID,
)
await telemetry_recorder.record_stream_stats(
ctx,
dict(context.get("original_headers") or {}),
dict(context.get("original_request_body") or {}),
request_start_time,
)
async def _record_gateway_passthrough_chat_stream_success(
payload: GatewayStreamReportRequest,
*,
db: Session,
) -> None:
from src.api.handlers.base.stream_context import StreamContext
from src.api.handlers.base.stream_processor import StreamProcessor
from src.api.handlers.base.stream_telemetry import StreamTelemetryRecorder
from src.api.handlers.claude import ClaudeChatAdapter
from src.api.handlers.gemini import GeminiChatAdapter
from src.api.handlers.openai import OpenAIChatAdapter
from src.services.system.config import SystemConfigService
gateway_module = _gateway_module()
if payload.status_code >= 400:
return
context = dict(payload.report_context or {})
user_id = str(context.get("user_id") or "").strip()
api_key_id = str(context.get("api_key_id") or "").strip()
client_api_format = str(context.get("client_api_format") or "").strip().lower()
if not user_id or not api_key_id or not client_api_format:
return
body_bytes = gateway_module._extract_gateway_report_body_bytes(payload)
if not body_bytes:
return
user = db.query(User).filter(User.id == user_id).first()
api_key = db.query(ApiKey).filter(ApiKey.id == api_key_id).first()
if not user or not api_key:
return
if client_api_format == "claude:chat":
adapter = ClaudeChatAdapter()
elif client_api_format == "gemini:chat":
adapter = GeminiChatAdapter()
elif client_api_format == "openai:chat":
adapter = OpenAIChatAdapter()
else:
return
handler = adapter._create_handler(
db=db,
user=user,
api_key=api_key,
request_id=str(context.get("request_id") or payload.trace_id or uuid.uuid4().hex[:8]),
client_ip="127.0.0.1",
user_agent=str(
(context.get("original_headers") or {}).get("user-agent") or "aether-gateway"
),
start_time=time.time(),
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
)
ctx = StreamContext(
model=str(context.get("model") or "unknown"),
api_format=handler.allowed_api_formats[0],
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
)
ctx.request_id = handler.request_id
ctx.client_api_format = client_api_format
ctx.provider_type = str(context.get("provider_name") or "unknown")
ctx.update_provider_info(
provider_name=str(context.get("provider_name") or "unknown"),
provider_id=str(context.get("provider_id") or ""),
endpoint_id=str(context.get("endpoint_id") or ""),
key_id=str(context.get("key_id") or ""),
provider_api_format=str(context.get("provider_api_format") or client_api_format),
)
ctx.mapped_model = str(context.get("mapped_model") or "") or None
ctx.provider_request_headers = dict(context.get("provider_request_headers") or {})
ctx.provider_request_body = context.get("provider_request_body")
ctx.response_headers = dict(payload.headers or {})
ctx.status_code = payload.status_code
ctx.record_parsed_chunks = SystemConfigService.should_log_body(db)
if str(context.get("candidate_id") or "").strip():
ctx.attempt_id = str(context.get("candidate_id"))
proxy_info = context.get("proxy_info")
if isinstance(proxy_info, dict):
ctx.proxy_info = dict(proxy_info)
ctx.set_proxy_timing(ctx.response_headers)
telemetry = payload.telemetry if isinstance(payload.telemetry, dict) else {}
ttfb_ms = telemetry.get("ttfb_ms")
if ttfb_ms is not None:
try:
ctx.first_byte_time_ms = max(int(ttfb_ms), 0)
except (TypeError, ValueError):
ctx.first_byte_time_ms = None
if ctx.proxy_info is not None and ctx.first_byte_time_ms is not None:
ctx.set_ttfb_ms(ctx.first_byte_time_ms)
stream_processor = StreamProcessor(
request_id=ctx.request_id,
default_parser=handler.parser,
on_streaming_start=None,
)
stream = stream_processor.create_response_stream(
ctx,
_iter_gateway_report_body_chunks(body_bytes),
_GatewayReportStreamContext(),
start_time=time.time(),
)
async for _chunk in stream:
pass
elapsed_ms = telemetry.get("elapsed_ms")
try:
response_elapsed_ms = max(int(elapsed_ms or 0), 0)
except (TypeError, ValueError):
response_elapsed_ms = 0
request_start_time = (
time.time() - (response_elapsed_ms / 1000.0) if response_elapsed_ms > 0 else time.time()
)
telemetry_recorder = StreamTelemetryRecorder(
request_id=ctx.request_id,
user_id=str(user.id),
api_key_id=str(api_key.id),
client_ip="127.0.0.1",
format_id=handler.FORMAT_ID,
)
await telemetry_recorder.record_stream_stats(
ctx,
dict(context.get("original_headers") or {}),
dict(context.get("original_request_body") or {}),
request_start_time,
)
async def _record_gateway_passthrough_cli_stream_success(
payload: GatewayStreamReportRequest,
*,
db: Session,
) -> None:
from src.api.handlers.base.stream_context import StreamContext
from src.api.handlers.claude_cli import ClaudeCliAdapter
from src.api.handlers.gemini_cli import GeminiCliAdapter
from src.api.handlers.openai_cli import OpenAICliAdapter
from src.services.system.config import SystemConfigService
gateway_module = _gateway_module()
if payload.status_code >= 400:
return
context = dict(payload.report_context or {})
user_id = str(context.get("user_id") or "").strip()
api_key_id = str(context.get("api_key_id") or "").strip()
client_api_format = str(context.get("client_api_format") or "").strip().lower()
if not user_id or not api_key_id or not client_api_format:
return
body_bytes = gateway_module._extract_gateway_report_body_bytes(payload)
if not body_bytes:
return
user = db.query(User).filter(User.id == user_id).first()
api_key = db.query(ApiKey).filter(ApiKey.id == api_key_id).first()
if not user or not api_key:
return
if client_api_format == "claude:cli":
adapter = ClaudeCliAdapter()
elif client_api_format == "gemini:cli":
adapter = GeminiCliAdapter()
elif client_api_format == "openai:cli":
adapter = OpenAICliAdapter()
else:
return
handler = adapter.HANDLER_CLASS(
db=db,
user=user,
api_key=api_key,
request_id=str(context.get("request_id") or payload.trace_id or uuid.uuid4().hex[:8]),
client_ip="127.0.0.1",
user_agent=str(
(context.get("original_headers") or {}).get("user-agent") or "aether-gateway"
),
start_time=time.time(),
allowed_api_formats=adapter.allowed_api_formats,
adapter_detector=adapter.detect_capability_requirements,
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
)
ctx = StreamContext(
model=str(context.get("model") or "unknown"),
api_format=handler.primary_api_format,
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
request_id=handler.request_id,
user_id=user.id,
api_key_id=api_key.id,
)
ctx.client_api_format = client_api_format
ctx.provider_type = str(context.get("provider_name") or "unknown")
ctx.update_provider_info(
provider_name=str(context.get("provider_name") or "unknown"),
provider_id=str(context.get("provider_id") or ""),
endpoint_id=str(context.get("endpoint_id") or ""),
key_id=str(context.get("key_id") or ""),
provider_api_format=str(context.get("provider_api_format") or client_api_format),
)
ctx.mapped_model = str(context.get("mapped_model") or "") or None
ctx.provider_request_headers = dict(context.get("provider_request_headers") or {})
ctx.provider_request_body = context.get("provider_request_body")
ctx.response_headers = dict(payload.headers or {})
ctx.status_code = payload.status_code
ctx.record_parsed_chunks = SystemConfigService.should_log_body(db)
if str(context.get("candidate_id") or "").strip():
ctx.attempt_id = str(context.get("candidate_id"))
proxy_info = context.get("proxy_info")
if isinstance(proxy_info, dict):
ctx.proxy_info = dict(proxy_info)
ctx.set_proxy_timing(ctx.response_headers)
telemetry = payload.telemetry if isinstance(payload.telemetry, dict) else {}
ttfb_ms = telemetry.get("ttfb_ms")
if ttfb_ms is not None:
try:
ctx.first_byte_time_ms = max(int(ttfb_ms), 0)
except (TypeError, ValueError):
ctx.first_byte_time_ms = None
if ctx.proxy_info is not None and ctx.first_byte_time_ms is not None:
ctx.set_ttfb_ms(ctx.first_byte_time_ms)
elapsed_ms = telemetry.get("elapsed_ms")
try:
response_elapsed_ms = max(int(elapsed_ms or 0), 0)
except (TypeError, ValueError):
response_elapsed_ms = 0
if response_elapsed_ms > 0:
handler.start_time = time.time() - (response_elapsed_ms / 1000.0)
stream = handler._create_response_stream_with_prefetch(
ctx,
_iter_gateway_report_body_chunks(body_bytes),
_GatewayReportStreamContext(),
[],
)
async for _chunk in stream:
pass
await handler._record_stream_stats(
ctx,
dict(context.get("original_headers") or {}),
dict(context.get("original_request_body") or {}),
)
@@ -1,394 +0,0 @@
from __future__ import annotations
import inspect
import time
import uuid
from datetime import datetime, timezone
from typing import Any
from fastapi import BackgroundTasks
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.models.database import ApiKey, RequestCandidate, User
from .gateway_contract import GatewayStreamReportRequest, GatewaySyncReportRequest
from .gateway_reporting_common import _gateway_module
def _postprocess_gateway_report_provider_response(
context: dict[str, Any],
provider_response_json: dict[str, Any],
) -> None:
if str(context.get("envelope_name") or "").strip().lower() != "antigravity:v1internal":
return
try:
from src.services.provider.adapters.antigravity.envelope import (
_inject_claude_tool_ids_response,
cache_thought_signatures,
)
model = str(context.get("mapped_model") or context.get("model") or "")
_inject_claude_tool_ids_response(provider_response_json, model)
cache_thought_signatures(model, provider_response_json)
except Exception:
return
async def _record_gateway_openai_chat_sync_success(
payload: GatewaySyncReportRequest,
*,
db: Session,
) -> None:
from src.api.handlers.openai import OpenAIChatAdapter
if payload.status_code >= 400:
return
if not isinstance(payload.body_json, dict) or payload.body_json.get("error") is not None:
return
context = dict(payload.report_context or {})
user_id = str(context.get("user_id") or "").strip()
api_key_id = str(context.get("api_key_id") or "").strip()
if not user_id or not api_key_id:
return
user = db.query(User).filter(User.id == user_id).first()
api_key = db.query(ApiKey).filter(ApiKey.id == api_key_id).first()
if not user or not api_key:
return
adapter = OpenAIChatAdapter()
handler = adapter._create_handler(
db=db,
user=user,
api_key=api_key,
request_id=str(context.get("request_id") or payload.trace_id or uuid.uuid4().hex[:8]),
client_ip="127.0.0.1",
user_agent=str(
(context.get("original_headers") or {}).get("user-agent") or "aether-gateway"
),
start_time=time.time(),
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
)
provider_response_json = dict(payload.body_json)
_postprocess_gateway_report_provider_response(context, provider_response_json)
client_response_json = (
dict(payload.client_body_json) if isinstance(payload.client_body_json, dict) else None
)
response_json = handler._normalize_response(client_response_json or provider_response_json)
usage_info = handler._extract_usage(response_json)
request_metadata: dict[str, Any] = {
"gateway_direct_executor": True,
"phase": "3c_trial",
}
proxy_info = context.get("proxy_info")
if isinstance(proxy_info, dict):
request_metadata["proxy"] = proxy_info
response_time_ms = 0
if isinstance(payload.telemetry, dict):
raw_elapsed_ms = payload.telemetry.get("elapsed_ms")
try:
response_time_ms = max(int(raw_elapsed_ms or 0), 0)
except (TypeError, ValueError):
response_time_ms = 0
await handler.telemetry.record_success(
provider=str(context.get("provider_name") or "openai"),
model=str(context.get("model") or "unknown"),
input_tokens=int(usage_info.get("input_tokens", 0) or 0),
output_tokens=int(usage_info.get("output_tokens", 0) or 0),
response_time_ms=response_time_ms,
status_code=payload.status_code,
request_headers=dict(context.get("original_headers") or {}),
request_body=dict(context.get("original_request_body") or {}),
response_headers=dict(payload.headers or {}),
client_response_headers=dict(payload.headers or {}),
response_body=provider_response_json if client_response_json else response_json,
client_response_body=response_json if client_response_json else None,
provider_request_headers=dict(context.get("provider_request_headers") or {}),
provider_request_body=context.get("provider_request_body"),
is_stream=False,
provider_id=str(context.get("provider_id") or "") or None,
provider_endpoint_id=str(context.get("endpoint_id") or "") or None,
provider_api_key_id=str(context.get("key_id") or "") or None,
api_format=str(context.get("client_api_format") or "openai:chat"),
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
endpoint_api_format=str(context.get("provider_api_format") or "") or None,
has_format_conversion=client_response_json is not None,
target_model=str(context.get("mapped_model") or "") or None,
request_metadata=request_metadata,
)
async def _record_gateway_passthrough_chat_sync_success(
payload: GatewaySyncReportRequest,
*,
db: Session,
) -> None:
from src.api.handlers.claude import ClaudeChatAdapter
from src.api.handlers.gemini import GeminiChatAdapter
from src.api.handlers.openai import OpenAIChatAdapter
if payload.status_code >= 400:
return
if not isinstance(payload.body_json, dict) or payload.body_json.get("error") is not None:
return
context = dict(payload.report_context or {})
user_id = str(context.get("user_id") or "").strip()
api_key_id = str(context.get("api_key_id") or "").strip()
client_api_format = str(context.get("client_api_format") or "").strip().lower()
if not user_id or not api_key_id or not client_api_format:
return
user = db.query(User).filter(User.id == user_id).first()
api_key = db.query(ApiKey).filter(ApiKey.id == api_key_id).first()
if not user or not api_key:
return
if client_api_format == "claude:chat":
adapter = ClaudeChatAdapter()
elif client_api_format == "gemini:chat":
adapter = GeminiChatAdapter()
elif client_api_format == "openai:chat":
adapter = OpenAIChatAdapter()
else:
return
handler = adapter._create_handler(
db=db,
user=user,
api_key=api_key,
request_id=str(context.get("request_id") or payload.trace_id or uuid.uuid4().hex[:8]),
client_ip="127.0.0.1",
user_agent=str(
(context.get("original_headers") or {}).get("user-agent") or "aether-gateway"
),
start_time=time.time(),
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
)
provider_response_json = dict(payload.body_json)
_postprocess_gateway_report_provider_response(context, provider_response_json)
response_json = handler._normalize_response(provider_response_json)
usage_info = handler._extract_usage(response_json)
request_metadata: dict[str, Any] = {
"gateway_direct_executor": True,
"phase": "3c_trial",
}
proxy_info = context.get("proxy_info")
if isinstance(proxy_info, dict):
request_metadata["proxy"] = proxy_info
response_time_ms = 0
if isinstance(payload.telemetry, dict):
raw_elapsed_ms = payload.telemetry.get("elapsed_ms")
try:
response_time_ms = max(int(raw_elapsed_ms or 0), 0)
except (TypeError, ValueError):
response_time_ms = 0
await handler.telemetry.record_success(
provider=str(context.get("provider_name") or "unknown"),
model=str(context.get("model") or "unknown"),
input_tokens=int(usage_info.get("input_tokens", 0) or 0),
output_tokens=int(usage_info.get("output_tokens", 0) or 0),
response_time_ms=response_time_ms,
status_code=payload.status_code,
request_headers=dict(context.get("original_headers") or {}),
request_body=dict(context.get("original_request_body") or {}),
response_headers=dict(payload.headers or {}),
client_response_headers=dict(payload.headers or {}),
response_body=response_json,
provider_request_headers=dict(context.get("provider_request_headers") or {}),
provider_request_body=context.get("provider_request_body"),
cache_creation_tokens=int(
usage_info.get("cache_creation_input_tokens", 0)
or usage_info.get("cache_creation_tokens", 0)
or 0
),
cache_read_tokens=int(
usage_info.get("cache_read_input_tokens", 0)
or usage_info.get("cache_read_tokens", 0)
or 0
),
is_stream=False,
provider_id=str(context.get("provider_id") or "") or None,
provider_endpoint_id=str(context.get("endpoint_id") or "") or None,
provider_api_key_id=str(context.get("key_id") or "") or None,
api_format=client_api_format,
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
endpoint_api_format=str(context.get("provider_api_format") or "") or None,
has_format_conversion=False,
target_model=str(context.get("mapped_model") or "") or None,
request_metadata=request_metadata,
)
async def _record_gateway_passthrough_cli_sync_success(
payload: GatewaySyncReportRequest,
*,
db: Session,
) -> None:
from src.api.handlers.claude_cli import ClaudeCliAdapter
from src.api.handlers.gemini_cli import GeminiCliAdapter
from src.api.handlers.openai_cli import OpenAICliAdapter, OpenAICompactAdapter
if payload.status_code >= 400:
return
if not isinstance(payload.body_json, dict) or payload.body_json.get("error") is not None:
return
context = dict(payload.report_context or {})
user_id = str(context.get("user_id") or "").strip()
api_key_id = str(context.get("api_key_id") or "").strip()
client_api_format = str(context.get("client_api_format") or "openai:cli").strip().lower()
if not user_id or not api_key_id:
return
user = db.query(User).filter(User.id == user_id).first()
api_key = db.query(ApiKey).filter(ApiKey.id == api_key_id).first()
if not user or not api_key:
return
if client_api_format == "openai:compact":
adapter = OpenAICompactAdapter()
elif client_api_format == "claude:cli":
adapter = ClaudeCliAdapter()
elif client_api_format == "gemini:cli":
adapter = GeminiCliAdapter()
else:
adapter = OpenAICliAdapter()
handler = adapter.HANDLER_CLASS(
db=db,
user=user,
api_key=api_key,
request_id=str(context.get("request_id") or payload.trace_id or uuid.uuid4().hex[:8]),
client_ip="127.0.0.1",
user_agent=str(
(context.get("original_headers") or {}).get("user-agent") or "aether-gateway"
),
start_time=time.time(),
allowed_api_formats=adapter.allowed_api_formats,
adapter_detector=adapter.detect_capability_requirements,
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
)
provider_response_json = dict(payload.body_json)
_postprocess_gateway_report_provider_response(context, provider_response_json)
client_response_json = (
dict(payload.client_body_json) if isinstance(payload.client_body_json, dict) else None
)
response_json = dict(client_response_json or provider_response_json)
usage_info = handler.parser.extract_usage_from_response(response_json)
response_metadata = handler._extract_response_metadata(response_json)
request_metadata: dict[str, Any] = {
"gateway_direct_executor": True,
"phase": "3c_trial",
}
proxy_info = context.get("proxy_info")
if isinstance(proxy_info, dict):
request_metadata["proxy"] = proxy_info
response_time_ms = 0
if isinstance(payload.telemetry, dict):
raw_elapsed_ms = payload.telemetry.get("elapsed_ms")
try:
response_time_ms = max(int(raw_elapsed_ms or 0), 0)
except (TypeError, ValueError):
response_time_ms = 0
await handler.telemetry.record_success(
provider=str(context.get("provider_name") or "openai"),
model=str(context.get("model") or "unknown"),
input_tokens=int(usage_info.get("input_tokens", 0) or 0),
output_tokens=int(usage_info.get("output_tokens", 0) or 0),
response_time_ms=response_time_ms,
status_code=payload.status_code,
request_headers=dict(context.get("original_headers") or {}),
request_body=dict(context.get("original_request_body") or {}),
response_headers=dict(payload.headers or {}),
client_response_headers=dict(payload.headers or {}),
response_body=provider_response_json if client_response_json else response_json,
client_response_body=response_json if client_response_json else None,
provider_request_headers=dict(context.get("provider_request_headers") or {}),
provider_request_body=context.get("provider_request_body"),
cache_creation_tokens=int(usage_info.get("cache_creation_tokens", 0) or 0),
cache_read_tokens=int(usage_info.get("cache_read_tokens", 0) or 0),
is_stream=False,
provider_id=str(context.get("provider_id") or "") or None,
provider_endpoint_id=str(context.get("endpoint_id") or "") or None,
provider_api_key_id=str(context.get("key_id") or "") or None,
api_format=client_api_format,
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
endpoint_api_format=str(context.get("provider_api_format") or "") or None,
has_format_conversion=client_response_json is not None,
target_model=str(context.get("mapped_model") or "") or None,
response_metadata=response_metadata if response_metadata else None,
request_metadata=request_metadata,
)
async def _record_gateway_openai_video_delete_sync_success(
payload: GatewaySyncReportRequest,
*,
db: Session,
) -> None:
gateway_module = _gateway_module()
await gateway_module._finalize_gateway_openai_video_delete_sync(payload, db=db)
async def _record_gateway_openai_video_cancel_sync_success(
payload: GatewaySyncReportRequest,
*,
db: Session,
) -> None:
gateway_module = _gateway_module()
await gateway_module._finalize_gateway_openai_video_cancel_sync(payload, db=db)
async def _record_gateway_gemini_video_cancel_sync_success(
payload: GatewaySyncReportRequest,
*,
db: Session,
) -> None:
gateway_module = _gateway_module()
await gateway_module._finalize_gateway_gemini_video_cancel_sync(payload, db=db)
async def _record_gateway_openai_video_create_sync_success(
payload: GatewaySyncReportRequest,
*,
db: Session,
) -> None:
gateway_module = _gateway_module()
await gateway_module._finalize_gateway_openai_video_create_sync(payload, db=db)
async def _record_gateway_openai_video_remix_sync_success(
payload: GatewaySyncReportRequest,
*,
db: Session,
) -> None:
gateway_module = _gateway_module()
await gateway_module._finalize_gateway_openai_video_remix_sync(payload, db=db)
async def _record_gateway_gemini_video_create_sync_success(
payload: GatewaySyncReportRequest,
*,
db: Session,
) -> None:
gateway_module = _gateway_module()
await gateway_module._finalize_gateway_gemini_video_create_sync(payload, db=db)
@@ -1,125 +0,0 @@
from __future__ import annotations
import inspect
import time
import uuid
from datetime import datetime, timezone
from typing import Any
from fastapi import BackgroundTasks
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.models.database import ApiKey, RequestCandidate, User
from .gateway_contract import GatewayStreamReportRequest, GatewaySyncReportRequest
from .gateway_reporting_common import _gateway_module
def _build_gateway_sync_telemetry_writer(
*,
db: Session,
request_id: str,
user_id: str,
api_key_id: str,
fallback_telemetry: Any,
) -> Any:
from src.config.settings import config
from src.services.system.config import SystemConfigService
from src.services.usage.telemetry_writer import DbTelemetryWriter, QueueTelemetryWriter
if config.usage_queue_enabled and user_id and api_key_id:
try:
log_level = SystemConfigService.get_request_record_level(db).value
sensitive_headers = SystemConfigService.get_sensitive_headers(db) or []
max_request_body_size = int(
SystemConfigService.get_config(db, "max_request_body_size", 5242880) or 0
)
max_response_body_size = int(
SystemConfigService.get_config(db, "max_response_body_size", 5242880) or 0
)
return QueueTelemetryWriter(
request_id=request_id,
user_id=user_id,
api_key_id=api_key_id,
log_level=log_level,
sensitive_headers=sensitive_headers,
max_request_body_size=max_request_body_size,
max_response_body_size=max_response_body_size,
)
except Exception:
return DbTelemetryWriter(fallback_telemetry)
return DbTelemetryWriter(fallback_telemetry)
async def _dispatch_gateway_sync_telemetry(
*,
telemetry_writer: Any,
operation: str,
**kwargs: Any,
) -> None:
gateway_module = _gateway_module()
submitter = getattr(telemetry_writer, operation, None)
if not callable(submitter):
raise AttributeError(f"Telemetry writer missing operation: {operation}")
if bool(getattr(telemetry_writer, "supports_background_submission", lambda: False)()):
request_id = str(getattr(telemetry_writer, "request_id", "") or "unknown")
async def _run_in_background() -> None:
try:
await submitter(**kwargs)
except Exception as exc:
logger.warning(
"[gateway] background telemetry submission failed: request_id={}, operation={}, error={}",
request_id,
operation,
exc,
)
task = gateway_module.safe_create_task(_run_in_background())
if task is None:
await submitter(**kwargs)
return
await submitter(**kwargs)
async def _schedule_gateway_sync_telemetry(
*,
background_tasks: BackgroundTasks | None,
telemetry_writer: Any,
operation: str,
**kwargs: Any,
) -> None:
gateway_module = _gateway_module()
if background_tasks is not None:
background_tasks.add_task(
gateway_module._dispatch_gateway_sync_telemetry,
telemetry_writer=telemetry_writer,
operation=operation,
**kwargs,
)
return
await gateway_module._dispatch_gateway_sync_telemetry(
telemetry_writer=telemetry_writer,
operation=operation,
**kwargs,
)
def _build_gateway_usage_metadata(
*,
request_metadata: dict[str, Any] | None = None,
response_metadata: dict[str, Any] | None = None,
) -> dict[str, Any] | None:
metadata: dict[str, Any] | None = None
if request_metadata:
metadata = dict(request_metadata)
if response_metadata:
metadata.setdefault("response", response_metadata)
elif response_metadata:
metadata = dict(response_metadata)
return metadata
@@ -1,386 +0,0 @@
from __future__ import annotations
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Request
from fastapi.responses import JSONResponse, Response
from sqlalchemy.orm import Session
from src.database import get_db
from . import gateway as gateway_impl
from .gateway_contract import (
CONTROL_EXECUTED_HEADER,
GatewayAuthContextRequest,
GatewayExecuteRequest,
GatewayResolveRequest,
GatewayStreamReportRequest,
GatewaySyncReportRequest,
classify_gateway_route,
)
router = APIRouter(
prefix="/api/internal/gateway",
tags=["Internal - Gateway"],
include_in_schema=False,
)
def _ensure_legacy_internal_gateway_request(request: Request) -> Response | None:
if gateway_impl._request_allows_legacy_chat_cli_internal_gateway(request):
return None
return gateway_impl._build_retired_internal_gateway_response()
@router.post("/resolve")
async def resolve_gateway_route(
request: Request, payload: GatewayResolveRequest
) -> Response:
gateway_impl.ensure_loopback(request)
retired = _ensure_legacy_internal_gateway_request(request)
if retired is not None:
return retired
decision = classify_gateway_route(payload.method, payload.path, payload.headers)
decision.auth_context = await gateway_impl._resolve_auth_context(payload, decision)
return JSONResponse(status_code=200, content=decision.model_dump(exclude_none=True))
@router.post("/auth-context")
async def resolve_gateway_auth_context(
request: Request,
payload: GatewayAuthContextRequest,
) -> Response:
gateway_impl.ensure_loopback(request)
retired = _ensure_legacy_internal_gateway_request(request)
if retired is not None:
return retired
auth_context = await gateway_impl._resolve_auth_context_signature(
headers=payload.headers,
query_string=payload.query_string,
auth_endpoint_signature=payload.auth_endpoint_signature,
)
return JSONResponse(status_code=200, content={"auth_context": auth_context})
@router.post("/decision-sync")
async def decide_gateway_sync(
request: Request,
payload: GatewayExecuteRequest,
db: Session = Depends(get_db),
) -> Response:
gateway_impl.ensure_loopback(request)
retired = _ensure_legacy_internal_gateway_request(request)
if retired is not None:
return retired
decision = classify_gateway_route(payload.method, payload.path, payload.headers)
if not gateway_impl._allows_legacy_chat_cli_internal_route(request, decision):
return gateway_impl._build_retired_internal_gateway_response()
auth_context = await gateway_impl._resolve_gateway_execute_auth_context(
payload=payload,
decision=decision,
)
if auth_context is None or not auth_context.access_allowed:
return JSONResponse(status_code=200, content={"action": "fallback_plan"})
try:
resolved = await gateway_impl._build_gateway_sync_decision_response(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
except HTTPException as exc:
headers = dict(exc.headers or {})
headers[CONTROL_EXECUTED_HEADER] = "true"
return JSONResponse(
status_code=exc.status_code,
content={"detail": exc.detail},
headers=headers,
)
if resolved is None:
body: dict[str, object] = {"action": "fallback_plan"}
if payload.auth_context is None:
body["auth_context"] = auth_context.model_dump(exclude_none=True)
return JSONResponse(status_code=200, content=body)
if payload.auth_context is not None:
resolved.auth_context = None
return JSONResponse(status_code=200, content=resolved.model_dump(exclude_none=True))
@router.post("/decision-stream")
async def decide_gateway_stream(
request: Request,
payload: GatewayExecuteRequest,
db: Session = Depends(get_db),
) -> Response:
gateway_impl.ensure_loopback(request)
retired = _ensure_legacy_internal_gateway_request(request)
if retired is not None:
return retired
decision = classify_gateway_route(payload.method, payload.path, payload.headers)
if not gateway_impl._allows_legacy_chat_cli_internal_route(request, decision):
return gateway_impl._build_retired_internal_gateway_response()
auth_context = await gateway_impl._resolve_gateway_execute_auth_context(
payload=payload,
decision=decision,
)
if auth_context is None or not auth_context.access_allowed:
return JSONResponse(status_code=200, content={"action": "fallback_plan"})
try:
resolved = await gateway_impl._build_gateway_stream_decision_response(
request=request,
payload=payload,
db=db,
auth_context=auth_context,
decision=decision,
)
except HTTPException as exc:
headers = dict(exc.headers or {})
headers[CONTROL_EXECUTED_HEADER] = "true"
return JSONResponse(
status_code=exc.status_code,
content={"detail": exc.detail},
headers=headers,
)
if resolved is None:
body: dict[str, object] = {"action": "fallback_plan"}
if payload.auth_context is None:
body["auth_context"] = auth_context.model_dump(exclude_none=True)
return JSONResponse(status_code=200, content=body)
if payload.auth_context is not None:
resolved.auth_context = None
return JSONResponse(status_code=200, content=resolved.model_dump(exclude_none=True))
@router.post("/execute-sync")
async def execute_gateway_sync(
request: Request,
payload: GatewayExecuteRequest,
db: Session = Depends(get_db),
) -> Response:
retired = _ensure_legacy_internal_gateway_request(request)
if retired is not None:
return retired
return await gateway_impl._execute_gateway_control_request(
request=request,
payload=payload,
db=db,
require_stream=False,
)
@router.post("/execute-stream")
async def execute_gateway_stream(
request: Request,
payload: GatewayExecuteRequest,
db: Session = Depends(get_db),
) -> Response:
retired = _ensure_legacy_internal_gateway_request(request)
if retired is not None:
return retired
return await gateway_impl._execute_gateway_control_request(
request=request,
payload=payload,
db=db,
require_stream=True,
)
@router.post("/plan-stream")
async def plan_gateway_stream(
request: Request,
payload: GatewayExecuteRequest,
db: Session = Depends(get_db),
) -> Response:
gateway_impl.ensure_loopback(request)
retired = _ensure_legacy_internal_gateway_request(request)
if retired is not None:
return retired
decision = classify_gateway_route(payload.method, payload.path, payload.headers)
if not gateway_impl._allows_legacy_chat_cli_internal_route(request, decision):
return gateway_impl._build_retired_internal_gateway_response()
try:
planned = await gateway_impl._build_gateway_stream_plan_response(
request=request,
payload=payload,
db=db,
)
except HTTPException as exc:
headers = dict(exc.headers or {})
headers[CONTROL_EXECUTED_HEADER] = "true"
return JSONResponse(
status_code=exc.status_code,
content={"detail": exc.detail},
headers=headers,
)
if planned is None:
response = await gateway_impl._execute_gateway_control_request(
request=request,
payload=payload,
db=db,
require_stream=True,
)
if gateway_impl._is_gateway_control_executed_response(response):
return response
return gateway_impl._build_proxy_public_fallback_response()
if payload.auth_context is not None:
planned.auth_context = None
return JSONResponse(status_code=200, content=planned.model_dump(exclude_none=True))
@router.post("/plan-sync")
async def plan_gateway_sync(
request: Request,
payload: GatewayExecuteRequest,
db: Session = Depends(get_db),
) -> Response:
gateway_impl.ensure_loopback(request)
retired = _ensure_legacy_internal_gateway_request(request)
if retired is not None:
return retired
decision = classify_gateway_route(payload.method, payload.path, payload.headers)
if not gateway_impl._allows_legacy_chat_cli_internal_route(request, decision):
return gateway_impl._build_retired_internal_gateway_response()
try:
planned = await gateway_impl._build_gateway_sync_plan_response(
request=request,
payload=payload,
db=db,
)
except HTTPException as exc:
headers = dict(exc.headers or {})
headers[CONTROL_EXECUTED_HEADER] = "true"
return JSONResponse(
status_code=exc.status_code,
content={"detail": exc.detail},
headers=headers,
)
if planned is None:
response = await gateway_impl._execute_gateway_control_request(
request=request,
payload=payload,
db=db,
require_stream=False,
)
if gateway_impl._is_gateway_control_executed_response(response):
return response
return gateway_impl._build_proxy_public_fallback_response()
if payload.auth_context is not None:
planned.auth_context = None
return JSONResponse(status_code=200, content=planned.model_dump(exclude_none=True))
@router.post("/report-sync")
async def report_gateway_sync(
request: Request,
payload: GatewaySyncReportRequest,
background_tasks: BackgroundTasks,
) -> Response:
gateway_impl.ensure_loopback(request)
retired = _ensure_legacy_internal_gateway_request(request)
if retired is not None:
return retired
if not gateway_impl._allows_legacy_chat_cli_report_route(request, payload.report_kind):
return gateway_impl._build_retired_internal_gateway_response()
payload_copy = payload.model_copy(deep=True)
if gateway_impl._gateway_sync_report_requires_inline(payload_copy):
db, cleanup = gateway_impl._resolve_gateway_background_db(getattr(request, "app", None))
try:
await gateway_impl._run_gateway_sync_report_background(payload_copy, db)
finally:
if cleanup is not None:
cleanup()
return JSONResponse(status_code=200, content={"ok": True})
background_tasks.add_task(
gateway_impl._run_gateway_sync_report_background_with_session,
payload_copy,
getattr(request, "app", None),
)
return JSONResponse(status_code=200, content={"ok": True})
@router.post("/finalize-sync")
async def finalize_gateway_sync(
request: Request,
payload: GatewaySyncReportRequest,
background_tasks: BackgroundTasks,
) -> Response:
gateway_impl.ensure_loopback(request)
retired = _ensure_legacy_internal_gateway_request(request)
if retired is not None:
return retired
if not gateway_impl._allows_legacy_chat_cli_report_route(request, payload.report_kind):
return gateway_impl._build_retired_internal_gateway_response()
fast_response = await gateway_impl._maybe_build_gateway_core_sync_fast_success_response(payload)
if fast_response is not None:
if payload.report_kind in {
"openai_chat_sync_finalize",
"claude_chat_sync_finalize",
"gemini_chat_sync_finalize",
}:
background_tasks.add_task(
gateway_impl._run_gateway_chat_sync_finalize_background_with_session,
payload.model_copy(deep=True),
)
elif payload.report_kind in {
"openai_cli_sync_finalize",
"openai_compact_sync_finalize",
"claude_cli_sync_finalize",
"gemini_cli_sync_finalize",
}:
background_tasks.add_task(
gateway_impl._run_gateway_cli_sync_finalize_background_with_session,
payload.model_copy(deep=True),
)
fast_response.headers[CONTROL_EXECUTED_HEADER] = "true"
return fast_response
db, cleanup = gateway_impl._resolve_gateway_finalize_db(request)
try:
response = await gateway_impl._finalize_gateway_sync_response(
payload,
db=db,
background_tasks=background_tasks,
)
except Exception:
if cleanup is not None:
cleanup()
raise
if cleanup is not None:
background_tasks.add_task(cleanup)
response.headers[CONTROL_EXECUTED_HEADER] = "true"
return response
@router.post("/report-stream")
async def report_gateway_stream(
request: Request,
payload: GatewayStreamReportRequest,
background_tasks: BackgroundTasks,
) -> Response:
gateway_impl.ensure_loopback(request)
retired = _ensure_legacy_internal_gateway_request(request)
if retired is not None:
return retired
if not gateway_impl._allows_legacy_chat_cli_report_route(request, payload.report_kind):
return gateway_impl._build_retired_internal_gateway_response()
background_tasks.add_task(
gateway_impl._run_gateway_stream_report_background_with_session,
payload.model_copy(deep=True),
getattr(request, "app", None),
)
return JSONResponse(status_code=200, content={"ok": True})
@@ -1,594 +0,0 @@
from __future__ import annotations
import base64
import time
import uuid
from dataclasses import asdict
from datetime import datetime, timezone
from typing import Any
from fastapi import BackgroundTasks, HTTPException, Request
from fastapi.responses import JSONResponse, Response
from sqlalchemy.orm import Session
from src.api.base.context import ApiRequestContext
from src.api.base.pipeline import get_pipeline
from src.core.api_format.headers import extract_client_api_key_for_endpoint_with_query
from src.core.api_format.metadata import get_auth_config_for_endpoint
from src.core.crypto import crypto_service
from src.core.http_compression import normalize_content_encoding
from src.core.logger import logger
from src.database import create_session, get_db
from src.models.database import ApiKey, RequestCandidate, User
from src.services.auth.service import AuthService
from src.utils.async_utils import safe_create_task
from .common import ensure_loopback
from .gateway_contract import (
_GEMINI_FILES_DOWNLOAD_ROUTE_RE,
_GEMINI_FILES_RESOURCE_ROUTE_RE,
_GEMINI_MODEL_OPERATION_CANCEL_RE,
_GEMINI_OPERATION_CANCEL_RE,
_GEMINI_SYNC_ROUTE_RE,
_GEMINI_VIDEO_CREATE_ROUTE_RE,
_GEMINI_VIDEO_MODEL_OPERATION_ANY_RE,
_OPENAI_VIDEO_CANCEL_ROUTE_RE,
_OPENAI_VIDEO_CONTENT_ROUTE_RE,
_OPENAI_VIDEO_REMIX_ROUTE_RE,
_OPENAI_VIDEO_TASK_ROUTE_RE,
CONTROL_ACTION_HEADER,
CONTROL_ACTION_PROXY_PUBLIC,
CONTROL_EXECUTED_HEADER,
GatewayAuthContext,
GatewayExecuteRequest,
GatewayExecutionDecisionResponse,
GatewayExecutionPlanResponse,
GatewayResolveRequest,
GatewayRouteDecision,
GatewayStreamReportRequest,
GatewaySyncReportRequest,
classify_gateway_route,
)
class _GatewayProxy:
def __getattr__(self, name: str) -> Any:
from . import gateway as gateway_module
return getattr(gateway_module, name)
gateway_module = _GatewayProxy()
LEGACY_CHAT_CLI_CONTROL_EXECUTE_FALLBACK_HEADER = "x-aether-control-execute-fallback"
LEGACY_CHAT_CLI_INTERNAL_GATEWAY_HEADER = "x-aether-legacy-internal-gateway"
_LEGACY_CHAT_CLI_INTERNAL_GATEWAY_TRUE_VALUES = {"1", "true", "yes", "on"}
_LEGACY_CHAT_CLI_REPORT_KIND_PREFIXES = (
"openai_chat_",
"claude_chat_",
"gemini_chat_",
"openai_cli_",
"openai_compact_",
"claude_cli_",
"gemini_cli_",
)
def _is_legacy_chat_cli_control_route(decision: GatewayRouteDecision) -> bool:
return decision.route_class == "ai_public" and decision.route_kind in {
"chat",
"cli",
"compact",
}
def _request_allows_legacy_chat_cli_internal_gateway(request: Request) -> bool:
for key, value in request.headers.items():
normalized_key = str(key or "").strip().lower()
if normalized_key not in {
LEGACY_CHAT_CLI_INTERNAL_GATEWAY_HEADER,
LEGACY_CHAT_CLI_CONTROL_EXECUTE_FALLBACK_HEADER,
}:
continue
return (
str(value or "").strip().lower()
in _LEGACY_CHAT_CLI_INTERNAL_GATEWAY_TRUE_VALUES
)
return False
def _allows_legacy_chat_cli_internal_route(
request: Request,
decision: GatewayRouteDecision,
) -> bool:
if not _is_legacy_chat_cli_control_route(decision):
return True
return _request_allows_legacy_chat_cli_internal_gateway(request)
def _is_legacy_chat_cli_report_kind(report_kind: str | None) -> bool:
normalized = str(report_kind or "").strip().lower()
if not normalized:
return False
return normalized.startswith(_LEGACY_CHAT_CLI_REPORT_KIND_PREFIXES)
def _allows_legacy_chat_cli_report_route(request: Request, report_kind: str | None) -> bool:
if not _is_legacy_chat_cli_report_kind(report_kind):
return True
return _request_allows_legacy_chat_cli_internal_gateway(request)
def _build_retired_internal_gateway_response() -> JSONResponse:
return JSONResponse(
status_code=410,
content={"detail": "legacy internal gateway route removed; use public proxy"},
)
def _allows_legacy_chat_cli_control_execute(
request: Request,
payload: GatewayExecuteRequest,
decision: GatewayRouteDecision,
) -> bool:
if not _is_legacy_chat_cli_control_route(decision):
return True
if not _allows_legacy_chat_cli_internal_route(request, decision):
return False
for key, value in (payload.headers or {}).items():
if str(key or "").strip().lower() != LEGACY_CHAT_CLI_CONTROL_EXECUTE_FALLBACK_HEADER:
continue
return str(value or "").strip().lower() in {"1", "true", "yes", "on"}
return False
async def _execute_gateway_control_request(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
require_stream: bool,
) -> Response:
gateway_module.ensure_loopback(request)
decision = classify_gateway_route(payload.method, payload.path, payload.headers)
if _is_gemini_files_route(decision):
return await _execute_gateway_files_control_request(
request=request,
payload=payload,
db=db,
require_stream=require_stream,
)
if not _allows_legacy_chat_cli_control_execute(request, payload, decision):
return gateway_module._build_retired_internal_gateway_response()
adapter, path_params = gateway_module._resolve_gateway_sync_adapter(decision, payload.path)
if adapter is None:
return gateway_module._build_proxy_public_fallback_response()
is_stream_request = _is_stream_request_payload(payload.body_json, path_params)
if is_stream_request != require_stream:
return gateway_module._build_proxy_public_fallback_response()
auth_context = await gateway_module._resolve_gateway_execute_auth_context(
payload=payload,
decision=decision,
)
if auth_context is None or not auth_context.access_allowed:
return gateway_module._build_proxy_public_fallback_response()
try:
effective_request = (
gateway_module._build_gateway_forward_request(request=request, payload=payload)
if _is_video_route(decision)
else request
)
user, api_key = gateway_module._load_gateway_auth_models(db, auth_context)
pipeline = gateway_module.get_pipeline()
await pipeline._check_user_rate_limit(effective_request, db, user, api_key)
context = gateway_module._build_gateway_request_context(
request=effective_request,
payload=payload,
db=db,
user=user,
api_key=api_key,
adapter=adapter,
path_params=path_params,
balance_remaining=auth_context.balance_remaining,
)
authorize_result = adapter.authorize(context)
if hasattr(authorize_result, "__await__"):
await authorize_result
response = await adapter.handle(context)
except HTTPException as exc:
headers = dict(exc.headers or {})
headers[CONTROL_EXECUTED_HEADER] = "true"
return JSONResponse(
status_code=exc.status_code,
content={"detail": exc.detail},
headers=headers,
)
response.headers[CONTROL_EXECUTED_HEADER] = "true"
return response
def _is_gateway_control_executed_response(response: Response) -> bool:
return str(response.headers.get(CONTROL_EXECUTED_HEADER) or "").strip().lower() == "true"
async def _execute_gateway_files_control_request(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
require_stream: bool,
) -> Response:
if require_stream:
return gateway_module._build_proxy_public_fallback_response()
decision = classify_gateway_route(payload.method, payload.path, payload.headers)
auth_context = await gateway_module._resolve_gateway_execute_auth_context(
payload=payload,
decision=decision,
)
if auth_context is None or not auth_context.access_allowed:
return gateway_module._build_proxy_public_fallback_response()
try:
user, api_key = gateway_module._load_gateway_auth_models(db, auth_context)
gateway_request = gateway_module._build_gateway_forward_request(
request=request, payload=payload
)
pipeline = gateway_module.get_pipeline()
await pipeline._check_user_rate_limit(gateway_request, db, user, api_key)
response = await _dispatch_gateway_files_handler(gateway_request, payload)
except HTTPException as exc:
headers = dict(exc.headers or {})
headers[CONTROL_EXECUTED_HEADER] = "true"
return JSONResponse(
status_code=exc.status_code,
content={"detail": exc.detail},
headers=headers,
)
response.headers[CONTROL_EXECUTED_HEADER] = "true"
return response
def _is_gemini_files_route(decision: GatewayRouteDecision) -> bool:
return (
decision.route_class == "ai_public"
and decision.route_family == "gemini"
and decision.route_kind == "files"
)
def _is_video_route(decision: GatewayRouteDecision) -> bool:
return decision.route_class == "ai_public" and decision.route_kind == "video"
async def _dispatch_gateway_files_handler(
request: Request,
payload: GatewayExecuteRequest,
) -> Response:
from src.api.public.gemini_files import (
delete_file,
download_file,
get_file,
list_files,
upload_file,
)
method = str(payload.method or "").strip().upper()
path = str(payload.path or "").strip()
if method == "POST" and path == "/upload/v1beta/files":
return await upload_file(request)
if method == "GET" and path == "/v1beta/files":
query_params = gateway_module._parse_query_string(payload.query_string)
page_size: int | None = None
if query_params.get("pageSize") not in {None, ""}:
try:
page_size = int(str(query_params["pageSize"]))
except ValueError as exc:
raise HTTPException(status_code=400, detail="Invalid pageSize") from exc
return await list_files(
request,
pageSize=page_size,
pageToken=query_params.get("pageToken"),
)
download_match = _GEMINI_FILES_DOWNLOAD_ROUTE_RE.match(path)
if method == "GET" and download_match:
return await download_file(download_match.group("file_id"), request)
resource_match = _GEMINI_FILES_RESOURCE_ROUTE_RE.match(path)
if resource_match:
file_name = resource_match.group("file_name")
if method == "GET":
return await get_file(file_name, request)
if method == "DELETE":
return await delete_file(file_name, request)
return gateway_module._build_proxy_public_fallback_response()
def _extract_gemini_path_params(path: str) -> dict[str, Any]:
match = _GEMINI_SYNC_ROUTE_RE.match(str(path or "").strip())
if not match:
return {}
action = str(match.group("action") or "").strip()
return {
"model": str(match.group("model") or "").strip(),
"stream": action == "streamGenerateContent",
}
def _extract_openai_video_path_params(path: str) -> dict[str, Any]:
normalized_path = str(path or "").strip()
for route_re, extra in (
(_OPENAI_VIDEO_CANCEL_ROUTE_RE, {"action": "cancel"}),
(_OPENAI_VIDEO_REMIX_ROUTE_RE, {}),
(_OPENAI_VIDEO_CONTENT_ROUTE_RE, {}),
(_OPENAI_VIDEO_TASK_ROUTE_RE, {}),
):
match = route_re.match(normalized_path)
if match:
params = {"task_id": str(match.group("task_id") or "").strip()}
params.update(extra)
return params
return {}
def _extract_gemini_video_path_params(path: str) -> dict[str, Any]:
normalized_path = str(path or "").strip()
create_match = _GEMINI_VIDEO_CREATE_ROUTE_RE.match(normalized_path)
if create_match:
return {"model": str(create_match.group("model") or "").strip()}
model_operation_match = _GEMINI_VIDEO_MODEL_OPERATION_ANY_RE.match(normalized_path)
if model_operation_match:
operation_name = (
f"models/{str(model_operation_match.group('model') or '').strip()}/operations/"
f"{str(model_operation_match.group('operation_id') or '').strip()}"
)
params = {"task_id": operation_name}
if model_operation_match.group("cancel"):
params["action"] = "cancel"
return params
if normalized_path == "/v1beta/operations":
return {}
if normalized_path.startswith("/v1beta/operations/"):
operation_id = normalized_path[len("/v1beta/operations/") :].strip()
if operation_id.endswith(":cancel"):
return {"task_id": operation_id[: -len(":cancel")], "action": "cancel"}
return {"task_id": operation_id}
return {}
def _is_stream_request_payload(
body_json: dict[str, Any],
path_params: dict[str, Any] | None = None,
) -> bool:
if bool((path_params or {}).get("stream")):
return True
return bool(body_json.get("stream"))
def _load_gateway_auth_models(
db: Session,
auth_context: GatewayAuthContext,
) -> tuple[User, ApiKey]:
user = db.query(User).filter(User.id == auth_context.user_id).first()
api_key = db.query(ApiKey).filter(ApiKey.id == auth_context.api_key_id).first()
if not user or not api_key:
raise HTTPException(status_code=401, detail="无效的API密钥")
if not user.is_active or user.is_deleted:
raise HTTPException(status_code=401, detail="无效的API密钥")
if not api_key.is_active:
raise HTTPException(status_code=401, detail="无效的API密钥")
if api_key.is_locked and not api_key.is_standalone:
raise HTTPException(status_code=403, detail="该密钥已被管理员锁定,请联系管理员")
if str(api_key.user_id) != str(user.id):
raise HTTPException(status_code=401, detail="无效的API密钥")
return user, api_key
def _build_gateway_request_context(
*,
request: Request,
payload: GatewayExecuteRequest,
db: Session,
user: User,
api_key: ApiKey,
adapter: Any,
path_params: dict[str, Any],
balance_remaining: float | None,
) -> ApiRequestContext:
request_id = str(
payload.trace_id or getattr(request.state, "request_id", "") or uuid.uuid4().hex[:8]
)
request.state.request_id = request_id
request.state.user_id = user.id
request.state.api_key_id = api_key.id
request.state.prefetched_balance_remaining = balance_remaining
original_headers = {str(k): str(v) for k, v in (payload.headers or {}).items()}
client_accept_encoding = str(original_headers.get("accept-encoding") or "").strip() or None
raw_body = gateway_module._extract_gateway_raw_body(payload)
return ApiRequestContext(
request=request,
db=db,
user=user,
api_key=api_key,
request_id=request_id,
start_time=time.time(),
request_method=str(payload.method or request.method or "GET").upper(),
request_path=str(payload.path or request.url.path or "/"),
client_ip=request.client.host if request.client else "127.0.0.1",
user_agent=str(original_headers.get("user-agent") or "unknown"),
original_headers=original_headers,
query_params=gateway_module._parse_query_string(payload.query_string),
request_content_type=str(original_headers.get("content-type") or "").strip() or None,
raw_body=raw_body,
json_body=(dict(payload.body_json) if payload.body_json else None),
balance_remaining=balance_remaining,
mode=getattr(getattr(adapter, "mode", None), "value", "standard"),
api_format_hint=(
adapter.allowed_api_formats[0]
if getattr(adapter, "allowed_api_formats", None)
else None
),
path_params=dict(path_params or {}),
client_content_encoding=normalize_content_encoding(
original_headers.get("content-encoding")
),
client_accept_encoding=client_accept_encoding,
)
def _build_gateway_forward_request(
*,
request: Request,
payload: GatewayExecuteRequest,
) -> Request:
body = gateway_module._decode_gateway_body(payload)
header_items = [
(str(key).encode("latin-1"), str(value).encode("latin-1"))
for key, value in (payload.headers or {}).items()
]
scope = {
"type": "http",
"http_version": "1.1",
"method": str(payload.method or "GET").upper(),
"scheme": "http",
"path": str(payload.path or "/"),
"raw_path": str(payload.path or "/").encode("utf-8"),
"query_string": str(payload.query_string or "").encode("utf-8"),
"headers": header_items,
"client": ("127.0.0.1", 0),
"server": ("127.0.0.1", 80),
"app": request.app,
"state": {"request_id": str(payload.trace_id or uuid.uuid4().hex[:8])},
}
received = False
async def receive() -> dict[str, object]:
nonlocal received
if received:
return {"type": "http.request", "body": b"", "more_body": False}
received = True
return {"type": "http.request", "body": body, "more_body": False}
return Request(scope, receive)
def _build_proxy_public_fallback_response() -> JSONResponse:
return JSONResponse(
status_code=409,
content={"action": CONTROL_ACTION_PROXY_PUBLIC},
headers={CONTROL_ACTION_HEADER: CONTROL_ACTION_PROXY_PUBLIC},
)
def _stream_executor_requires_python_rewrite(
*,
envelope: Any = None,
needs_conversion: bool,
provider_api_format: str | None = None,
client_api_format: str | None = None,
) -> bool:
provider_api_format = str(provider_api_format or "").strip().lower()
client_api_format = str(client_api_format or "").strip().lower()
if needs_conversion:
if (
envelope is None
and (
(
provider_api_format in {"claude:chat", "gemini:chat"}
and client_api_format == "openai:chat"
)
or (
provider_api_format in {"claude:cli", "gemini:cli"}
and client_api_format in {"openai:cli", "openai:compact"}
)
)
) or (
str(getattr(envelope, "name", "") or "").strip().lower() == "antigravity:v1internal"
and (
(provider_api_format == "gemini:chat" and client_api_format == "openai:chat")
or (
provider_api_format == "gemini:cli"
and client_api_format in {"openai:cli", "openai:compact"}
)
)
):
return False
return True
if envelope is None:
return False
try:
requires_rewrite = bool(envelope.force_stream_rewrite())
except Exception:
return True
if not requires_rewrite:
return False
envelope_name = str(getattr(envelope, "name", "") or "").strip().lower()
if (
envelope_name == "antigravity:v1internal"
and provider_api_format == client_api_format
and provider_api_format in {"gemini:chat", "gemini:cli"}
):
return False
if (
envelope_name == "kiro:generateassistantresponse"
and provider_api_format == "claude:cli"
and client_api_format == "claude:cli"
):
return False
return True
def _serialize_gateway_sync_proxy(proxy: Any) -> dict[str, Any] | None:
if proxy is None:
return None
raw = asdict(proxy)
return {key: value for key, value in raw.items() if value is not None}
def _serialize_gateway_sync_timeouts(timeouts: Any) -> dict[str, Any] | None:
if timeouts is None:
return None
raw = asdict(timeouts)
return {key: value for key, value in raw.items() if value is not None}
def _extract_gateway_upstream_auth(
provider_request_headers: dict[str, str],
*,
provider_api_format: str,
key: Any,
) -> tuple[str, str]:
normalized_headers = {
str(header_name).strip().lower(): str(header_value).strip()
for header_name, header_value in (provider_request_headers or {}).items()
if str(header_name).strip() and str(header_value).strip()
}
for header_name in ("authorization", "x-api-key", "x-goog-api-key"):
header_value = normalized_headers.get(header_name)
if header_value:
return header_name, header_value
auth_header, auth_type = get_auth_config_for_endpoint(provider_api_format)
decrypted_key = crypto_service.decrypt(key.api_key)
auth_value = f"Bearer {decrypted_key}" if auth_type == "bearer" else decrypted_key
return str(auth_header or "").strip() or "authorization", auth_value
File diff suppressed because it is too large Load Diff
@@ -1,40 +0,0 @@
"""User monitoring routers."""
from __future__ import annotations
from fastapi import APIRouter
from starlette.routing import BaseRoute
from .user import router as monitoring_router
_RUST_OWNED_MONITORING_ROUTE_SIGNATURES = frozenset(
{
("GET", "/api/monitoring/my-audit-logs"),
("GET", "/api/monitoring/rate-limit-status"),
}
)
def _route_is_rust_owned(route: BaseRoute) -> bool:
path = getattr(route, "path", None)
methods = getattr(route, "methods", None)
if not isinstance(path, str) or not methods:
return False
return any(
(method, path) in _RUST_OWNED_MONITORING_ROUTE_SIGNATURES
for method in methods
if method not in {"HEAD", "OPTIONS"}
)
def _build_python_monitoring_router() -> APIRouter:
router = APIRouter()
router.include_router(monitoring_router)
router.routes = [route for route in router.routes if not _route_is_rust_owned(route)]
return router
python_monitoring_router = _build_python_monitoring_router()
router = python_monitoring_router
__all__ = ["python_monitoring_router", "router"]
@@ -1,39 +0,0 @@
"""Payment API routes."""
from __future__ import annotations
from fastapi import APIRouter
from starlette.routing import BaseRoute
from .routes import router as payment_router
_RUST_OWNED_PAYMENT_ROUTE_SIGNATURES = frozenset(
{
("POST", "/api/payment/callback/{payment_method}"),
}
)
def _route_is_rust_owned(route: BaseRoute) -> bool:
path = getattr(route, "path", None)
methods = getattr(route, "methods", None)
if not isinstance(path, str) or not methods:
return False
return any(
(method, path) in _RUST_OWNED_PAYMENT_ROUTE_SIGNATURES
for method in methods
if method not in {"HEAD", "OPTIONS"}
)
def _build_python_payment_router() -> APIRouter:
router = APIRouter()
router.include_router(payment_router)
router.routes = [route for route in router.routes if not _route_is_rust_owned(route)]
return router
python_payment_router = _build_python_payment_router()
router = python_payment_router
__all__ = ["python_payment_router", "router"]
-24
View File
@@ -1,24 +0,0 @@
"""Public-facing API routers.
Keep the compatibility frontdoor surface explicit so Rust can take ownership of
that manifest later without re-auditing the entire Python public app shell.
"""
from __future__ import annotations
from fastapi import APIRouter
from .support import python_public_support_router
router = APIRouter()
router.include_router(python_public_support_router)
__all__ = ["frontdoor_compat_router", "python_public_support_router", "router"]
def __getattr__(name: str) -> object:
if name == "frontdoor_compat_router":
from .compat import frontdoor_compat_router
return frontdoor_compat_router
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
@@ -1,193 +0,0 @@
"""
能力配置公共 API
提供系统支持的能力列表,供前端展示和配置使用。
"""
from dataclasses import dataclass
from typing import Any
from fastapi import APIRouter, Depends, Request
from sqlalchemy.orm import Session
from src.api.base.adapter import ApiAdapter, ApiMode
from src.api.base.context import ApiRequestContext
from src.api.base.pipeline import get_pipeline
from src.core.key_capabilities import (
get_all_capabilities,
get_user_configurable_capabilities,
)
from src.database import get_db
router = APIRouter(prefix="/api/capabilities", tags=["System Catalog"])
pipeline = get_pipeline()
def _serialize_capability(cap: Any) -> dict[str, Any]:
return {
"name": cap.name,
"display_name": cap.display_name,
"short_name": cap.short_name,
"description": cap.description,
"match_mode": cap.match_mode.value,
"config_mode": cap.config_mode.value,
}
class PublicCapabilitiesApiAdapter(ApiAdapter):
mode = ApiMode.PUBLIC
def authorize(self, context: ApiRequestContext) -> None: # type: ignore[override]
return None
class PublicCapabilitiesListAdapter(PublicCapabilitiesApiAdapter):
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
del context
return {"capabilities": [_serialize_capability(cap) for cap in get_all_capabilities()]}
class PublicUserConfigurableCapabilitiesAdapter(PublicCapabilitiesApiAdapter):
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
del context
return {
"capabilities": [
_serialize_capability(cap) for cap in get_user_configurable_capabilities()
]
}
@dataclass
class PublicModelCapabilitiesAdapter(PublicCapabilitiesApiAdapter):
model_name: str
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
from src.models.database import GlobalModel
global_model = (
context.db.query(GlobalModel)
.filter(GlobalModel.name == self.model_name, GlobalModel.is_active == True)
.first()
)
if not global_model:
return {
"model": self.model_name,
"supported_capabilities": [],
"capability_details": [],
"error": "模型不存在",
}
supported_caps = global_model.supported_capabilities or []
all_caps = {cap.name: cap for cap in get_all_capabilities()}
capability_details = [
{
"name": cap.name,
"display_name": cap.display_name,
"description": cap.description,
"match_mode": cap.match_mode.value,
"config_mode": cap.config_mode.value,
}
for cap_name in supported_caps
if (cap := all_caps.get(cap_name)) is not None
]
return {
"model": self.model_name,
"global_model_id": str(global_model.id),
"global_model_name": global_model.name,
"supported_capabilities": supported_caps,
"capability_details": capability_details,
}
@router.get("")
async def list_capabilities(request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取所有能力定义
返回系统中定义的所有能力(capabilities),包括用户可配置和系统内部使用的能力。
能力用于描述模型支持的功能特性,如视觉输入、函数调用、流式输出等。
**返回字段**
- capabilities: 能力列表,每个能力包含:
- name: 能力的唯一标识符(如 vision、function_calling)
- display_name: 能力的显示名称(如"视觉输入"、"函数调用")
- short_name: 能力的简短名称(如"视觉"、"函数")
- description: 能力的详细描述
- match_mode: 匹配模式(exact 精确匹配,fuzzy 模糊匹配,prefix 前缀匹配等)
- config_mode: 配置模式(user_configurable 用户可配置,system_only 仅系统使用)
"""
adapter = PublicCapabilitiesListAdapter()
return await pipeline.run(
adapter=adapter,
http_request=request,
db=db,
mode=ApiMode.PUBLIC,
)
@router.get("/user-configurable")
async def list_user_configurable_capabilities(
request: Request,
db: Session = Depends(get_db),
) -> Any:
"""
获取用户可配置的能力列表
返回允许用户在 API Key 中配置的能力列表,用于前端展示配置选项。
用户可以通过配置这些能力来限制或指定 API Key 可以访问的模型功能。
**返回字段**
- capabilities: 用户可配置的能力列表,每个能力包含:
- name: 能力的唯一标识符
- display_name: 能力的显示名称
- short_name: 能力的简短名称
- description: 能力的详细描述
- match_mode: 匹配模式(exact、fuzzy、prefix 等)
- config_mode: 配置模式(此接口返回的都是 user_configurable)
"""
adapter = PublicUserConfigurableCapabilitiesAdapter()
return await pipeline.run(
adapter=adapter,
http_request=request,
db=db,
mode=ApiMode.PUBLIC,
)
@router.get("/model/{model_name}")
async def get_model_supported_capabilities(
model_name: str,
request: Request,
db: Session = Depends(get_db),
) -> Any:
"""
获取指定模型支持的能力列表
根据全局模型名称(GlobalModel.name)查询该模型支持的能力,
并返回每个能力的详细定义。只查询活跃的全局模型。
**路径参数**
- model_name: 全局模型名称(如 claude-sonnet-4-20250514,必须是 GlobalModel.name)
**返回字段**
- model: 查询的模型名称
- global_model_id: 全局模型的 UUID
- global_model_name: 全局模型的标准名称
- supported_capabilities: 该模型支持的能力名称列表
- capability_details: 支持的能力详细信息列表,每个能力包含:
- name: 能力标识符
- display_name: 能力显示名称
- description: 能力描述
- match_mode: 匹配模式
- config_mode: 配置模式
- error: 错误信息(仅在模型不存在时返回)
"""
adapter = PublicModelCapabilitiesAdapter(model_name=model_name)
return await pipeline.run(
adapter=adapter,
http_request=request,
db=db,
mode=ApiMode.PUBLIC,
)
-26
View File
@@ -1,26 +0,0 @@
"""Rust-frontdoor-owned public compatibility route definitions."""
from fastapi import APIRouter
from .claude import router as claude_router
from .gemini import router as gemini_router
from .gemini_files import router as gemini_files_router
from .openai import router as openai_router
from .videos import router as videos_router
def build_frontdoor_compat_router() -> APIRouter:
"""Return public compat routes that Rust frontdoor owns at host level."""
compat_router = APIRouter()
compat_router.include_router(videos_router, tags=["Video Generation"])
compat_router.include_router(claude_router, tags=["Claude API"])
compat_router.include_router(openai_router)
compat_router.include_router(gemini_router, tags=["Gemini API"])
compat_router.include_router(gemini_files_router, tags=["Gemini Files API"])
return compat_router
frontdoor_compat_router = build_frontdoor_compat_router()
__all__ = ["build_frontdoor_compat_router", "frontdoor_compat_router"]
-62
View File
@@ -1,62 +0,0 @@
"""公开模块状态 API(供登录页等使用)"""
from typing import Any
from fastapi import APIRouter, Depends, Request
from pydantic import BaseModel
from sqlalchemy.orm import Session
from src.api.base.adapter import ApiAdapter, ApiMode
from src.api.base.context import ApiRequestContext
from src.api.base.pipeline import get_pipeline
from src.core.modules import get_module_registry
from src.database import get_db
router = APIRouter(prefix="/api/modules", tags=["Modules"])
pipeline = get_pipeline()
class AuthModuleInfo(BaseModel):
"""认证模块简要信息"""
name: str
display_name: str
active: bool
class PublicModulesApiAdapter(ApiAdapter):
mode = ApiMode.PUBLIC
def authorize(self, context: ApiRequestContext) -> None: # type: ignore[override]
return None
class PublicAuthModulesStatusAdapter(PublicModulesApiAdapter):
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
registry = get_module_registry()
auth_modules = registry.get_auth_modules_status(context.db)
return [
AuthModuleInfo(
name=status.name,
display_name=status.display_name,
active=status.active,
)
for status in auth_modules
]
@router.get("/auth-status", response_model=list[AuthModuleInfo])
async def get_auth_modules_status(request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取认证模块状态(公开接口)
供登录页使用,返回所有可用的认证模块及其激活状态。
不需要认证即可访问。
**返回字段**:
- `name`: 模块名称
- `display_name`: 显示名称
- `active`: 是否激活
"""
adapter = PublicAuthModulesStatusAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=ApiMode.PUBLIC)
-18
View File
@@ -1,18 +0,0 @@
"""Python-hosted public support route definitions."""
from fastapi import APIRouter
from .catalog import python_host_router as catalog_python_host_router
def build_python_public_support_router() -> APIRouter:
"""Return public routes that still belong to the Python host."""
support_router = APIRouter()
support_router.include_router(catalog_python_host_router)
return support_router
python_public_support_router = build_python_public_support_router()
__all__ = ["build_python_public_support_router", "python_public_support_router"]
@@ -1,694 +0,0 @@
"""
System Catalog / 健康检查相关端点
这些是系统工具端点,不需要复杂的 Adapter 抽象。
"""
from __future__ import annotations
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Any
import httpx
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from sqlalchemy import func
from sqlalchemy.orm import Session, load_only, selectinload
from src.api.base.adapter import ApiAdapter, ApiMode
from src.api.base.context import ApiRequestContext
from src.api.base.pipeline import get_pipeline
from src.api.handlers.base.request_builder import (
PassthroughRequestBuilder,
build_test_request_body,
get_provider_auth,
)
from src.clients.redis_client import get_redis_client
from src.config.settings import config
from src.core.logger import logger
from src.database import get_db
from src.database.database import get_pool_status
from src.models.database import GlobalModel, Model, Provider, ProviderAPIKey, ProviderEndpoint
from src.services.provider.provider_context import resolve_provider_proxy
from src.services.provider.transport import build_provider_url
router = APIRouter(tags=["System Catalog"])
pipeline = get_pipeline()
class PublicSystemCatalogApiAdapter(ApiAdapter):
mode = ApiMode.PUBLIC
def authorize(self, context: ApiRequestContext) -> None: # type: ignore[override]
return None
# ============== 辅助函数 ==============
def _as_bool(value: str | None, default: bool) -> bool:
"""将字符串转换为布尔值"""
if value is None:
return default
return value.lower() in {"1", "true", "yes", "on"}
def _serialize_provider(
provider: Provider,
include_models: bool,
include_endpoints: bool,
) -> dict[str, Any]:
"""序列化 Provider 对象"""
provider_data: dict[str, Any] = {
"id": provider.id,
"name": provider.name,
"is_active": provider.is_active,
"provider_priority": provider.provider_priority,
}
if include_endpoints:
provider_data["endpoints"] = [
{
"id": endpoint.id,
"base_url": endpoint.base_url,
"api_format": endpoint.api_format if endpoint.api_format else None,
"is_active": endpoint.is_active,
}
for endpoint in provider.endpoints or []
]
if include_models:
provider_data["models"] = [
{
"id": model.id,
"name": (
model.global_model.name if model.global_model else model.provider_model_name
),
"display_name": (
model.global_model.display_name
if model.global_model
else model.provider_model_name
),
"is_active": model.is_active,
"supports_streaming": model.supports_streaming,
}
for model in provider.models or []
if model.is_active
]
return provider_data
def _select_provider(db: Session, provider_name: str | None) -> Provider | None:
"""选择 Provider(按 provider_priority 优先级选择)"""
query = db.query(Provider).filter(Provider.is_active.is_(True))
if provider_name:
provider = query.filter(Provider.name == provider_name).first()
if provider:
return provider
# 按优先级选择(provider_priority 最小的优先)
return query.order_by(Provider.provider_priority.asc()).first()
async def _build_test_connection_transport_context(
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
) -> tuple[dict[str, Any] | None, dict[str, Any] | None, Any]:
from src.services.proxy_node.resolver import (
build_proxy_url_async,
get_system_proxy_config_async,
resolve_delegate_config_async,
resolve_effective_proxy,
resolve_proxy_info_async,
)
from src.services.request.execution_runtime_plan import ExecutionProxySnapshot
try:
effective_proxy = resolve_effective_proxy(
resolve_provider_proxy(endpoint=endpoint, key=key),
getattr(key, "proxy", None),
)
if not effective_proxy or not effective_proxy.get("enabled", True):
effective_proxy = await get_system_proxy_config_async()
delegate_cfg = await resolve_delegate_config_async(effective_proxy)
proxy_url: str | None = None
if effective_proxy and not (delegate_cfg and delegate_cfg.get("tunnel")):
proxy_url = await build_proxy_url_async(effective_proxy)
proxy_info = await resolve_proxy_info_async(effective_proxy)
proxy_snapshot = ExecutionProxySnapshot.from_proxy_info(
proxy_info,
proxy_url=proxy_url,
mode_override="tunnel" if delegate_cfg and delegate_cfg.get("tunnel") else None,
node_id_override=(
str(delegate_cfg.get("node_id") or "").strip() or None
if delegate_cfg and delegate_cfg.get("tunnel")
else None
),
)
return effective_proxy, delegate_cfg, proxy_snapshot
except Exception as exc:
logger.warning(
"Failed to build test-connection transport context endpoint={} key={}: {}",
getattr(endpoint, "id", None),
getattr(key, "id", None),
exc,
)
return None, None, None
async def _try_rust_test_connection_response(
*,
request_id: str,
url: str,
headers: dict[str, str],
body: dict[str, Any],
provider_name: str,
provider_id: str | None,
endpoint_id: str | None,
key_id: str | None,
api_format: str,
model_name: str,
proxy_snapshot: Any,
) -> httpx.Response:
import json
from src.services.request.execution_runtime_plan import (
ExecutionPlan,
ExecutionPlanTimeouts,
build_execution_plan_body,
)
from src.services.request.execution_runtime_client import (
ExecutionRuntimeClient,
ExecutionRuntimeClientError,
)
if config.execution_runtime_backend != "rust":
raise HTTPException(
status_code=503,
detail="System catalog test-connection requires Rust executor",
)
try:
result = await ExecutionRuntimeClient().execute_sync_json(
ExecutionPlan(
request_id=request_id,
candidate_id=None,
provider_name=provider_name,
provider_id=str(provider_id or ""),
endpoint_id=str(endpoint_id or ""),
key_id=str(key_id or ""),
method="POST",
url=url,
headers=dict(headers),
body=build_execution_plan_body(body, content_type="application/json"),
stream=False,
provider_api_format=api_format,
client_api_format=api_format,
model_name=model_name,
content_type="application/json",
proxy=proxy_snapshot,
timeouts=ExecutionPlanTimeouts(
connect_ms=30_000,
read_ms=30_000,
write_ms=30_000,
pool_ms=30_000,
total_ms=30_000,
),
)
)
except (ExecutionRuntimeClientError, httpx.HTTPError, json.JSONDecodeError) as exc:
logger.warning("Rust test-connection unavailable url={}: {}", url, exc)
raise HTTPException(
status_code=503,
detail="System catalog test-connection requires Rust executor",
) from exc
response_headers = dict(result.headers)
if result.response_json is not None:
response_headers.setdefault("content-type", "application/json")
response_body = json.dumps(result.response_json, ensure_ascii=False).encode("utf-8")
elif result.response_body_bytes is not None:
response_body = result.response_body_bytes
else:
response_body = b""
return httpx.Response(
status_code=result.status_code,
request=httpx.Request("POST", url, headers=headers),
headers=response_headers,
content=response_body,
)
async def _service_health_response(db: Session) -> dict[str, Any]:
active_providers = (
db.query(func.count(Provider.id)).filter(Provider.is_active.is_(True)).scalar() or 0
)
active_models = db.query(func.count(Model.id)).filter(Model.is_active.is_(True)).scalar() or 0
redis_info: dict[str, Any] = {"status": "unknown"}
try:
redis = await get_redis_client()
if redis:
await redis.ping()
redis_info = {"status": "ok"}
else:
redis_info = {"status": "degraded", "message": "Redis client not initialized"}
except Exception as exc:
redis_info = {"status": "error", "message": str(exc)}
return {
"status": "ok",
"timestamp": datetime.now(timezone.utc).isoformat(),
"stats": {
"active_providers": active_providers,
"active_models": active_models,
},
"dependencies": {
"database": {"status": "ok"},
"redis": redis_info,
},
}
def _health_check_response() -> dict[str, Any]:
try:
pool_status = get_pool_status()
pool_health = {
"checked_out": pool_status["checked_out"],
"pool_size": pool_status["pool_size"],
"overflow": pool_status["overflow"],
"max_capacity": pool_status["max_capacity"],
"usage_rate": (
f"{(pool_status['checked_out'] / pool_status['max_capacity'] * 100):.1f}%"
if pool_status["max_capacity"] > 0
else "0.0%"
),
}
except Exception as e:
pool_health = {"error": str(e)}
return {
"status": "healthy",
"timestamp": datetime.now(timezone.utc).isoformat(),
"database_pool": pool_health,
}
def _root_response(db: Session) -> dict[str, Any]:
top_provider = (
db.query(Provider)
.options(load_only(Provider.id, Provider.name, Provider.provider_priority))
.filter(Provider.is_active.is_(True))
.order_by(Provider.provider_priority.asc())
.first()
)
active_providers = (
db.query(func.count(Provider.id)).filter(Provider.is_active.is_(True)).scalar() or 0
)
return {
"message": "AI Proxy with Modular Architecture v4.0.0",
"status": "running",
"current_provider": top_provider.name if top_provider else "None",
"available_providers": active_providers,
"config": {},
"endpoints": {
"messages": "/v1/messages",
"count_tokens": "/v1/messages/count_tokens",
"health": "/v1/health",
"providers": "/v1/providers",
"test_connection": "/v1/test-connection",
},
}
def _list_providers_response(
db: Session,
*,
include_models: bool,
include_endpoints: bool,
active_only: bool,
) -> dict[str, Any]:
load_options = [
load_only(Provider.id, Provider.name, Provider.is_active, Provider.provider_priority)
]
if include_models:
load_options.append(
selectinload(Provider.models)
.load_only(
Model.id,
Model.provider_model_name,
Model.is_active,
Model.supports_streaming,
Model.global_model_id,
)
.selectinload(Model.global_model)
.load_only(GlobalModel.id, GlobalModel.name, GlobalModel.display_name)
)
if include_endpoints:
load_options.append(
selectinload(Provider.endpoints).load_only(
ProviderEndpoint.id,
ProviderEndpoint.base_url,
ProviderEndpoint.api_format,
ProviderEndpoint.is_active,
)
)
base_query = db.query(Provider)
if load_options:
base_query = base_query.options(*load_options)
if active_only:
base_query = base_query.filter(Provider.is_active.is_(True))
base_query = base_query.order_by(Provider.provider_priority.asc(), Provider.name.asc())
providers = base_query.all()
return {
"providers": [
_serialize_provider(provider, include_models, include_endpoints)
for provider in providers
]
}
def _provider_detail_response(
db: Session,
*,
provider_identifier: str,
include_models: bool,
include_endpoints: bool,
) -> dict[str, Any]:
load_options = [
load_only(Provider.id, Provider.name, Provider.is_active, Provider.provider_priority)
]
if include_models:
load_options.append(
selectinload(Provider.models)
.load_only(
Model.id,
Model.provider_model_name,
Model.is_active,
Model.supports_streaming,
Model.global_model_id,
)
.selectinload(Model.global_model)
.load_only(GlobalModel.id, GlobalModel.name, GlobalModel.display_name)
)
if include_endpoints:
load_options.append(
selectinload(Provider.endpoints).load_only(
ProviderEndpoint.id,
ProviderEndpoint.base_url,
ProviderEndpoint.api_format,
ProviderEndpoint.is_active,
)
)
base_query = db.query(Provider)
if load_options:
base_query = base_query.options(*load_options)
provider = base_query.filter(
(Provider.id == provider_identifier) | (Provider.name == provider_identifier)
).first()
if not provider:
raise HTTPException(status_code=404, detail="Provider not found")
return _serialize_provider(provider, include_models, include_endpoints)
async def _test_connection_response(
*,
request: Request,
db: Session,
provider: str | None,
model: str,
api_format: str | None,
) -> dict[str, Any]:
selected_provider = _select_provider(db, provider)
if not selected_provider:
raise HTTPException(status_code=503, detail="No active provider available")
active_endpoints: list[ProviderEndpoint] = [
ep for ep in (selected_provider.endpoints or []) if getattr(ep, "is_active", False)
]
if not active_endpoints:
raise HTTPException(status_code=503, detail="Provider has no active endpoints")
if api_format:
endpoint = next(
(ep for ep in active_endpoints if (ep.api_format or "") == api_format),
None,
)
if not endpoint:
raise HTTPException(
status_code=400,
detail=f"Provider has no active endpoint for api_format={api_format}",
)
format_value = api_format
else:
endpoint = active_endpoints[0]
format_value = endpoint.api_format or "claude:chat"
active_keys: list[ProviderAPIKey] = [
k for k in (selected_provider.api_keys or []) if getattr(k, "is_active", False)
]
if not active_keys:
raise HTTPException(status_code=503, detail="Provider has no active api keys")
def _key_supports_format(k: ProviderAPIKey) -> bool:
formats = getattr(k, "api_formats", None)
if formats is None:
return True
if isinstance(formats, list):
return str(format_value) in {str(x) for x in formats}
return True
key = next((k for k in active_keys if _key_supports_format(k)), active_keys[0])
payload = build_test_request_body(
format_value,
request_data={
"model": model,
"messages": [{"role": "user", "content": "Health check"}],
"max_tokens": 5,
},
)
try:
auth_info = await get_provider_auth(endpoint, key)
request_builder = PassthroughRequestBuilder()
provider_payload, provider_headers = request_builder.build(
payload,
{},
endpoint,
key,
is_stream=False,
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
)
url = build_provider_url(
endpoint,
query_params=dict(request.query_params),
path_params={"model": model},
is_stream=False,
key=key,
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
)
proxy_config, delegate_cfg, proxy_snapshot = await _build_test_connection_transport_context(
endpoint,
key,
)
resp = await _try_rust_test_connection_response(
request_id=f"test-connection:{selected_provider.id}:{model}",
url=url,
headers=provider_headers,
body=provider_payload,
provider_name=selected_provider.name,
provider_id=getattr(selected_provider, "id", None),
endpoint_id=getattr(endpoint, "id", None),
key_id=getattr(key, "id", None),
api_format=format_value,
model_name=model,
proxy_snapshot=proxy_snapshot,
)
if resp is None:
raise HTTPException(
status_code=503,
detail="System catalog test-connection requires Rust executor",
)
resp.raise_for_status()
response = resp.json()
return {
"status": "success",
"provider": selected_provider.name,
"endpoint_id": getattr(endpoint, "id", None),
"api_format": format_value,
"timestamp": datetime.now(timezone.utc).isoformat(),
"response_id": response.get("id", "unknown"),
}
except HTTPException:
raise
except Exception as exc:
logger.error(f"API connectivity test failed: {exc}")
raise HTTPException(status_code=503, detail=str(exc))
class PublicServiceHealthAdapter(PublicSystemCatalogApiAdapter):
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
return await _service_health_response(context.db)
class PublicSimpleHealthCheckAdapter(PublicSystemCatalogApiAdapter):
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
del context
return _health_check_response()
class PublicRootCatalogAdapter(PublicSystemCatalogApiAdapter):
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
return _root_response(context.db)
@dataclass
class PublicProvidersListAdapter(PublicSystemCatalogApiAdapter):
include_models: bool
include_endpoints: bool
active_only: bool
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
return _list_providers_response(
context.db,
include_models=self.include_models,
include_endpoints=self.include_endpoints,
active_only=self.active_only,
)
@dataclass
class PublicProviderDetailAdapter(PublicSystemCatalogApiAdapter):
provider_identifier: str
include_models: bool
include_endpoints: bool
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
return _provider_detail_response(
context.db,
provider_identifier=self.provider_identifier,
include_models=self.include_models,
include_endpoints=self.include_endpoints,
)
@dataclass
class PublicTestConnectionAdapter(PublicSystemCatalogApiAdapter):
provider: str | None
model: str
api_format: str | None
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
return await _test_connection_response(
request=context.request,
db=context.db,
provider=self.provider,
model=self.model,
api_format=self.api_format,
)
# ============== 端点 ==============
@router.get("/v1/health")
async def service_health(request: Request, db: Session = Depends(get_db)) -> Any:
"""返回服务健康状态与依赖信息"""
adapter = PublicServiceHealthAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=ApiMode.PUBLIC)
@router.get("/health")
async def health_check(request: Request, db: Session = Depends(get_db)) -> Any:
"""简单健康检查端点(无需认证)"""
adapter = PublicSimpleHealthCheckAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=ApiMode.PUBLIC)
@router.get("/")
async def root(request: Request, db: Session = Depends(get_db)) -> Any:
"""Root endpoint - 服务信息概览"""
adapter = PublicRootCatalogAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=ApiMode.PUBLIC)
@router.get("/v1/providers")
async def list_providers(
request: Request,
db: Session = Depends(get_db),
include_models: bool = Query(False),
include_endpoints: bool = Query(False),
active_only: bool = Query(True),
) -> Any:
"""列出所有 Provider"""
adapter = PublicProvidersListAdapter(
include_models=include_models,
include_endpoints=include_endpoints,
active_only=active_only,
)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=ApiMode.PUBLIC)
@router.get("/v1/providers/{provider_identifier}")
async def provider_detail(
provider_identifier: str,
request: Request,
db: Session = Depends(get_db),
include_models: bool = Query(False),
include_endpoints: bool = Query(False),
) -> Any:
"""获取单个 Provider 详情"""
adapter = PublicProviderDetailAdapter(
provider_identifier=provider_identifier,
include_models=include_models,
include_endpoints=include_endpoints,
)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=ApiMode.PUBLIC)
@router.get("/v1/test-connection")
async def test_connection(
request: Request,
db: Session = Depends(get_db),
provider: str | None = Query(None),
model: str = Query("claude-3-haiku-20240307"),
api_format: str | None = Query(None),
) -> Any:
"""测试 Provider 连接"""
adapter = PublicTestConnectionAdapter(
provider=provider,
model=model,
api_format=api_format,
)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=ApiMode.PUBLIC)
@router.get("/test-connection")
async def test_connection_legacy(
request: Request,
db: Session = Depends(get_db),
provider: str | None = Query(None),
model: str = Query("claude-3-haiku-20240307"),
api_format: str | None = Query(None),
) -> Any:
"""测试 Provider 连接(legacy alias,已弃用)"""
del request, db, provider, model, api_format
raise HTTPException(
status_code=410,
detail="Deprecated endpoint. Please use /v1/test-connection.",
)
@@ -1,67 +0,0 @@
"""Routes for authenticated user self-service APIs."""
from __future__ import annotations
from fastapi import APIRouter
from starlette.routing import BaseRoute
from .routes import router as me_router
_RUST_OWNED_USER_ME_ROUTE_SIGNATURES = frozenset(
{
("GET", "/api/users/me"),
("PUT", "/api/users/me"),
("PATCH", "/api/users/me/password"),
("GET", "/api/users/me/sessions"),
("DELETE", "/api/users/me/sessions/others"),
("PATCH", "/api/users/me/sessions/{session_id}"),
("DELETE", "/api/users/me/sessions/{session_id}"),
("GET", "/api/users/me/api-keys"),
("POST", "/api/users/me/api-keys"),
("GET", "/api/users/me/api-keys/{key_id}"),
("DELETE", "/api/users/me/api-keys/{key_id}"),
("PUT", "/api/users/me/api-keys/{key_id}"),
("PATCH", "/api/users/me/api-keys/{key_id}"),
("GET", "/api/users/me/usage"),
("GET", "/api/users/me/usage/active"),
("GET", "/api/users/me/usage/interval-timeline"),
("GET", "/api/users/me/usage/heatmap"),
("GET", "/api/users/me/providers"),
("GET", "/api/users/me/available-models"),
("GET", "/api/users/me/endpoint-status"),
("PUT", "/api/users/me/api-keys/{api_key_id}/providers"),
("PUT", "/api/users/me/api-keys/{api_key_id}/capabilities"),
("GET", "/api/users/me/preferences"),
("PUT", "/api/users/me/preferences"),
("GET", "/api/users/me/model-capabilities"),
("PUT", "/api/users/me/model-capabilities"),
}
)
def _route_is_rust_owned(route: BaseRoute) -> bool:
path = getattr(route, "path", None)
methods = getattr(route, "methods", None)
if not isinstance(path, str) or not methods:
return False
return any(
(method, path) in _RUST_OWNED_USER_ME_ROUTE_SIGNATURES
for method in methods
if method not in {"HEAD", "OPTIONS"}
)
def _build_python_user_me_router() -> APIRouter:
router = APIRouter()
router.include_router(me_router)
router.routes = [route for route in router.routes if not _route_is_rust_owned(route)]
return router
python_user_me_router = _build_python_user_me_router()
router = python_user_me_router
# 注意:management_tokens_router 已迁移到模块系统,由 ModuleRegistry 动态注册
# 当 MANAGEMENT_TOKENS_AVAILABLE=true 时注册
__all__ = ["python_user_me_router", "router"]
-48
View File
@@ -1,48 +0,0 @@
"""Wallet API routes."""
from __future__ import annotations
from fastapi import APIRouter
from starlette.routing import BaseRoute
from .routes import router as wallet_router
_RUST_OWNED_WALLET_ROUTE_SIGNATURES = frozenset(
{
("GET", "/api/wallet/balance"),
("GET", "/api/wallet/transactions"),
("GET", "/api/wallet/flow"),
("GET", "/api/wallet/today-cost"),
("GET", "/api/wallet/recharge"),
("POST", "/api/wallet/recharge"),
("GET", "/api/wallet/recharge/{order_id}"),
("GET", "/api/wallet/refunds"),
("POST", "/api/wallet/refunds"),
("GET", "/api/wallet/refunds/{refund_id}"),
}
)
def _route_is_rust_owned(route: BaseRoute) -> bool:
path = getattr(route, "path", None)
methods = getattr(route, "methods", None)
if not isinstance(path, str) or not methods:
return False
return any(
(method, path) in _RUST_OWNED_WALLET_ROUTE_SIGNATURES
for method in methods
if method not in {"HEAD", "OPTIONS"}
)
def _build_python_wallet_router() -> APIRouter:
router = APIRouter()
router.include_router(wallet_router)
router.routes = [route for route in router.routes if not _route_is_rust_owned(route)]
return router
python_wallet_router = _build_python_wallet_router()
router = python_wallet_router
__all__ = ["python_wallet_router", "router"]
@@ -1,515 +0,0 @@
"""
上游模型获取公共模块
提供从上游 API 获取模型列表的公共函数;通用 api_format 抓取能力统一来自 core.api_format.capabilities,供以下场景使用:
- 定时任务自动获取(fetch_scheduler.py)
- 管理后台手动查询(provider_query.py)
"""
import asyncio
import json
from collections.abc import Iterable
from dataclasses import dataclass
from typing import Any, Awaitable, Callable
import httpx
from src.config.settings import config
from src.core.api_format import get_extra_headers_from_endpoint
from src.core.api_format.capabilities import BROWSER_FINGERPRINT_HEADERS
from src.core.api_format.headers import build_adapter_headers_for_endpoint
from src.core.logger import logger
# 并发请求限制
MAX_CONCURRENT_REQUESTS = 5
# 模型获取格式优先级:同族内优先使用 chat 端点,若无则回退到 cli 端点
MODEL_FETCH_FORMAT_PRIORITY: list[tuple[str, ...]] = [
("openai:chat", "openai:cli", "openai:compact"),
("claude:chat", "claude:cli"),
("gemini:chat", "gemini:cli"),
]
# Return tuple signature:
# (models, errors, has_success, upstream_metadata)
_ModelsFetcher = Callable[
["UpstreamModelsFetchContext", float],
Awaitable[tuple[list[dict], list[str], bool, dict[str, Any] | None]],
]
@dataclass(frozen=True, slots=True)
class EndpointFetchConfig:
"""端点获取配置(纯数据,不依赖 DB session)。
从 ProviderEndpoint ORM 对象提取必要字段,确保在 DB session 关闭后
仍可安全使用(避免 DetachedInstanceError)。
"""
base_url: str
extra_headers: dict[str, str] | None = None
def build_format_to_config(endpoints: Iterable[Any]) -> dict[str, EndpointFetchConfig]:
"""将活跃的 ProviderEndpoint 转换为 api_format -> EndpointFetchConfig 映射。
应在 DB session 活跃时调用,提取 ORM 对象上的 base_url 和 header_rules,
转换为 session 无关的纯数据结构。
"""
result: dict[str, EndpointFetchConfig] = {}
for ep in endpoints:
if not getattr(ep, "is_active", False):
continue
result[ep.api_format] = EndpointFetchConfig(
base_url=ep.base_url,
extra_headers=get_extra_headers_from_endpoint(ep),
)
return result
@dataclass(frozen=True)
class UpstreamModelsFetchContext:
"""上游模型获取上下文(Key 级别)。"""
provider_type: str
api_key_value: str
format_to_endpoint: dict[str, EndpointFetchConfig]
proxy_config: dict[str, Any] | None = None
auth_config: dict[str, Any] | None = None
class UpstreamModelsFetcherRegistry:
"""按 provider_type 注册上游模型获取策略,避免到处写特判。"""
_fetchers: dict[str, _ModelsFetcher] = {}
@classmethod
def register(cls, *, provider_types: list[str], fetcher: _ModelsFetcher) -> None:
for pt in provider_types:
if not pt:
continue
cls._fetchers[pt.lower()] = fetcher
@classmethod
def get(cls, provider_type: str) -> _ModelsFetcher | None:
if not provider_type:
return None
return cls._fetchers.get(provider_type.lower())
async def _fetch_models_default(
ctx: UpstreamModelsFetchContext,
timeout_seconds: float,
) -> tuple[list[dict], list[str], bool, dict[str, Any] | None]:
endpoint_configs = build_all_format_configs(ctx.api_key_value, ctx.format_to_endpoint)
models, errors, has_success = await fetch_models_from_endpoints(
endpoint_configs, timeout=timeout_seconds, proxy_config=ctx.proxy_config
)
return models, errors, has_success, None
async def fetch_models_for_key(
ctx: UpstreamModelsFetchContext,
*,
timeout_seconds: float = 30.0,
) -> tuple[list[dict], list[str], bool, dict[str, Any] | None]:
"""统一入口:按 provider_type 选择策略获取模型列表(可附带 upstream_metadata)。"""
# Ensure provider plugins (including custom model fetchers) are registered.
from src.services.provider.envelope import ensure_providers_bootstrapped
ensure_providers_bootstrapped(provider_types=[ctx.provider_type] if ctx.provider_type else None)
fetcher = UpstreamModelsFetcherRegistry.get(ctx.provider_type) or _fetch_models_default
return await fetcher(ctx, timeout_seconds)
def merge_upstream_metadata(
current: dict[str, Any] | None,
incoming: dict[str, Any],
) -> dict[str, Any]:
"""合并上游元数据,对 quota_by_model 做模型级深度合并。
以上游返回为准(全量替换),不再保留上游已下架的模型。
仅对上游返回的模型补充旧数据中的 reset_time(当新数据缺少时)。
"""
merged: dict[str, Any] = dict(current) if isinstance(current, dict) else {}
for ns_key, ns_val in incoming.items():
old_ns = merged.get(ns_key)
if (
isinstance(ns_val, dict)
and isinstance(old_ns, dict)
and "quota_by_model" in ns_val
and "quota_by_model" in old_ns
):
old_qbm = old_ns["quota_by_model"]
new_qbm = ns_val["quota_by_model"]
if isinstance(old_qbm, dict) and isinstance(new_qbm, dict):
# 保留新数据中已有模型的旧 reset_time
for model_id, new_info in new_qbm.items():
if not isinstance(new_info, dict):
continue
old_info = old_qbm.get(model_id)
if (
isinstance(old_info, dict)
and "reset_time" in old_info
and "reset_time" not in new_info
):
new_info["reset_time"] = old_info["reset_time"]
merged[ns_key] = ns_val
return merged
def build_all_format_configs(
api_key_value: str,
format_to_endpoint: dict[str, EndpointFetchConfig],
) -> list[dict]:
"""
构建所有 API 格式的端点配置
只对实际配置了端点的格式构建请求配置,不同端点的 base_url 可能不同,
不应使用某个端点的 base_url 去尝试其他格式。
Args:
api_key_value: 解密后的 API Key
format_to_endpoint: API 格式到 EndpointFetchConfig 的映射
Returns:
端点配置列表,每个配置包含 api_key, base_url, api_format, extra_headers
"""
if not format_to_endpoint:
return []
# 同族内优先使用 chat 端点,若无则回退到 cli 端点
configs: list[dict] = []
for candidates in MODEL_FETCH_FORMAT_PRIORITY:
fmt = next((f for f in candidates if f in format_to_endpoint), None)
if fmt is not None:
cfg = format_to_endpoint[fmt]
base_url = str(getattr(cfg, "base_url", "") or "")
extra_headers: dict[str, str] | None
if isinstance(cfg, EndpointFetchConfig):
extra_headers = cfg.extra_headers
else:
# 允许直接传递类似 ProviderEndpoint 的对象(测试/独立使用场景)。
candidate_extra = getattr(cfg, "extra_headers", None)
if isinstance(candidate_extra, dict):
extra_headers = {str(k): str(v) for k, v in candidate_extra.items() if k}
else:
extra_headers = get_extra_headers_from_endpoint(cfg)
configs.append(
{
"api_key": api_key_value,
"base_url": base_url,
"api_format": fmt,
"extra_headers": extra_headers,
}
)
return configs
async def _build_models_proxy_snapshot(
proxy_config: dict[str, Any] | None,
) -> Any:
from src.services.request.execution_runtime_plan import build_proxy_snapshot
return await build_proxy_snapshot(proxy_config, label="upstream model fetch")
def _build_model_fetch_url(
api_format: str,
base_url: str,
api_key: str,
after_id: str | None = None,
) -> str:
normalized_api_format = str(api_format or "").strip().lower()
base_url = str(base_url or "").rstrip("/")
if normalized_api_format.startswith("gemini:"):
if base_url.endswith("/v1beta"):
return f"{base_url}/models?key={api_key}"
return f"{base_url}/v1beta/models?key={api_key}"
if base_url.endswith("/v1"):
url = f"{base_url}/models"
else:
url = f"{base_url}/v1/models"
if normalized_api_format.startswith("claude:"):
if after_id:
return f"{url}?limit=100&after_id={after_id}"
return f"{url}?limit=100"
return url
def _build_model_fetch_headers(
api_format: str,
api_key: str,
extra_headers: dict[str, str] | None,
) -> dict[str, str]:
normalized_api_format = str(api_format or "").strip().lower()
if normalized_api_format == "openai:cli":
headers = {"User-Agent": config.internal_user_agent_openai_cli}
if extra_headers:
headers.update(extra_headers)
return build_adapter_headers_for_endpoint("openai:chat", api_key, headers)
if normalized_api_format == "openai:compact":
headers = {"User-Agent": config.internal_user_agent_openai_cli}
if extra_headers:
headers.update(extra_headers)
return build_adapter_headers_for_endpoint("openai:chat", api_key, headers)
if normalized_api_format == "claude:cli":
headers = {"User-Agent": config.internal_user_agent_claude_cli}
if extra_headers:
headers.update(extra_headers)
return build_adapter_headers_for_endpoint(normalized_api_format, api_key, headers)
if normalized_api_format == "claude:chat":
headers = build_adapter_headers_for_endpoint(normalized_api_format, api_key, extra_headers)
if "authorization" not in {str(k).lower() for k in headers}:
headers["Authorization"] = f"Bearer {api_key}"
return headers
if normalized_api_format == "gemini:cli":
headers = {
**BROWSER_FINGERPRINT_HEADERS,
"User-Agent": config.internal_user_agent_gemini_cli,
}
if extra_headers:
headers.update(extra_headers)
return headers
if normalized_api_format == "gemini:chat":
headers = {**BROWSER_FINGERPRINT_HEADERS}
if extra_headers:
headers.update(extra_headers)
return headers
return build_adapter_headers_for_endpoint(normalized_api_format, api_key, extra_headers)
def _parse_model_fetch_response(
api_format: str,
payload: Any,
) -> tuple[list[dict[str, Any]], bool, str | None]:
normalized_api_format = str(api_format or "").strip().lower()
if normalized_api_format.startswith("gemini:"):
if isinstance(payload, dict) and isinstance(payload.get("models"), list):
out: list[dict[str, Any]] = []
for model in payload["models"]:
if not isinstance(model, dict):
continue
out.append(
{
"id": str(model.get("name", "")).replace("models/", ""),
"owned_by": "google",
"display_name": model.get("displayName", ""),
"api_format": normalized_api_format,
}
)
return out, False, None
return [], False, None
page_models: list[dict[str, Any]] = []
has_more = False
if isinstance(payload, dict) and isinstance(payload.get("data"), list):
page_models = [m for m in payload["data"] if isinstance(m, dict)]
has_more = bool(payload.get("has_more"))
elif isinstance(payload, list):
page_models = [m for m in payload if isinstance(m, dict)]
for model in page_models:
model.setdefault("api_format", normalized_api_format)
next_cursor = None
if normalized_api_format.startswith("claude:") and isinstance(payload, dict):
raw_last_id = payload.get("last_id")
if has_more and isinstance(raw_last_id, str) and raw_last_id.strip():
next_cursor = raw_last_id.strip()
return page_models, has_more, next_cursor
def _extract_rust_error_message(result: Any) -> str:
payload = getattr(result, "response_json", None)
if isinstance(payload, dict):
err = payload.get("error")
if isinstance(err, dict):
message = str(err.get("message") or "").strip()
if message:
return message[:500]
message = str(payload.get("message") or "").strip()
if message:
return message[:500]
try:
return json.dumps(payload, ensure_ascii=False)[:500]
except Exception:
pass
body_bytes = getattr(result, "response_body_bytes", None)
if body_bytes:
try:
return body_bytes.decode("utf-8", errors="replace")[:500]
except Exception:
return "(binary body)"
return "(empty)"
async def _try_rust_fetch_models_for_api_format(
*,
api_format: str,
base_url: str,
api_key: str,
extra_headers: dict[str, str] | None,
timeout: float,
proxy_config: dict[str, Any] | None,
) -> tuple[list[dict[str, Any]], str | None, bool] | None:
from src.services.request.execution_runtime_plan import (
ExecutionPlan,
ExecutionPlanBody,
ExecutionPlanTimeouts,
)
from src.services.request.execution_runtime_client import (
ExecutionRuntimeClient,
ExecutionRuntimeClientError,
)
if config.execution_runtime_backend != "rust":
return None
try:
proxy_snapshot = await _build_models_proxy_snapshot(proxy_config)
headers = _build_model_fetch_headers(api_format, api_key, extra_headers)
models: list[dict[str, Any]] = []
next_cursor: str | None = None
seen_ids: set[str] = set()
for _ in range(20):
url = _build_model_fetch_url(api_format, base_url, api_key, next_cursor)
result = await ExecutionRuntimeClient().execute_sync_json(
ExecutionPlan(
request_id=f"model-fetch:{api_format}:{next_cursor or 'root'}",
candidate_id=None,
provider_name=str(api_format.split(":", 1)[0] or "unknown"),
provider_id="",
endpoint_id="",
key_id="",
method="GET",
url=url,
headers=dict(headers),
body=ExecutionPlanBody(),
stream=False,
provider_api_format=api_format,
client_api_format=api_format,
model_name="models",
proxy=proxy_snapshot,
timeouts=ExecutionPlanTimeouts(
connect_ms=min(int(timeout * 1000), 30_000),
read_ms=int(timeout * 1000),
write_ms=int(timeout * 1000),
pool_ms=min(int(timeout * 1000), 30_000),
total_ms=int(timeout * 1000),
),
)
)
if result.status_code != 200:
error_body = _extract_rust_error_message(result)
return [], f"HTTP {result.status_code}: {error_body}", False
payload = result.response_json
page_models, has_more, next_after_id = _parse_model_fetch_response(api_format, payload)
for model in page_models:
model_id = model.get("id")
if isinstance(model_id, str) and model_id and model_id in seen_ids:
continue
if isinstance(model_id, str) and model_id:
seen_ids.add(model_id)
models.append(model)
if not has_more or not next_after_id or next_after_id == next_cursor:
return models, None, True
next_cursor = next_after_id
return models, None, True
except (ExecutionRuntimeClientError, httpx.HTTPError, json.JSONDecodeError) as exc:
logger.warning("Rust model fetch fallback for {}: {}", api_format, exc)
return None
except Exception as exc:
logger.warning("Rust model fetch unexpected fallback for {}: {}", api_format, exc)
return None
async def fetch_models_from_endpoints(
endpoint_configs: list[dict],
timeout: float = 30.0,
proxy_config: dict[str, Any] | None = None,
) -> tuple[list[dict], list[str], bool]:
"""
从多个端点并发获取模型
Args:
endpoint_configs: 端点配置列表,每个配置包含 api_key, base_url, api_format, extra_headers
timeout: 请求超时时间(秒)
proxy_config: 代理配置(可选),支持系统默认回退
Returns:
(模型列表, 错误列表, 是否有成功)
"""
all_models: list[dict] = []
errors: list[str] = []
has_success = False
semaphore = asyncio.Semaphore(MAX_CONCURRENT_REQUESTS)
async def fetch_one(config: dict) -> tuple[list, str | None, bool]:
base_url = config["base_url"]
if not base_url:
return [], None, False
base_url = base_url.rstrip("/")
api_format = config["api_format"]
api_key_value = config["api_key"]
extra_headers = config.get("extra_headers")
try:
async with semaphore:
rust_result = await _try_rust_fetch_models_for_api_format(
api_format=api_format,
base_url=base_url,
api_key=api_key_value,
extra_headers=extra_headers,
timeout=timeout,
proxy_config=proxy_config,
)
if rust_result is None:
return [], f"{api_format}: rust executor unavailable", False
models, error, success = rust_result
for m in models:
if "api_format" not in m:
m["api_format"] = api_format
# 即使返回空列表,只要没有错误也算成功
return models, error, success
except httpx.TimeoutException:
logger.warning("获取 {} 模型超时", api_format)
return [], f"{api_format}: timeout", False
except Exception:
logger.exception("获取 {} 模型出错", api_format)
return [], f"{api_format}: error", False
results = await asyncio.gather(*[fetch_one(c) for c in endpoint_configs])
for models, error, success in results:
all_models.extend(models)
if error:
errors.append(error)
if success:
has_success = True
return all_models, errors, has_success
@@ -1,108 +0,0 @@
"""Shared Rust executor HTTP helper for Antigravity side calls."""
from __future__ import annotations
import json
from typing import Any
import httpx
from src.config.settings import config
from src.core.logger import logger
async def execute_antigravity_rust_http_request(
*,
method: str,
url: str,
headers: dict[str, str],
body: Any,
proxy_config: dict[str, Any] | None,
request_id: str,
provider_api_format: str,
timeout_seconds: float,
content_type: str | None = None,
) -> httpx.Response | None:
from src.services.request.execution_runtime_plan import (
ExecutionPlan,
ExecutionPlanBody,
ExecutionPlanTimeouts,
build_execution_plan_body,
build_proxy_snapshot,
)
from src.services.request.execution_runtime_client import (
ExecutionRuntimeClient,
ExecutionRuntimeClientError,
)
if config.execution_runtime_backend != "rust":
return None
final_headers = dict(headers)
if (
body is not None
and content_type
and not any(str(key).lower() == "content-type" for key in final_headers)
):
final_headers["content-type"] = content_type
timeout_ms = max(int(timeout_seconds * 1000), 1_000)
try:
proxy_snapshot = await build_proxy_snapshot(proxy_config, label="Antigravity")
result = await ExecutionRuntimeClient().execute_sync_json(
ExecutionPlan(
request_id=request_id,
candidate_id=None,
provider_name="antigravity",
provider_id="",
endpoint_id="",
key_id="",
method=method,
url=url,
headers=final_headers,
body=(
build_execution_plan_body(body, content_type=content_type)
if body is not None
else ExecutionPlanBody()
),
stream=False,
provider_api_format=provider_api_format,
client_api_format=provider_api_format,
model_name="antigravity",
content_type=content_type,
proxy=proxy_snapshot,
timeouts=ExecutionPlanTimeouts(
connect_ms=timeout_ms,
read_ms=timeout_ms,
write_ms=timeout_ms,
pool_ms=timeout_ms,
total_ms=timeout_ms,
),
)
)
except (ExecutionRuntimeClientError, httpx.HTTPError, json.JSONDecodeError) as exc:
logger.warning("Antigravity Rust HTTP fallback {} {}: {}", method, url, exc)
return None
except Exception as exc:
logger.warning("Antigravity Rust HTTP unexpected fallback {} {}: {}", method, url, exc)
return None
response_headers = dict(result.headers)
if result.response_json is not None:
response_headers.setdefault("content-type", "application/json")
response_body = json.dumps(result.response_json, ensure_ascii=False).encode("utf-8")
elif result.response_body_bytes is not None:
response_body = result.response_body_bytes
else:
response_body = b""
return httpx.Response(
status_code=result.status_code,
request=httpx.Request(method, url, headers=final_headers),
headers=response_headers,
content=response_body,
)
__all__ = ["execute_antigravity_rust_http_request"]
@@ -1,108 +0,0 @@
"""Shared Rust executor HTTP helper for Gemini CLI side calls."""
from __future__ import annotations
import json
from typing import Any
import httpx
from src.config.settings import config
from src.core.logger import logger
async def execute_gemini_cli_rust_http_request(
*,
method: str,
url: str,
headers: dict[str, str],
body: Any,
proxy_config: dict[str, Any] | None,
request_id: str,
provider_api_format: str,
timeout_seconds: float,
content_type: str | None = None,
) -> httpx.Response | None:
from src.services.request.execution_runtime_plan import (
ExecutionPlan,
ExecutionPlanBody,
ExecutionPlanTimeouts,
build_execution_plan_body,
build_proxy_snapshot,
)
from src.services.request.execution_runtime_client import (
ExecutionRuntimeClient,
ExecutionRuntimeClientError,
)
if config.execution_runtime_backend != "rust":
return None
final_headers = dict(headers)
if (
body is not None
and content_type
and not any(str(key).lower() == "content-type" for key in final_headers)
):
final_headers["content-type"] = content_type
timeout_ms = max(int(timeout_seconds * 1000), 1_000)
try:
proxy_snapshot = await build_proxy_snapshot(proxy_config, label="GeminiCLI")
result = await ExecutionRuntimeClient().execute_sync_json(
ExecutionPlan(
request_id=request_id,
candidate_id=None,
provider_name="gemini_cli",
provider_id="",
endpoint_id="",
key_id="",
method=method,
url=url,
headers=final_headers,
body=(
build_execution_plan_body(body, content_type=content_type)
if body is not None
else ExecutionPlanBody()
),
stream=False,
provider_api_format=provider_api_format,
client_api_format=provider_api_format,
model_name="gemini_cli",
content_type=content_type,
proxy=proxy_snapshot,
timeouts=ExecutionPlanTimeouts(
connect_ms=timeout_ms,
read_ms=timeout_ms,
write_ms=timeout_ms,
pool_ms=timeout_ms,
total_ms=timeout_ms,
),
)
)
except (ExecutionRuntimeClientError, httpx.HTTPError, json.JSONDecodeError) as exc:
logger.warning("GeminiCLI Rust HTTP fallback {} {}: {}", method, url, exc)
return None
except Exception as exc:
logger.warning("GeminiCLI Rust HTTP unexpected fallback {} {}: {}", method, url, exc)
return None
response_headers = dict(result.headers)
if result.response_json is not None:
response_headers.setdefault("content-type", "application/json")
response_body = json.dumps(result.response_json, ensure_ascii=False).encode("utf-8")
elif result.response_body_bytes is not None:
response_body = result.response_body_bytes
else:
response_body = b""
return httpx.Response(
status_code=result.status_code,
request=httpx.Request(method, url, headers=final_headers),
headers=response_headers,
content=response_body,
)
__all__ = ["execute_gemini_cli_rust_http_request"]
@@ -1,108 +0,0 @@
"""Shared Rust executor HTTP helper for Kiro provider side calls."""
from __future__ import annotations
import json
from typing import Any
import httpx
from src.config.settings import config
from src.core.logger import logger
async def execute_kiro_rust_http_request(
*,
method: str,
url: str,
headers: dict[str, str],
body: Any,
proxy_config: dict[str, Any] | None,
request_id: str,
provider_api_format: str,
content_type: str | None = None,
timeout_seconds: float = 30.0,
) -> httpx.Response | None:
from src.services.request.execution_runtime_plan import (
ExecutionPlan,
ExecutionPlanBody,
ExecutionPlanTimeouts,
build_execution_plan_body,
build_proxy_snapshot,
)
from src.services.request.execution_runtime_client import (
ExecutionRuntimeClient,
ExecutionRuntimeClientError,
)
if config.execution_runtime_backend != "rust":
return None
final_headers = dict(headers)
if (
body is not None
and content_type
and not any(str(key).lower() == "content-type" for key in final_headers)
):
final_headers["content-type"] = content_type
timeout_ms = max(int(timeout_seconds * 1000), 1_000)
try:
proxy_snapshot = await build_proxy_snapshot(proxy_config, label="Kiro")
result = await ExecutionRuntimeClient().execute_sync_json(
ExecutionPlan(
request_id=request_id,
candidate_id=None,
provider_name="kiro",
provider_id="",
endpoint_id="",
key_id="",
method=method,
url=url,
headers=final_headers,
body=(
build_execution_plan_body(body, content_type=content_type)
if body is not None
else ExecutionPlanBody()
),
stream=False,
provider_api_format=provider_api_format,
client_api_format=provider_api_format,
model_name="kiro",
content_type=content_type,
proxy=proxy_snapshot,
timeouts=ExecutionPlanTimeouts(
connect_ms=timeout_ms,
read_ms=timeout_ms,
write_ms=timeout_ms,
pool_ms=timeout_ms,
total_ms=timeout_ms,
),
)
)
except (ExecutionRuntimeClientError, httpx.HTTPError, json.JSONDecodeError) as exc:
logger.warning("Kiro Rust HTTP fallback {} {}: {}", method, url, exc)
return None
except Exception as exc:
logger.warning("Kiro Rust HTTP unexpected fallback {} {}: {}", method, url, exc)
return None
response_headers = dict(result.headers)
if result.response_json is not None:
response_headers.setdefault("content-type", "application/json")
response_body = json.dumps(result.response_json, ensure_ascii=False).encode("utf-8")
elif result.response_body_bytes is not None:
response_body = result.response_body_bytes
else:
response_body = b""
return httpx.Response(
status_code=result.status_code,
request=httpx.Request(method, url, headers=final_headers),
headers=response_headers,
content=response_body,
)
__all__ = ["execute_kiro_rust_http_request"]
@@ -1,67 +0,0 @@
"""
Tunnel relay 配置。
控制 worker 是否通过本机 gateway relay 转发 tunnel 帧。
"""
from __future__ import annotations
import os
from dataclasses import dataclass
from urllib.parse import quote
_DEFAULT_DOCKER_TUNNEL_URL = "http://127.0.0.1:8084"
_DOCKER_TUNNEL_CONNECT_TIMEOUT_SECONDS = 5.0
_DOCKER_TUNNEL_URL_ENV_KEYS = ("AETHER_TUNNEL_BASE_URL",)
@dataclass(frozen=True)
class TunnelRelayConfig:
enabled: bool
url: str
connect_timeout_seconds: float
@property
def local_relay_base_url(self) -> str:
return f"{self.url.rstrip('/')}/api/internal/tunnel/relay"
def local_relay_url(self, node_id: str) -> str:
return f"{self.local_relay_base_url}/{quote(node_id, safe='')}"
_tunnel_relay_config: TunnelRelayConfig | None = None
def _is_docker_runtime() -> bool:
if os.getenv("DOCKER_CONTAINER", "").strip().lower() == "true":
return True
return os.path.exists("/.dockerenv")
def _resolve_docker_tunnel_url() -> str:
for key in _DOCKER_TUNNEL_URL_ENV_KEYS:
value = os.getenv(key, "").strip()
if value:
return value.rstrip("/")
return _DEFAULT_DOCKER_TUNNEL_URL
def get_tunnel_relay_config() -> TunnelRelayConfig:
"""读取 tunnel relay 配置(进程内缓存)。"""
global _tunnel_relay_config
if _tunnel_relay_config is not None:
return _tunnel_relay_config
docker_runtime = _is_docker_runtime()
_tunnel_relay_config = TunnelRelayConfig(
enabled=docker_runtime,
url=_resolve_docker_tunnel_url(),
connect_timeout_seconds=_DOCKER_TUNNEL_CONNECT_TIMEOUT_SECONDS,
)
return _tunnel_relay_config
def reset_tunnel_relay_config_cache() -> None:
"""测试或热更新场景下清理配置缓存。"""
global _tunnel_relay_config
_tunnel_relay_config = None
@@ -1,277 +0,0 @@
"""
Rust execution runtime 客户端主入口。
旧的 `rust_executor_client.py` 仍然保留,作为兼容入口。
"""
from __future__ import annotations
import base64
import json
from collections.abc import AsyncIterator
from dataclasses import dataclass, field
from typing import Any
import httpx
from src.config.settings import config
from src.services.request.execution_runtime_plan import ExecutionPlan
class ExecutionRuntimeClientError(RuntimeError):
"""Rust execution runtime 客户端错误。"""
@dataclass(slots=True)
class ExecutionRuntimeSyncResult:
status_code: int
response_json: Any = None
headers: dict[str, str] = field(default_factory=dict)
provider_response_json: Any = None
response_body_bytes: bytes | None = None
@dataclass(slots=True)
class ExecutionRuntimeStreamResult:
status_code: int
headers: dict[str, str]
byte_iterator: AsyncIterator[bytes]
response_ctx: Any
class _ExecutionRuntimeManagedStreamContext:
def __init__(self, client: httpx.AsyncClient, response_ctx: Any) -> None:
self._client = client
self._response_ctx = response_ctx
self._closed = False
async def __aexit__(self, exc_type: object, exc: object, tb: object) -> None:
if self._closed:
return
self._closed = True
try:
await self._response_ctx.__aexit__(exc_type, exc, tb)
finally:
await self._client.aclose()
class ExecutionRuntimeClient:
"""Python 控制面访问 Rust execution runtime 的轻量客户端。"""
def __init__(
self,
*,
transport: str | None = None,
base_url: str | None = None,
socket_path: str | None = None,
request_timeout: float | None = None,
) -> None:
self.transport = (transport or config.execution_runtime_transport).strip().lower()
self.base_url = (base_url or config.execution_runtime_base_url).strip()
self.socket_path = (socket_path or config.execution_runtime_socket_path).strip()
self.request_timeout = (
request_timeout
if request_timeout is not None
else config.execution_runtime_request_timeout
)
def _build_client(self, *, streaming: bool = False) -> httpx.AsyncClient:
if streaming:
timeout = httpx.Timeout(
connect=min(self.request_timeout, 30.0),
read=None,
write=self.request_timeout,
pool=self.request_timeout,
)
else:
timeout = httpx.Timeout(self.request_timeout)
if self.transport == "unix_socket":
if not self.socket_path:
raise ExecutionRuntimeClientError(
"EXECUTION_RUNTIME_SOCKET_PATH is required for unix_socket"
)
transport = httpx.AsyncHTTPTransport(uds=self.socket_path, retries=0)
return httpx.AsyncClient(
transport=transport,
base_url=self.base_url,
timeout=timeout,
)
if self.transport != "tcp":
raise ExecutionRuntimeClientError(
f"Unsupported execution-runtime transport: {self.transport}"
)
return httpx.AsyncClient(
base_url=self.base_url,
timeout=timeout,
transport=httpx.AsyncHTTPTransport(retries=0),
)
async def execute_sync_json(self, plan: ExecutionPlan) -> ExecutionRuntimeSyncResult:
async with self._build_client() as client:
response = await client.post(
"/v1/execute/sync",
json=plan.to_payload(),
)
response.raise_for_status()
payload = response.json()
status_code = int(payload.get("status_code") or 200)
headers = payload.get("headers") or {}
if not isinstance(headers, dict):
raise ExecutionRuntimeClientError("Execution runtime response headers must be an object")
response_json = payload.get("response_json")
provider_response_json = payload.get("provider_response_json")
body_payload = payload.get("body")
body_bytes_b64 = None
if response_json is None and isinstance(body_payload, dict):
response_json = body_payload.get("json_body")
body_bytes_b64 = body_payload.get("body_bytes_b64")
response_body_bytes: bytes | None = None
if body_bytes_b64 is not None:
if not isinstance(body_bytes_b64, str):
raise ExecutionRuntimeClientError(
"Execution runtime body_bytes_b64 must be a string"
)
try:
response_body_bytes = base64.b64decode(body_bytes_b64)
except Exception as exc: # noqa: BLE001
raise ExecutionRuntimeClientError(
"Execution runtime body_bytes_b64 must be valid base64"
) from exc
return ExecutionRuntimeSyncResult(
status_code=status_code,
response_json=response_json,
headers={str(k): str(v) for k, v in headers.items()},
provider_response_json=provider_response_json,
response_body_bytes=response_body_bytes,
)
async def execute_stream(self, plan: ExecutionPlan) -> ExecutionRuntimeStreamResult:
client = self._build_client(streaming=True)
response_ctx = client.stream(
"POST",
"/v1/execute/stream",
json=plan.to_payload(),
)
try:
response = await response_ctx.__aenter__()
response.raise_for_status()
line_iter = response.aiter_lines()
headers_frame = await self._read_first_stream_frame(line_iter)
payload = headers_frame.get("payload")
if not isinstance(payload, dict) or payload.get("kind") != "headers":
raise ExecutionRuntimeClientError(
"Execution runtime stream must start with headers frame"
)
status_code = int(payload.get("status_code") or 200)
headers = payload.get("headers") or {}
if not isinstance(headers, dict):
raise ExecutionRuntimeClientError(
"Execution runtime stream headers must be an object"
)
async def _byte_iter() -> AsyncIterator[bytes]:
async for line in line_iter:
if not line:
continue
frame = self._decode_stream_frame(line)
frame_payload = frame["payload"]
kind = str(frame_payload.get("kind") or "").strip().lower()
if kind == "data":
chunk_b64 = frame_payload.get("chunk_b64")
if isinstance(chunk_b64, str):
if chunk_b64:
try:
yield base64.b64decode(chunk_b64)
except Exception as exc: # noqa: BLE001
raise ExecutionRuntimeClientError(
"Execution runtime stream chunk_b64 must be valid base64"
) from exc
continue
text = frame_payload.get("text")
if isinstance(text, str):
if text:
yield text.encode("utf-8")
continue
if kind == "error":
error = frame_payload.get("error") or {}
message = str(error.get("message") or "execution runtime stream error")
raise httpx.ReadError(message)
if kind == "telemetry":
continue
if kind == "eof":
break
raise ExecutionRuntimeClientError(
f"Unexpected execution runtime stream frame kind: {kind}"
)
return ExecutionRuntimeStreamResult(
status_code=status_code,
headers={str(k): str(v) for k, v in headers.items()},
byte_iterator=_byte_iter(),
response_ctx=_ExecutionRuntimeManagedStreamContext(client, response_ctx),
)
except Exception:
try:
await response_ctx.__aexit__(None, None, None)
finally:
await client.aclose()
raise
@staticmethod
def _decode_stream_frame(line: str) -> dict[str, Any]:
try:
frame = json.loads(line)
except json.JSONDecodeError as exc:
raise ExecutionRuntimeClientError(
"Execution runtime stream frame must be valid JSON"
) from exc
if not isinstance(frame, dict):
raise ExecutionRuntimeClientError("Execution runtime stream frame must be an object")
payload = frame.get("payload")
if not isinstance(payload, dict):
raise ExecutionRuntimeClientError(
"Execution runtime stream frame payload must be an object"
)
return frame
async def _read_first_stream_frame(
self,
line_iter: AsyncIterator[str],
) -> dict[str, Any]:
async for line in line_iter:
if not line:
continue
return self._decode_stream_frame(line)
raise ExecutionRuntimeClientError(
"Execution runtime stream ended before headers frame"
)
# Compatibility aliases for older call sites still using executor terminology.
RustExecutorClientError = ExecutionRuntimeClientError
RustExecutorSyncResult = ExecutionRuntimeSyncResult
RustExecutorStreamResult = ExecutionRuntimeStreamResult
RustExecutorClient = ExecutionRuntimeClient
__all__ = [
"ExecutionRuntimeClient",
"ExecutionRuntimeClientError",
"ExecutionRuntimeStreamResult",
"ExecutionRuntimeSyncResult",
"RustExecutorClient",
"RustExecutorClientError",
"RustExecutorStreamResult",
"RustExecutorSyncResult",
]
@@ -1,274 +0,0 @@
"""
Execution runtime 计划契约
用于在 Python 控制面和 Rust execution runtime 之间传递稳定的请求执行信息。
当前阶段先服务于非流式 chat 路径的计划构建与本地执行拆分。
"""
from __future__ import annotations
import base64
import json
from dataclasses import asdict, dataclass
from typing import Any
def _drop_none(value: Any) -> Any:
"""递归移除 None 字段,便于序列化为紧凑 payload。"""
if isinstance(value, dict):
return {key: _drop_none(item) for key, item in value.items() if item is not None}
if isinstance(value, list):
return [_drop_none(item) for item in value]
return value
@dataclass(slots=True)
class ExecutionPlanTimeouts:
connect_ms: int | None = None
read_ms: int | None = None
write_ms: int | None = None
pool_ms: int | None = None
total_ms: int | None = None
@dataclass(slots=True)
class ExecutionPlanBody:
json_body: Any = None
body_bytes_b64: str | None = None
body_ref: str | None = None
@dataclass(slots=True)
class ExecutionProxySnapshot:
enabled: bool
mode: str | None = None
node_id: str | None = None
label: str | None = None
url: str | None = None
extra: dict[str, Any] | None = None
@classmethod
def from_proxy_info(
cls,
proxy_info: dict[str, Any] | None,
*,
proxy_url: str | None = None,
mode_override: str | None = None,
node_id_override: str | None = None,
extra: dict[str, Any] | None = None,
) -> ExecutionProxySnapshot | None:
if not proxy_info and not proxy_url:
return None
mode = (
mode_override
or str(proxy_info.get("type") or proxy_info.get("mode") or "").strip()
or None
)
if not mode and proxy_url:
mode = proxy_url.split("://", 1)[0].strip().lower() or None
return cls(
enabled=True,
mode=mode,
node_id=node_id_override
or str((proxy_info or {}).get("node_id") or "").strip()
or None,
label=str((proxy_info or {}).get("label") or "").strip() or None,
url=str(proxy_url or (proxy_info or {}).get("url") or "").strip() or None,
extra=extra or None,
)
@dataclass(slots=True)
class ExecutionPlan:
request_id: str
candidate_id: str | None
provider_name: str
provider_id: str
endpoint_id: str
key_id: str
method: str
url: str
headers: dict[str, str]
body: ExecutionPlanBody
stream: bool
provider_api_format: str
client_api_format: str
model_name: str
content_type: str | None = None
content_encoding: str | None = None
proxy: ExecutionProxySnapshot | None = None
tls_profile: str | None = None
timeouts: ExecutionPlanTimeouts | None = None
def to_payload(self) -> dict[str, Any]:
return _drop_none(asdict(self))
_REMOTE_EXECUTION_RUNTIME_BYPASS_FORMATS = {"openai:cli", "openai:compact"}
def should_bypass_remote_execution_runtime_url(
url: str | None,
*,
provider_api_format: str | None = None,
client_api_format: str | None = None,
) -> bool:
normalized_url = str(url or "").strip().lower()
if not normalized_url:
return False
normalized_provider_api_format = str(provider_api_format or "").strip().lower()
normalized_client_api_format = str(client_api_format or "").strip().lower()
if (
normalized_provider_api_format not in _REMOTE_EXECUTION_RUNTIME_BYPASS_FORMATS
and normalized_client_api_format not in _REMOTE_EXECUTION_RUNTIME_BYPASS_FORMATS
):
return False
return "/backend-api/codex" in normalized_url or "/backendapi/codex" in normalized_url
def should_bypass_remote_execution_runtime(contract: ExecutionPlan) -> bool:
return should_bypass_remote_execution_runtime_url(
contract.url,
provider_api_format=contract.provider_api_format,
client_api_format=contract.client_api_format,
)
def is_remote_execution_runtime_proxy_supported(proxy: ExecutionProxySnapshot | None) -> bool:
proxy_mode = str((proxy.mode if proxy else "") or "").strip().lower()
return (
proxy is None
or (str(proxy.url or "").strip() != "" and proxy_mode not in {"tunnel"})
or (proxy_mode == "tunnel" and str(proxy.node_id or "").strip() != "")
)
def is_remote_execution_runtime_contract_eligible(contract: ExecutionPlan) -> bool:
content_encoding = str(contract.content_encoding or "").strip().lower()
has_json_body = contract.body.json_body is not None
has_raw_body = bool(str(contract.body.body_bytes_b64 or "").strip())
has_body = has_json_body or has_raw_body
return (
((not has_body and content_encoding == "") or has_body)
and (not has_json_body or content_encoding in {"", "gzip"} or has_raw_body)
and is_remote_execution_runtime_proxy_supported(contract.proxy)
and not should_bypass_remote_execution_runtime(contract)
)
def build_execution_plan_body(
payload: Any,
*,
content_type: str | None = None,
) -> ExecutionPlanBody:
normalized_content_type = str(content_type or "").strip().lower()
if isinstance(payload, dict):
return ExecutionPlanBody(json_body=payload)
if isinstance(payload, list) and "json" in normalized_content_type:
return ExecutionPlanBody(json_body=payload)
if isinstance(payload, (bytes, bytearray, memoryview)):
return ExecutionPlanBody(body_bytes_b64=base64.b64encode(bytes(payload)).decode("ascii"))
if isinstance(payload, str):
return ExecutionPlanBody(
body_bytes_b64=base64.b64encode(payload.encode("utf-8")).decode("ascii")
)
if payload is None:
return ExecutionPlanBody()
return ExecutionPlanBody(
body_bytes_b64=base64.b64encode(
json.dumps(payload, ensure_ascii=False, separators=(",", ":")).encode("utf-8")
).decode("ascii")
)
@dataclass(slots=True)
class PreparedExecutionPlan:
"""本地执行所需的运行时上下文;`contract` 是可序列化的稳定边界。"""
contract: ExecutionPlan
payload: dict[str, Any]
headers: dict[str, str]
upstream_is_stream: bool
needs_conversion: bool
provider_type: str
request_timeout: float
delegate_config: dict[str, Any] | None = None
proxy_config: dict[str, Any] | None = None
envelope: Any = None
selected_base_url: str | None = None
client_content_encoding: str | None = None
proxy_info: dict[str, Any] | None = None
@property
def remote_eligible(self) -> bool:
return is_remote_execution_runtime_contract_eligible(self.contract)
async def build_proxy_snapshot(
proxy_config: dict[str, Any] | None,
*,
label: str = "adapter",
) -> ExecutionProxySnapshot | None:
if not proxy_config:
return None
try:
from src.services.proxy_node.resolver import (
build_proxy_url_async,
resolve_delegate_config_async,
resolve_proxy_info_async,
)
delegate_cfg = await resolve_delegate_config_async(proxy_config)
proxy_url: str | None = None
if proxy_config and not (delegate_cfg and delegate_cfg.get("tunnel")):
proxy_url = await build_proxy_url_async(proxy_config)
proxy_info = await resolve_proxy_info_async(proxy_config)
return ExecutionProxySnapshot.from_proxy_info(
proxy_info,
proxy_url=proxy_url,
mode_override="tunnel" if delegate_cfg and delegate_cfg.get("tunnel") else None,
node_id_override=(
str(delegate_cfg.get("node_id") or "").strip() or None
if delegate_cfg and delegate_cfg.get("tunnel")
else None
),
)
except Exception as exc:
from src.core.logger import logger
logger.warning("Build {} proxy snapshot failed: {}", label, exc)
return None
# Compatibility aliases for older call sites still using executor terminology.
should_bypass_remote_executor_url = should_bypass_remote_execution_runtime_url
should_bypass_remote_executor = should_bypass_remote_execution_runtime
is_remote_proxy_supported = is_remote_execution_runtime_proxy_supported
is_remote_contract_eligible = is_remote_execution_runtime_contract_eligible
__all__ = [
"ExecutionPlan",
"ExecutionPlanBody",
"ExecutionPlanTimeouts",
"ExecutionProxySnapshot",
"PreparedExecutionPlan",
"build_execution_plan_body",
"build_proxy_snapshot",
"is_remote_contract_eligible",
"is_remote_execution_runtime_contract_eligible",
"is_remote_execution_runtime_proxy_supported",
"is_remote_proxy_supported",
"should_bypass_remote_execution_runtime",
"should_bypass_remote_execution_runtime_url",
"should_bypass_remote_executor",
"should_bypass_remote_executor_url",
]
@@ -1,41 +0,0 @@
"""
旧 executor 命名兼容入口。
当前主实现已经迁到 `execution_runtime_plan.py`。
"""
from src.services.request.execution_runtime_plan import (
ExecutionPlan,
ExecutionPlanBody,
ExecutionPlanTimeouts,
ExecutionProxySnapshot,
PreparedExecutionPlan,
build_execution_plan_body,
build_proxy_snapshot,
is_remote_contract_eligible,
is_remote_execution_runtime_contract_eligible,
is_remote_execution_runtime_proxy_supported,
is_remote_proxy_supported,
should_bypass_remote_execution_runtime,
should_bypass_remote_execution_runtime_url,
should_bypass_remote_executor,
should_bypass_remote_executor_url,
)
__all__ = [
"ExecutionPlan",
"ExecutionPlanBody",
"ExecutionPlanTimeouts",
"ExecutionProxySnapshot",
"PreparedExecutionPlan",
"build_execution_plan_body",
"build_proxy_snapshot",
"is_remote_contract_eligible",
"is_remote_execution_runtime_contract_eligible",
"is_remote_execution_runtime_proxy_supported",
"is_remote_proxy_supported",
"should_bypass_remote_execution_runtime",
"should_bypass_remote_execution_runtime_url",
"should_bypass_remote_executor",
"should_bypass_remote_executor_url",
]
@@ -1,27 +0,0 @@
"""
旧 executor 命名兼容入口。
当前主实现已经迁到 `execution_runtime_client.py`。
"""
from src.services.request.execution_runtime_client import (
ExecutionRuntimeClient,
ExecutionRuntimeClientError,
ExecutionRuntimeStreamResult,
ExecutionRuntimeSyncResult,
RustExecutorClient,
RustExecutorClientError,
RustExecutorStreamResult,
RustExecutorSyncResult,
)
__all__ = [
"ExecutionRuntimeClient",
"ExecutionRuntimeClientError",
"ExecutionRuntimeStreamResult",
"ExecutionRuntimeSyncResult",
"RustExecutorClient",
"RustExecutorClientError",
"RustExecutorStreamResult",
"RustExecutorSyncResult",
]
+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()
))
}
}
}
@@ -1,23 +1,21 @@
use std::collections::HashMap;
use std::sync::atomic::{AtomicBool, AtomicU32, AtomicU64, AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use std::time::Duration;
use aether_runtime::{BoundedQueueSender, MetricKind, MetricSample, QueueSendError, QueueSnapshot};
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 super::control_plane::ControlPlaneClient;
use super::protocol;
use crate::control_plane::ControlPlaneClient;
use crate::protocol;
const MAX_REQUEST_BODY_FRAME_SIZE: usize = 32 * 1024;
const SOFT_AVOID_QUEUE_PRESSURE_PERCENT: u64 = 50;
const SOFT_AVOID_STREAM_PRESSURE_PERCENT: u64 = 85;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SendStatus {
@@ -34,13 +32,13 @@ pub struct ConnConfig {
}
pub struct BoundedOutbound {
tx: BoundedQueueSender<Message>,
tx: mpsc::Sender<Message>,
close_tx: watch::Sender<bool>,
closing: AtomicBool,
}
impl BoundedOutbound {
pub fn new(tx: BoundedQueueSender<Message>, close_tx: watch::Sender<bool>) -> Self {
pub fn new(tx: mpsc::Sender<Message>, close_tx: watch::Sender<bool>) -> Self {
Self {
tx,
close_tx,
@@ -55,11 +53,11 @@ impl BoundedOutbound {
match self.tx.try_send(msg) {
Ok(()) => SendStatus::Queued,
Err(QueueSendError::Closed(_)) => {
Err(TrySendError::Closed(_)) => {
self.mark_closing();
SendStatus::Closed
}
Err(QueueSendError::Full(_)) => {
Err(TrySendError::Full(_)) => {
self.mark_closing();
SendStatus::Congested
}
@@ -77,10 +75,6 @@ impl BoundedOutbound {
let _ = self.close_tx.send(true);
true
}
pub fn snapshot(&self) -> QueueSnapshot {
self.tx.snapshot()
}
}
pub struct ProxyConn {
@@ -91,8 +85,6 @@ pub struct ProxyConn {
next_stream_id: AtomicU32,
pub stream_count: AtomicUsize,
pub max_streams: usize,
draining: AtomicBool,
congested_total: AtomicU64,
}
impl ProxyConn {
@@ -100,7 +92,7 @@ impl ProxyConn {
id: u64,
node_id: String,
node_name: String,
tx: BoundedQueueSender<Message>,
tx: mpsc::Sender<Message>,
close_tx: watch::Sender<bool>,
max_streams: usize,
) -> Self {
@@ -112,8 +104,6 @@ impl ProxyConn {
next_stream_id: AtomicU32::new(2),
stream_count: AtomicUsize::new(0),
max_streams,
draining: AtomicBool::new(false),
congested_total: AtomicU64::new(0),
}
}
@@ -169,26 +159,17 @@ impl ProxyConn {
}
pub fn is_available(&self) -> bool {
!self.outbound.is_closing() && !self.is_draining()
!self.outbound.is_closing()
}
pub fn request_close(&self) {
self.outbound.mark_closing();
}
pub fn mark_draining(&self) -> bool {
!self.draining.swap(true, Ordering::AcqRel)
}
pub fn is_draining(&self) -> bool {
self.draining.load(Ordering::Acquire)
}
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 {
self.congested_total.fetch_add(1, Ordering::Relaxed);
warn!(
conn_id = self.id,
node_id = %self.node_id,
@@ -199,62 +180,6 @@ impl ProxyConn {
}
status
}
fn snapshot(&self) -> ProxyConnSnapshot {
let outbound = self.outbound.snapshot();
let stream_count = self.stream_count.load(Ordering::Relaxed);
let queue_pressure_percent = percent_u64(outbound.depth, outbound.capacity);
let stream_pressure_percent = percent_u64(stream_count, self.max_streams);
let soft_avoid = queue_pressure_percent >= SOFT_AVOID_QUEUE_PRESSURE_PERCENT
|| stream_pressure_percent >= SOFT_AVOID_STREAM_PRESSURE_PERCENT;
ProxyConnSnapshot {
conn_id: self.id,
available: self.is_available(),
closing: self.outbound.is_closing(),
draining: self.is_draining(),
stream_count,
max_streams: self.max_streams,
stream_pressure_percent,
outbound,
queue_pressure_percent,
soft_avoid,
congested_total: self.congested_total.load(Ordering::Relaxed),
}
}
}
#[derive(Debug, Clone, Copy)]
struct ProxyConnSnapshot {
conn_id: u64,
available: bool,
closing: bool,
draining: bool,
stream_count: usize,
max_streams: usize,
stream_pressure_percent: u64,
outbound: QueueSnapshot,
queue_pressure_percent: u64,
soft_avoid: bool,
congested_total: u64,
}
#[derive(Clone)]
struct ProxyConnCandidate {
conn: Arc<ProxyConn>,
snapshot: ProxyConnSnapshot,
}
impl ProxyConnCandidate {
fn rank_key(&self) -> (u8, u64, u64, usize, usize, u64) {
(
u8::from(self.snapshot.soft_avoid),
self.snapshot.queue_pressure_percent,
self.snapshot.stream_pressure_percent,
self.snapshot.outbound.depth,
self.snapshot.stream_count,
self.snapshot.conn_id,
)
}
}
#[derive(Debug, Clone)]
@@ -343,21 +268,13 @@ impl LocalStream {
}
}
async fn push_body_chunk(&self, payload: Bytes) -> bool {
fn push_body_chunk(&self, payload: Bytes) -> bool {
if self.terminal.load(Ordering::Acquire) {
return false;
}
// Use a timeout to prevent a slow consumer from blocking the shared
// proxy-connection reader (head-of-line blocking across streams).
match tokio::time::timeout(
Duration::from_secs(5),
self.body_tx.send(LocalBodyEvent::Chunk(payload)),
)
.await
{
Ok(Ok(())) => true,
_ => false,
}
self.body_tx
.try_send(LocalBodyEvent::Chunk(payload))
.is_ok()
}
fn finish(&self) {
@@ -407,49 +324,10 @@ pub struct HubRouter {
next_conn_id: AtomicU64,
next_local_stream_id: AtomicU64,
control_plane: ControlPlaneClient,
node_status_tx: mpsc::UnboundedSender<NodeStatusEvent>,
soft_avoid_selection_total: AtomicU64,
selection_retry_total: AtomicU64,
selection_unavailable_total: AtomicU64,
}
#[derive(Debug)]
struct NodeStatusEvent {
node_id: String,
connected: bool,
conn_count: usize,
observed_at_unix_secs: u64,
}
impl HubRouter {
pub fn new(control_plane: ControlPlaneClient) -> Arc<Self> {
let (node_status_tx, mut node_status_rx) = mpsc::unbounded_channel::<NodeStatusEvent>();
let worker_control_plane = control_plane.clone();
if let Ok(handle) = tokio::runtime::Handle::try_current() {
handle.spawn(async move {
while let Some(event) = node_status_rx.recv().await {
if let Err(error) = worker_control_plane
.push_node_status(
&event.node_id,
event.connected,
event.conn_count,
event.observed_at_unix_secs,
)
.await
{
warn!(
node_id = %event.node_id,
connected = event.connected,
conn_count = event.conn_count,
observed_at_unix_secs = event.observed_at_unix_secs,
error = %error,
"failed to push node status to app control plane"
);
}
}
});
}
Arc::new(Self {
proxy_conns: RwLock::new(HashMap::new()),
proxy_conns_by_id: DashMap::new(),
@@ -458,10 +336,6 @@ impl HubRouter {
next_conn_id: AtomicU64::new(1),
next_local_stream_id: AtomicU64::new(1),
control_plane,
node_status_tx,
soft_avoid_selection_total: AtomicU64::new(0),
selection_retry_total: AtomicU64::new(0),
selection_unavailable_total: AtomicU64::new(0),
})
}
@@ -517,59 +391,32 @@ impl HubRouter {
self.notify_node_status(node_id.to_string(), pool_size > 0, pool_size);
}
pub fn request_close_all_proxies(&self) -> usize {
let conns = self
.proxy_conns_by_id
.iter()
.map(|entry| Arc::clone(entry.value()))
.collect::<Vec<_>>();
let total = conns.len();
for conn in conns {
conn.request_close();
}
total
}
fn notify_node_status(&self, node_id: String, connected: bool, conn_count: usize) {
let event = NodeStatusEvent {
node_id,
connected,
conn_count,
observed_at_unix_secs: current_unix_secs(),
};
if let Err(error) = self.node_status_tx.send(event) {
warn!(
node_id = %error.0.node_id,
connected = error.0.connected,
conn_count = error.0.conn_count,
observed_at_unix_secs = error.0.observed_at_unix_secs,
"node status worker unavailable"
);
}
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 ranked_proxy_conn_candidates(&self, node_id: &str) -> Vec<ProxyConnCandidate> {
let conns = {
let map = self.proxy_conns.read();
map.get(node_id)
.map(|entries| entries.to_vec())
.unwrap_or_default()
};
let mut candidates = conns
.into_iter()
.filter_map(|conn| {
let snapshot = conn.snapshot();
snapshot
.available
.then_some(ProxyConnCandidate { conn, snapshot })
})
.collect::<Vec<_>>();
candidates.sort_by_key(|candidate| candidate.rank_key());
candidates
}
pub fn has_local_proxy(&self, node_id: &str) -> bool {
!self.ranked_proxy_conn_candidates(node_id).is_empty()
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(
@@ -577,49 +424,12 @@ impl HubRouter {
node_id: &str,
meta: &protocol::RequestMeta,
) -> Result<Arc<LocalStream>, String> {
let candidates = self.ranked_proxy_conn_candidates(node_id);
if candidates.is_empty() {
self.selection_unavailable_total
.fetch_add(1, Ordering::Relaxed);
self.warn_no_available_proxy_connection(node_id);
return Err(format!("no proxy connection for node {node_id}"));
}
let mut skipped_candidates = 0usize;
let mut selected_candidate = None;
let mut proxy_stream_id = None;
for candidate in candidates {
match candidate.conn.alloc_stream_id() {
Some(stream_id) => {
proxy_stream_id = Some(stream_id);
selected_candidate = Some(candidate);
break;
}
None => skipped_candidates = skipped_candidates.saturating_add(1),
}
}
let Some(candidate) = selected_candidate else {
self.selection_unavailable_total
.fetch_add(1, Ordering::Relaxed);
return Err(format!("stream limit reached for node {node_id}"));
};
if skipped_candidates > 0 {
self.selection_retry_total.fetch_add(1, Ordering::Relaxed);
}
if candidate.snapshot.soft_avoid {
self.soft_avoid_selection_total
.fetch_add(1, Ordering::Relaxed);
debug!(
node_id = %node_id,
conn_id = candidate.snapshot.conn_id,
queue_pressure_percent = candidate.snapshot.queue_pressure_percent,
stream_pressure_percent = candidate.snapshot.stream_pressure_percent,
"selected high-pressure proxy connection because no lower-pressure alternative was available"
);
}
let proxy_conn = candidate.conn;
let proxy_stream_id = proxy_stream_id.expect("selected candidate should carry a stream id");
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
@@ -657,18 +467,7 @@ impl HubRouter {
self.proxy_to_local
.insert((proxy_conn.id, proxy_stream_id), local_stream_id);
let send_status = proxy_conn.send(Message::Binary(header_frame.into()));
debug!(
node_id = %node_id,
conn_id = proxy_conn.id,
proxy_stream_id = proxy_stream_id,
local_stream_id = local_stream_id,
stream_count = proxy_conn.stream_count.load(Ordering::Relaxed),
queue_depth = proxy_conn.outbound.snapshot().depth,
send_status = ?send_status,
"open_local_stream dispatched"
);
match send_status {
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);
@@ -678,28 +477,6 @@ impl HubRouter {
}
}
fn warn_no_available_proxy_connection(&self, node_id: &str) {
let conns = {
let map = self.proxy_conns.read();
map.get(node_id)
.map(|entries| entries.to_vec())
.unwrap_or_default()
};
if conns.is_empty() {
return;
}
let snapshots = conns.iter().map(|conn| conn.snapshot()).collect::<Vec<_>>();
warn!(
node_id = %node_id,
total_conns = snapshots.len(),
available = snapshots.iter().filter(|snapshot| snapshot.available).count(),
closing = snapshots.iter().filter(|snapshot| snapshot.closing).count(),
draining = snapshots.iter().filter(|snapshot| snapshot.draining).count(),
soft_avoid = snapshots.iter().filter(|snapshot| snapshot.soft_avoid).count(),
"no available proxy connection despite registered connections"
);
}
pub fn push_local_request_body(
&self,
local_stream_id: u64,
@@ -803,7 +580,7 @@ impl HubRouter {
self.route_response_headers(proxy_conn_id, header, data);
}
protocol::RESPONSE_BODY => {
self.route_response_body(proxy_conn_id, header, data).await;
self.route_response_body(proxy_conn_id, header, data);
}
protocol::STREAM_END => {
self.finish_proxy_stream(proxy_conn_id, header.stream_id);
@@ -828,16 +605,10 @@ impl HubRouter {
}
protocol::PONG => {}
protocol::GOAWAY => {
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
let first = pc.mark_draining();
if first {
warn!(
proxy_conn_id = proxy_conn_id,
node_id = %pc.node_id,
"received GOAWAY from proxy connection; marking connection draining"
);
}
}
warn!(
proxy_conn_id = proxy_conn_id,
"received GOAWAY from proxy connection"
);
}
_ => {
debug!(
@@ -879,12 +650,7 @@ impl HubRouter {
}
}
async fn route_response_body(
&self,
proxy_conn_id: u64,
header: protocol::FrameHeader,
data: &[u8],
) {
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;
};
@@ -902,7 +668,7 @@ impl HubRouter {
None => return,
};
if !stream.push_body_chunk(Bytes::from(payload)).await {
if !stream.push_body_chunk(Bytes::from(payload)) {
self.cancel_local_stream(local_id, "local relay response congested");
}
}
@@ -967,12 +733,8 @@ impl HubRouter {
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; keeping heartbeat pending"
);
return;
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) {
@@ -1004,223 +766,29 @@ impl HubRouter {
}
pub fn stats(&self) -> HubStats {
let proxy_conns = self
.proxy_conns_by_id
.iter()
.map(|entry| entry.value().snapshot())
.collect::<Vec<_>>();
let total_proxy = proxy_conns.len();
let nodes = self.proxy_conns.read().len();
let available_proxy_connections = proxy_conns
.iter()
.filter(|snapshot| snapshot.available)
.count();
let closing_proxy_connections = proxy_conns
.iter()
.filter(|snapshot| snapshot.closing)
.count();
let draining_proxy_connections = proxy_conns
.iter()
.filter(|snapshot| snapshot.draining)
.count();
let soft_avoid_proxy_connections = proxy_conns
.iter()
.filter(|snapshot| snapshot.available && snapshot.soft_avoid)
.count();
let outbound_queue_depth_total = proxy_conns
.iter()
.map(|snapshot| snapshot.outbound.depth)
.sum();
let outbound_queue_depth_max = proxy_conns
.iter()
.map(|snapshot| snapshot.outbound.depth)
.max()
.unwrap_or(0);
let outbound_queue_capacity_total = proxy_conns
.iter()
.map(|snapshot| snapshot.outbound.capacity)
.sum();
let outbound_queue_rejected_full_total = proxy_conns
.iter()
.map(|snapshot| snapshot.outbound.rejected_full_total)
.sum();
let outbound_queue_rejected_closed_total = proxy_conns
.iter()
.map(|snapshot| snapshot.outbound.rejected_closed_total)
.sum();
let proxy_connection_congested_total = proxy_conns
.iter()
.map(|snapshot| snapshot.congested_total)
.sum();
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,
available_proxy_connections,
closing_proxy_connections,
draining_proxy_connections,
soft_avoid_proxy_connections,
nodes,
active_streams: self.local_streams.len(),
outbound_queue_depth_total,
outbound_queue_depth_max,
outbound_queue_capacity_total,
outbound_queue_rejected_full_total,
outbound_queue_rejected_closed_total,
proxy_connection_congested_total,
soft_avoid_selection_total: self.soft_avoid_selection_total.load(Ordering::Relaxed),
selection_retry_total: self.selection_retry_total.load(Ordering::Relaxed),
selection_unavailable_total: self.selection_unavailable_total.load(Ordering::Relaxed),
}
}
}
fn current_unix_secs() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs()
}
fn percent_u64(value: usize, total: usize) -> u64 {
if total == 0 {
return 0;
}
((value as u128) * 100 / (total as u128)) as u64
}
#[derive(serde::Serialize)]
pub struct HubStats {
pub proxy_connections: usize,
pub available_proxy_connections: usize,
pub closing_proxy_connections: usize,
pub draining_proxy_connections: usize,
pub soft_avoid_proxy_connections: usize,
pub nodes: usize,
pub active_streams: usize,
pub outbound_queue_depth_total: usize,
pub outbound_queue_depth_max: usize,
pub outbound_queue_capacity_total: usize,
pub outbound_queue_rejected_full_total: u64,
pub outbound_queue_rejected_closed_total: u64,
pub proxy_connection_congested_total: u64,
pub soft_avoid_selection_total: u64,
pub selection_retry_total: u64,
pub selection_unavailable_total: u64,
}
impl HubStats {
pub fn to_metric_samples(&self) -> Vec<MetricSample> {
vec![
MetricSample::new(
"tunnel_proxy_connections",
"Current number of connected proxy sockets.",
MetricKind::Gauge,
self.proxy_connections as u64,
),
MetricSample::new(
"tunnel_proxy_connections_available",
"Current number of proxy connections available for new work.",
MetricKind::Gauge,
self.available_proxy_connections as u64,
),
MetricSample::new(
"tunnel_proxy_connections_closing",
"Current number of proxy connections marked closing.",
MetricKind::Gauge,
self.closing_proxy_connections as u64,
),
MetricSample::new(
"tunnel_proxy_connections_draining",
"Current number of proxy connections marked draining.",
MetricKind::Gauge,
self.draining_proxy_connections as u64,
),
MetricSample::new(
"tunnel_proxy_connections_soft_avoid",
"Current number of available proxy connections currently soft-avoided by the scheduler.",
MetricKind::Gauge,
self.soft_avoid_proxy_connections as u64,
),
MetricSample::new(
"tunnel_nodes",
"Current number of connected logical nodes.",
MetricKind::Gauge,
self.nodes as u64,
),
MetricSample::new(
"tunnel_active_streams",
"Current number of active local relay streams.",
MetricKind::Gauge,
self.active_streams as u64,
),
MetricSample::new(
"tunnel_proxy_outbound_queue_depth_total",
"Current aggregate depth across proxy outbound queues.",
MetricKind::Gauge,
self.outbound_queue_depth_total as u64,
),
MetricSample::new(
"tunnel_proxy_outbound_queue_depth_max",
"Current maximum depth observed on a single proxy outbound queue.",
MetricKind::Gauge,
self.outbound_queue_depth_max as u64,
),
MetricSample::new(
"tunnel_proxy_outbound_queue_capacity_total",
"Current aggregate capacity across proxy outbound queues.",
MetricKind::Gauge,
self.outbound_queue_capacity_total as u64,
),
MetricSample::new(
"tunnel_proxy_outbound_queue_rejected_full_total",
"Total proxy outbound queue sends rejected because a queue was full.",
MetricKind::Counter,
self.outbound_queue_rejected_full_total,
),
MetricSample::new(
"tunnel_proxy_outbound_queue_rejected_closed_total",
"Total proxy outbound queue sends rejected because a queue was closed.",
MetricKind::Counter,
self.outbound_queue_rejected_closed_total,
),
MetricSample::new(
"tunnel_proxy_connection_congested_total",
"Total number of times a proxy outbound queue became congested.",
MetricKind::Counter,
self.proxy_connection_congested_total,
),
MetricSample::new(
"tunnel_proxy_soft_avoid_selection_total",
"Total number of times the scheduler had to pick a high-pressure proxy connection.",
MetricKind::Counter,
self.soft_avoid_selection_total,
),
MetricSample::new(
"tunnel_proxy_selection_retry_total",
"Total number of times the scheduler retried a lower-ranked proxy connection after a race on stream allocation.",
MetricKind::Counter,
self.selection_retry_total,
),
MetricSample::new(
"tunnel_proxy_selection_unavailable_total",
"Total number of relay selections that failed because no proxy connection was available.",
MetricKind::Counter,
self.selection_unavailable_total,
),
]
}
}
#[cfg(test)]
mod tests {
use aether_runtime::bounded_queue;
use super::{protocol, ControlPlaneClient, HubRouter, ProxyConn, MAX_REQUEST_BODY_FRAME_SIZE};
use axum::extract::ws::Message;
use bytes::Bytes;
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::watch;
use super::*;
fn build_meta() -> protocol::RequestMeta {
protocol::RequestMeta {
@@ -1228,8 +796,6 @@ mod tests {
url: "https://example.com".to_string(),
headers: HashMap::new(),
timeout: 30,
follow_redirects: None,
http1_only: false,
}
}
@@ -1237,7 +803,7 @@ mod tests {
async fn cancel_local_stream_notifies_proxy() {
let hub = HubRouter::new(ControlPlaneClient::disabled());
let (proxy_tx, mut proxy_rx) = bounded_queue(8);
let (proxy_tx, mut proxy_rx) = mpsc::channel(8);
let (proxy_close_tx, _) = watch::channel(false);
let proxy = Arc::new(ProxyConn::new(
100,
@@ -1272,7 +838,7 @@ mod tests {
async fn push_local_request_body_splits_large_payload_and_marks_end() {
let hub = HubRouter::new(ControlPlaneClient::disabled());
let (proxy_tx, mut proxy_rx) = bounded_queue(8);
let (proxy_tx, mut proxy_rx) = mpsc::channel(8);
let (proxy_close_tx, _) = watch::channel(false);
let proxy = Arc::new(ProxyConn::new(
200,
@@ -1309,173 +875,4 @@ mod tests {
assert_eq!(second_header.msg_type, protocol::REQUEST_BODY);
assert_ne!(second_header.flags & protocol::FLAG_END_STREAM, 0);
}
#[tokio::test]
async fn goaway_marks_connection_draining_and_reroutes_new_streams() {
let hub = HubRouter::new(ControlPlaneClient::disabled());
let (proxy_one_tx, mut proxy_one_rx) = bounded_queue(8);
let (proxy_one_close_tx, _) = watch::channel(false);
let proxy_one = Arc::new(ProxyConn::new(
201,
"node-drain".to_string(),
"Node Drain".to_string(),
proxy_one_tx,
proxy_one_close_tx,
16,
));
hub.register_proxy(Arc::clone(&proxy_one));
let (proxy_two_tx, mut proxy_two_rx) = bounded_queue(8);
let (proxy_two_close_tx, _) = watch::channel(false);
let proxy_two = Arc::new(ProxyConn::new(
202,
"node-drain".to_string(),
"Node Drain".to_string(),
proxy_two_tx,
proxy_two_close_tx,
16,
));
hub.register_proxy(Arc::clone(&proxy_two));
let mut goaway = protocol::encode_goaway();
hub.handle_proxy_frame(201, &mut goaway).await;
assert!(
proxy_one.is_draining(),
"first connection should be draining"
);
assert!(
!proxy_two.is_draining(),
"second connection should remain schedulable"
);
let _stream = hub
.open_local_stream("node-drain", &build_meta())
.expect("open local stream");
assert!(
proxy_one_rx.try_recv().is_err(),
"draining connection should not receive new streams"
);
let routed = proxy_two_rx
.try_recv()
.expect("headers should route to second connection");
let routed_data = match routed {
Message::Binary(data) => data.to_vec(),
other => panic!("unexpected message: {other:?}"),
};
let header = protocol::FrameHeader::parse(&routed_data).expect("frame header");
assert_eq!(header.msg_type, protocol::REQUEST_HEADERS);
}
#[tokio::test]
async fn heartbeat_callback_failure_does_not_send_fake_ack() {
let hub = HubRouter::new(ControlPlaneClient::local(
|_payload| Box::pin(async { Err("db unavailable".to_string()) }),
|_node_id, _connected, _conn_count, _observed_at_unix_secs| Box::pin(async { Ok(()) }),
));
let (proxy_tx, mut proxy_rx) = bounded_queue(8);
let (proxy_close_tx, _) = watch::channel(false);
let proxy = Arc::new(ProxyConn::new(
300,
"node-3".to_string(),
"Node 3".to_string(),
proxy_tx,
proxy_close_tx,
16,
));
hub.register_proxy(proxy);
let payload = serde_json::to_vec(&serde_json::json!({
"node_id": "node-3",
"heartbeat_id": 99u64,
}))
.expect("payload should serialize");
let mut frame = protocol::encode_frame(1, protocol::HEARTBEAT_DATA, 0, &payload);
hub.handle_proxy_frame(300, &mut frame).await;
assert!(proxy_rx.try_recv().is_err());
}
#[tokio::test]
async fn second_stream_works_after_first_completes_via_stream_end() {
let hub = HubRouter::new(ControlPlaneClient::disabled());
let (proxy_tx, mut proxy_rx) = bounded_queue(8);
let (proxy_close_tx, _) = watch::channel(false);
let proxy = Arc::new(ProxyConn::new(
400,
"node-reuse".to_string(),
"Node Reuse".to_string(),
proxy_tx,
proxy_close_tx,
16,
));
hub.register_proxy(Arc::clone(&proxy));
// First request: open stream, send body, simulate proxy response + STREAM_END
let stream1 = hub
.open_local_stream("node-reuse", &build_meta())
.expect("open first stream");
let _ = proxy_rx.try_recv().expect("first headers frame");
hub.push_local_request_body(stream1.id, Bytes::new(), true)
.expect("first body");
let _ = proxy_rx.try_recv().expect("first body frame");
// Simulate proxy sending RESPONSE_HEADERS
let resp_meta = serde_json::to_vec(&serde_json::json!({
"status": 200,
"headers": []
}))
.unwrap();
let mut resp_headers_frame = protocol::encode_frame(
// Extract the proxy_stream_id from the request headers frame
2, // first stream_id allocated
protocol::RESPONSE_HEADERS,
0,
&resp_meta,
);
hub.handle_proxy_frame(400, &mut resp_headers_frame).await;
// Simulate proxy sending STREAM_END
let mut end_frame = protocol::encode_frame(2, protocol::STREAM_END, 0, &[]);
hub.handle_proxy_frame(400, &mut end_frame).await;
// Verify stream_count went back to 0
assert_eq!(
proxy
.stream_count
.load(std::sync::atomic::Ordering::Relaxed),
0,
"stream_count should be 0 after STREAM_END"
);
// Second request: should work
let stream2 = hub
.open_local_stream("node-reuse", &build_meta())
.expect("open second stream should succeed");
let second_headers = proxy_rx.try_recv().expect("second headers frame");
let second_data = match second_headers {
Message::Binary(data) => data.to_vec(),
other => panic!("unexpected message: {other:?}"),
};
let header =
protocol::FrameHeader::parse(&second_data).expect("second request header frame");
assert_eq!(header.msg_type, protocol::REQUEST_HEADERS);
assert_ne!(
header.stream_id, 2,
"second stream should have different stream_id"
);
// Simulate proxy response for second stream
let mut resp2_headers =
protocol::encode_frame(header.stream_id, protocol::RESPONSE_HEADERS, 0, &resp_meta);
hub.handle_proxy_frame(400, &mut resp2_headers).await;
let response = stream2
.wait_headers(std::time::Duration::from_secs(1))
.await
.expect("second stream should receive headers");
assert_eq!(response.status, 200);
}
}
+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))
}
}
@@ -5,14 +5,13 @@
use std::sync::Arc;
use std::time::Duration;
use aether_runtime::bounded_queue;
use axum::extract::ws::{Message, WebSocket};
use futures_util::{SinkExt, StreamExt};
use tokio::sync::watch;
use tokio::sync::{mpsc, watch};
use tracing::{debug, info, warn};
use super::hub::{ConnConfig, HubRouter, ProxyConn, SendStatus};
use super::protocol;
use crate::hub::{ConnConfig, HubRouter, ProxyConn, SendStatus};
use crate::protocol;
/// Maximum single frame size: 64 MB
const MAX_FRAME_SIZE: usize = 64 * 1024 * 1024;
@@ -28,7 +27,7 @@ pub async fn handle_proxy_connection(
let conn_id = hub.alloc_conn_id();
let (mut ws_tx, ws_rx) = ws.split();
let (tx, mut rx) = bounded_queue::<Message>(cfg.outbound_queue_capacity);
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(
@@ -42,51 +41,13 @@ pub async fn handle_proxy_connection(
hub.register_proxy(conn.clone());
let writer_conn_id = conn_id;
let writer_conn = conn.clone();
let writer = tokio::spawn(async move {
let mut frames_sent: u64 = 0;
loop {
tokio::select! {
msg = rx.recv() => match msg {
Some(msg) => {
let is_binary = matches!(&msg, Message::Binary(_));
let msg_len = match &msg {
Message::Binary(b) => b.len(),
_ => 0,
};
let send_result = tokio::time::timeout(
Duration::from_secs(15),
ws_tx.send(msg),
).await;
match send_result {
Ok(Ok(())) => {}
Ok(Err(e)) => {
warn!(
conn_id = writer_conn_id,
frames_sent = frames_sent,
error = %e,
"writer ws_tx.send failed"
);
break;
}
Err(_) => {
warn!(
conn_id = writer_conn_id,
frames_sent = frames_sent,
"writer ws_tx.send timed out"
);
break;
}
}
frames_sent += 1;
if is_binary && msg_len > protocol::HEADER_SIZE {
debug!(
conn_id = writer_conn_id,
size = msg_len,
frames_sent = frames_sent,
"writer sent binary frame"
);
if ws_tx.send(msg).await.is_err() {
break;
}
}
None => break,
@@ -98,13 +59,7 @@ pub async fn handle_proxy_connection(
}
}
}
info!(
conn_id = writer_conn_id,
frames_sent = frames_sent,
"writer task exiting"
);
writer_conn.request_close();
let _ = tokio::time::timeout(Duration::from_secs(5), ws_tx.close()).await;
let _ = ws_tx.close().await;
});
let ping_conn = conn.clone();
@@ -146,7 +101,6 @@ async fn run_proxy_reader(
) {
let idle_enabled = !idle_timeout.is_zero();
let mut oversized_count = 0u32;
let mut frames_received: u64 = 0;
loop {
let msg = if idle_enabled {
tokio::select! {
@@ -164,7 +118,6 @@ async fn run_proxy_reader(
match msg {
Some(Ok(Message::Binary(data))) => {
frames_received += 1;
let mut data = data.to_vec();
if data.len() > MAX_FRAME_SIZE {
oversized_count += 1;
@@ -190,26 +143,13 @@ async fn run_proxy_reader(
hub.handle_proxy_frame(conn.id, &mut data).await;
}
Some(Ok(Message::Close(_))) | None => {
info!(
conn_id = conn.id,
node_id = %conn.node_id,
frames_received = frames_received,
"proxy WebSocket closed"
);
info!(conn_id = conn.id, node_id = %conn.node_id, "proxy WebSocket closed");
break;
}
Some(Err(e)) => {
warn!(
conn_id = conn.id,
frames_received = frames_received,
error = %e,
"proxy WebSocket error"
);
warn!(conn_id = conn.id, error = %e, "proxy WebSocket error");
break;
}
Some(Ok(Message::Ping(payload))) => {
conn.send(Message::Pong(payload));
}
_ => {}
}
}
+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
@@ -1,15 +1,12 @@
[package]
name = "aether-proxy"
version = "0.3.2"
version = "0.2.5"
edition = "2021"
description = "Tunnel proxy for Aether"
[dependencies]
aether-contracts.workspace = true
aether-http.workspace = true
aether-runtime.workspace = true
tokio = { version = "1", features = ["full"] }
reqwest.workspace = true
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"
@@ -19,8 +16,9 @@ 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.workspace = true
serde_json = "1"
thiserror = "2"
bytes = "1"
sha2 = "0.10"
@@ -40,6 +38,7 @@ socket2 = { version = "0.5", features = ["all"] }
tower-service = "0.3"
webpki-roots = "0.26"
[dev-dependencies]
aether-gateway.workspace = true
axum.workspace = true
[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"]
@@ -6,26 +6,26 @@ Tunnel 模式下代理节点**无需对外监听端口**,仅需出站连接到
## 安装
`aether-proxy` 会根据宿主机自动选择服务管理器:
- 常规 Linux 发行版:`systemd`
- Alpine Linux:`OpenRC`
### 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 (GNU) | [aether-proxy-linux-amd64.tar.gz](https://github.com/fawney19/Aether/releases/download/proxy-v0.3.2/aether-proxy-linux-amd64.tar.gz) |
| Linux ARM64 (GNU) | [aether-proxy-linux-arm64.tar.gz](https://github.com/fawney19/Aether/releases/download/proxy-v0.3.2/aether-proxy-linux-arm64.tar.gz) |
| Linux x86_64 (musl) | [aether-proxy-linux-musl-amd64.tar.gz](https://github.com/fawney19/Aether/releases/download/proxy-v0.3.2/aether-proxy-linux-musl-amd64.tar.gz) |
| Linux ARM64 (musl) | [aether-proxy-linux-musl-arm64.tar.gz](https://github.com/fawney19/Aether/releases/download/proxy-v0.3.2/aether-proxy-linux-musl-arm64.tar.gz) |
| macOS x86_64 | [aether-proxy-macos-amd64.tar.gz](https://github.com/fawney19/Aether/releases/download/proxy-v0.3.2/aether-proxy-macos-amd64.tar.gz) |
| macOS ARM64 | [aether-proxy-macos-arm64.tar.gz](https://github.com/fawney19/Aether/releases/download/proxy-v0.3.2/aether-proxy-macos-arm64.tar.gz) |
| Windows x86_64 | [aether-proxy-windows-amd64.zip](https://github.com/fawney19/Aether/releases/download/proxy-v0.3.2/aether-proxy-windows-amd64.zip) |
| 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 -->
上表展示的是最新已发布版本的下载链接。从下一次 `proxy-v*` 发布开始,表格会自动补上 `Linux x86_64 (musl)` / `Linux ARM64 (musl)` 包,供 Alpine 等 musl 系统直接使用。
## 快速开始
```bash
@@ -34,7 +34,7 @@ sudo ./aether-proxy setup
# 2. 日常管理 (勾选 Install Service 作为系统服务的情况下)
aether-proxy status # 看状态
sudo aether-proxy logs # 看日志
aether-proxy logs # 看日志
sudo aether-proxy start # 启动服务
sudo aether-proxy stop # 停止服务
@@ -47,7 +47,7 @@ sudo aether-proxy setup
sudo aether-proxy uninstall
```
完成向导后, 配置自动保存到 `aether-proxy.toml`,如果启用了 Install Service,将自动注册并启动当前系统支持的服务(`systemd` 或 `OpenRC`)。
完成向导后, 配置自动保存到 `aether-proxy.toml`,如果启用了 Install Service,将自动注册并启动 systemd 服务。
### 直接运行
@@ -73,34 +73,26 @@ sudo aether-proxy uninstall
|------|----------|--------|------|
| `--aether-url` | `AETHER_PROXY_AETHER_URL` | **必填** | Aether 服务器地址 |
| `--management-token` | `AETHER_PROXY_MANAGEMENT_TOKEN` | **必填** | 管理员 Token(`ae_xxx` 格式) |
| `--node-name` | `AETHER_PROXY_NODE_NAME` | **必填** | 节点名称标识 |
| `--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` | `5` | 心跳间隔(秒) |
| `--heartbeat-interval` | `AETHER_PROXY_HEARTBEAT_INTERVAL` | `30` | 心跳间隔(秒) |
| `--allowed-ports` | `AETHER_PROXY_ALLOWED_PORTS` | `80,443,8080,8443` | 允许代理的目标端口 |
| `--allow-private-targets` | `AETHER_PROXY_ALLOW_PRIVATE_TARGETS` | `true` | 允许 private/reserved 目标地址,通过后仍受 `allowed_ports` 限制;设为 `false` 可恢复严格拦截 |
#### Tunnel 连接
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--tunnel-connections` | `AETHER_PROXY_TUNNEL_CONNECTIONS` | 自动(硬件估算) | 最小连接池大小;显式设置后默认固定为该值 |
| `--tunnel-connections-max` | `AETHER_PROXY_TUNNEL_CONNECTIONS_MAX` | 自动(硬件估算) | 连接池自动扩容上限;大于 `tunnel_connections` 时启用 autoscale |
| `--tunnel-connections` | `AETHER_PROXY_TUNNEL_CONNECTIONS` | `3` | 到 Aether 的连接池大小 |
| `--tunnel-max-streams` | `AETHER_PROXY_TUNNEL_MAX_STREAMS` | 自动(硬件估算) | 单连接最大并发 stream 数 |
| `--tunnel-ping-interval-ms` | `AETHER_PROXY_TUNNEL_PING_INTERVAL_MS` | `250` | fast-fail 探测周期(毫秒) |
| `--tunnel-connect-timeout-ms` | `AETHER_PROXY_TUNNEL_CONNECT_TIMEOUT_MS` | `800` | fast-reconnect 建连超时(毫秒) |
| `--tunnel-stale-timeout-ms` | `AETHER_PROXY_TUNNEL_STALE_TIMEOUT_MS` | `900` | 无入站数据断连阈值(毫秒) |
| `--tunnel-scale-check-interval-ms` | `AETHER_PROXY_TUNNEL_SCALE_CHECK_INTERVAL_MS` | `1000` | autoscale 采样周期(毫秒) |
| `--tunnel-scale-up-threshold-percent` | `AETHER_PROXY_TUNNEL_SCALE_UP_THRESHOLD_PERCENT` | `70` | 单 tunnel 占用率超过该值时扩容 |
| `--tunnel-scale-down-threshold-percent` | `AETHER_PROXY_TUNNEL_SCALE_DOWN_THRESHOLD_PERCENT` | `35` | 单 tunnel 占用率持续低于该值时允许缩容 |
| `--tunnel-scale-down-grace-secs` | `AETHER_PROXY_TUNNEL_SCALE_DOWN_GRACE_SECS` | `15` | 低负载持续时间达到该值后才回收次级 tunnel |
| `--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-reconnect-base-ms` | `AETHER_PROXY_TUNNEL_RECONNECT_BASE_MS` | `50` | 指数退避基础延迟(毫秒) |
| `--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` | 指数退避上限(毫秒) |
省略 `tunnel_connections` 时,proxy 会按设备能力自动计算一个基线值和扩容上限;如果显式设置了 `tunnel_connections` 但没有设置 `tunnel_connections_max`,则保持固定连接池,不自动扩缩。
#### 上游 HTTP 请求
| 参数 | 环境变量 | 默认值 | 说明 |
@@ -110,7 +102,6 @@ sudo aether-proxy uninstall
| `--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 |
| `--redirect-replay-budget-bytes` | `AETHER_PROXY_REDIRECT_REPLAY_BUDGET_BYTES` | `5M` | 307/308 请求体重放的预读预算,支持 `K/M/G`,`0` 表示禁用 body replay buffering |
#### Aether API 客户端
@@ -124,7 +115,6 @@ sudo aether-proxy uninstall
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--allow-private-targets` | `AETHER_PROXY_ALLOW_PRIVATE_TARGETS` | `true` | 默认允许 private/reserved 目标地址;设为 `false` 可恢复拦截,且仅影响重启后的进程 |
| `--dns-cache-ttl-secs` | `AETHER_PROXY_DNS_CACHE_TTL_SECS` | `60` | DNS 缓存 TTL(秒) |
| `--dns-cache-capacity` | `AETHER_PROXY_DNS_CACHE_CAPACITY` | `1024` | DNS 缓存容量(条目数) |
@@ -133,23 +123,11 @@ sudo aether-proxy uninstall
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--log-level` | `AETHER_PROXY_LOG_LEVEL` | `info` | 日志级别 |
| `--log-destination` | `AETHER_PROXY_LOG_DESTINATION` | `stdout` | 输出到 `stdout`、文件或两者同时输出 |
| `--log-dir` | `AETHER_PROXY_LOG_DIR` | 空 | 文件日志目录,`file/both` 时必填 |
| `--log-rotation` | `AETHER_PROXY_LOG_ROTATION` | `daily` | 文件日志按小时或按天轮转 |
| `--log-retention-days` | `AETHER_PROXY_LOG_RETENTION_DAYS` | `7` | 文件日志保留天数 |
| `--log-max-files` | `AETHER_PROXY_LOG_MAX_FILES` | `30` | 文件日志最多保留文件数 |
### 日志落点
- 默认 `AETHER_PROXY_LOG_DESTINATION=stdout`,日志交给容器日志驱动或宿主机服务管理器
- 需要落盘时改成 `file` 或 `both`,并设置 `AETHER_PROXY_LOG_DIR`;setup TUI 里用 `Save Logs to File` 开关即可
- 文件日志固定写普通文本,并支持 `hourly/daily` 轮转;默认按天轮换、保留 7 天,最多保留 30 个文件
- 以 `systemd` 或 `OpenRC` 安装时默认会额外打开文件日志到 `/var/log/aether-proxy`
- OpenRC 安装时,`aether-proxy logs` 实际读取 `/var/log/aether-proxy/current.log` 和 `/var/log/aether-proxy/error.log`;这些文件通常需要用 `sudo aether-proxy logs` 查看
| `--log-json` | `AETHER_PROXY_LOG_JSON` | `false` | JSON 格式日志 |
### 多服务器配置
在 `aether-proxy.toml` 中使用 `[[servers]]` 配置 Aether 服务器。即使只有一个服务器,也必须写成一个 `[[servers]]` 条目;旧的顶层单服务器写法已不再支持。
在 `aether-proxy.toml` 中使用 `[[servers]]` 配置多个 Aether 服务器:
```toml
[[servers]]
@@ -167,6 +145,7 @@ node_name = "jp-proxy-02"
推送 `proxy-v*` 格式的 tag,GitHub Actions 会自动:
- 编译所有平台二进制并发布到 Releases
- 构建 Docker 镜像并推送到 GHCR 和 Docker Hub
- 更新 README 中的下载链接表格
```bash
+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);
}
}
}
}

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