mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-05 00:47:48 +08:00
Compare commits
11
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4e96b9d870 | ||
|
|
d3c0b1aa7f | ||
|
|
66bfd3e592 | ||
|
|
392c557831 | ||
|
|
c34565b02b | ||
|
|
f57fe6e13e | ||
|
|
7c678b715f | ||
|
|
bd4e5f3a5d | ||
|
|
165d9eab8f | ||
|
|
dfb95f09e1 | ||
|
|
4d5c591654 |
+16
-4
@@ -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
@@ -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
|
||||
#
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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/
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
3.13
|
||||
Generated
-5356
File diff suppressed because it is too large
Load Diff
-95
@@ -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
@@ -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
@@ -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"]
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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,如果迁移涉及不可逆的数据变更(如删除列),可能无法完全恢复数据。因此强烈建议升级前备份。
|
||||
|
||||
---
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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"]
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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"]
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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"]
|
||||
@@ -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)
|
||||
@@ -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"]
|
||||
@@ -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",
|
||||
]
|
||||
@@ -0,0 +1,3 @@
|
||||
target/
|
||||
.git/
|
||||
.DS_Store
|
||||
Generated
+2012
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
@@ -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"]
|
||||
@@ -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` 固定版本)。
|
||||
Executable
+272
@@ -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
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
@@ -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")))
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
}
|
||||
+9
-69
@@ -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));
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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"
|
||||
@@ -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 => {},
|
||||
}
|
||||
}
|
||||
@@ -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
Reference in New Issue
Block a user