Compare commits

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

* chore: ignore local worktrees directory

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

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

---------

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

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

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

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

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

Closes #256

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

Closes #255

Co-authored-by: hemo94931 <[email protected]>
2026-03-23 17:31:29 +08:00
fawney19 46737d32f8 feat: 引入 status_snapshot 统一 provider key 状态管理
- 新增 StatusSnapshot 模型,聚合 OAuth / 账号 / 配额三维状态
- 新增 StatusSnapshotStore 负责快照的持久化与查询
- 重构 response_builder / endpoint_models,基于 snapshot 输出状态字段
- 前端抽取 providerKeyStatus / oauthRefreshFeedback 工具函数,
  统一 PoolManagement、ProviderDetailDrawer、BatchDialog 的状态展示
- errorParser 增加已知 OAuth 错误的友好提示
- refresher 适配 snapshot 写入,account_state 扩展状态分类
- 新增 alembic 迁移及存量数据回填脚本
- 补充前后端单元测试
2026-03-20 19:16:52 +08:00
fawney19 25d38ae632 feat(oauth): 账号封禁前置 OAuth 验证、抽取 provider_context、完善账号状态分类
- 新增 verify_oauth_before_account_block:在标记账号封禁前先尝试刷新 token,
  区分 OAuth 过期与真正的账号级封禁,避免误标
- 抽取 provider_context.py 统一解析 provider_type,解决 ORM detached 访问问题
- account_state 新增 workspace_deactivated 分类和 auto-removable 状态集合,
  补充中文验证关键词匹配
- OAuth refresh 成功后仅清除可恢复的 token 错误,不再自动清除账号级 block
- deploy.sh 依赖指纹改用纯 shell 实现,移除对 Python tomllib 的依赖
- 前端 Pool 管理页面新增筛选和批量操作优化
- 补充对应测试用例
2026-03-20 16:50:59 +08:00
fawney19 aa83b4a7a7 fix: 加固续租失败处理、verify_auth 异常捕获及调度器注册追踪
- task_coordinator: 续租连续失败 5 次后主动触发 lock_lost 回调,失败间加指数退避
- proxy_nodes: lock_lost 回调由 lambda 改为具名 async 函数,确保异步停止逻辑正确执行
- provider_ops: 将 prepare_verify_config 纳入外层 try,捕获 ValueError 并返回失败响应
- maintenance_scheduler: 用 _registered_job_ids 动态追踪已注册任务,stop 时按列表清理
- stats_aggregator: 内联 _do_aggregate 为 for/range(2) 循环,消除内嵌函数
2026-03-20 01:22:07 +08:00
fawney19 913ce2dbcb Merge pull request #252 from AAEE86/nn
fix(startup): 收口 leader 失锁后的后台任务
2026-03-20 01:13:01 +08:00
fawney19 cae5e520ac fix(frontend): 修复 restoreOriginalPlaceholder 递归调用、优化日志参数格式,补全 tsconfig lib 配置 2026-03-20 01:06:15 +08:00
fawney19 772f2ea601 fix(vertex): SA 认证注入代理配置,细化 token 获取异常处理
- _auth_service_account 接收 endpoint 参数,通过 _get_proxy_config 解析代理
- vertex_auth 区分 TimeoutException/RequestError/通用异常,提供可读错误信息
- 新增测试覆盖代理传递和超时场景
2026-03-20 00:49:44 +08:00
fawney19 28fa03451c feat(provider): 模型测试支持自定义请求头,优化对话框布局与并发策略
- 前后端新增 request_headers 字段,测试时可自定义额外请求头
- ModelTestDialog 拆分为请求头/请求体并排双栏布局,增加格式化与重置按钮
- 区分 Pool 托管(并发5)和单 Key Provider(并发1)的测试并发数
- JsonImportInput 新增 multiple prop 支持单文件模式
- KeyFormDialog Service Account 输入改用 JsonImportInput,支持拖拽导入
2026-03-20 00:24:56 +08:00
fawney19 6984984c22 feat(provider): 重构模型测试对话框,加固 Vertex AI 传输层
模型测试:
- 将消息输入替换为完整 JSON 请求体编辑器,支持格式化和校验
- 新增端点选择面板,测试前可选择目标端点
- 新增调试检查器,可查看每次尝试的请求/响应头和体
- 结果视图改用 HorizontalRequestTimeline 组件展示请求追踪
- endpoint_checker 返回完整调试数据,通过 candidate extra_data 持久化

Vertex AI:
- 改进上下文检测逻辑,不再仅依赖 provider_type,支持从 base_url 推断
- Service Account 密钥现支持自动拉取模型(使用 auth_config 而非 api_key)
- 移除 Gemini Developer API 回退,API Key 仅走 Express 模式
- 端点表单为 Vertex AI 显示格式特定的默认路径模板
- 密钥格式校验仅在 auth_type/api_formats 变更时执行

其他:
- 禁用 ClaudeCode 提供商类型创建入口
- Dialog 组件新增 closeOnBackdrop 属性
2026-03-19 23:52:17 +08:00
AAEE86 1209c835c7 fix(startup): 收口 leader 失锁后的后台任务
- 为后台调度器注册失锁回调并只在 stop 成功后清空生命周期引用
- 停止调度器时移除定时 job,补充启动与任务协调器回归测试
- 降低多 worker 下重复调度风险,保持停机收口与统计聚合回归一致
2026-03-19 22:26:51 +08:00
fawney19 e4ebd5cca1 refactor(oauth): LinuxDo 备用端点回退、Basic Auth 认证,修复 session 外访问 ORM 对象
- LinuxDo provider: token/userinfo 请求增加 backup 端点自动回退
- LinuxDo provider: token 请求改用 HTTP Basic Auth 认证
- 授权 URL 构建: scope 为空时不再发送该参数
- OAuthService: 引入 OAuthAuthenticatedUser 快照,避免 DB session 关闭后访问 ORM 对象
- OAuthService: _handle_login_sync 设置 expire_on_commit=False 防止属性过期
- 新增 LinuxDo provider 单元测试(Basic Auth、端点回退)
- 新增 _handle_login_sync 返回快照的集成测试
2026-03-19 20:32:33 +08:00
fawney19 f573110725 fix(frontend): 用 CSS text-security 替代 password 输入框,简化配额进度条 UI
- Input 组件 masked 模式改用 WebkitTextSecurity: disc 替代 type=password,
  避免浏览器密码管理器自动填充干扰
- LoginDialog/UserFormDialog/Settings/ProxyNodes 密码字段统一迁移到 masked 属性
- PoolManagement 配额进度条布局从 grid 改为 flex,移除未使用的
  getQuotaProgressDisplayClass/getQuotaProgressTooltip 函数

Close #249
Co-authored-by: AAEE86 <[email protected]>
2026-03-19 20:04:56 +08:00
fawney19andAAEE86 a8620e133a fix(kiro): 加固 Kiro adapter 错误处理与请求构建逻辑
- 提取 request.py 统一 URL/headers/payload 构建,消除 handler_adapter_base 与 envelope 的重复逻辑
- 新增 error_enhancer 模块,分类 HTTP 状态码与连接错误,增强上游错误诊断信息
- envelope 实现 extract_error_text / on_http_status / on_connection_error,透传错误上下文到 eventstream rewriter
- KiroRequestContext 扩展错误状态字段,支持网络诊断信息传递
- provider_oauth_utils 改用 importlib 动态加载,避免 core 层对 services 的静态依赖
- idc auth_method 下跳过 profileArn,修复 usage 查询参数

Closes #247

Co-authored-by: AAEE86 <[email protected]>
2026-03-19 13:54:40 +08:00
RWDai ddd6adbcf7 feat(provider): support custom prompts for model tests (#242) 2026-03-19 13:11:54 +08:00
fawney19 b570aaac48 fix: 补全 OpenAI 工具参数 schema 中缺失的 properties 字段
OpenAI API 要求 type=object 的 schema 节点必须声明 properties,
否则会拒绝请求。在 request_from_internal 出口处对工具参数 schema
进行深拷贝并递归补全缺失的空 properties,不影响内部表示。
2026-03-19 02:36:22 +08:00
fawney19 086efe6efe refactor(usage-queue): ACK 后立即删除消息,清理 consumer 元数据
- 消费成功/移入 DLQ 后立即 XDEL,避免主流 Redis 保留已入库历史
- 缩小 usage_queue_stream_maxlen 默认值:200000 -> 2000(仅作短暂缓冲)
- 启动时清理 pending=0 且长期闲置的旧 consumer(防 consumer group 元数据累积)
- 停机时主动 XGROUP DELCONSUMER 移除自身
- 移除 cache_fingerprint 模块及其对 telemetry/recording_helpers 的引用
- 同步更新相关测试,验证 xdel 调用及 consumer 生命周期行为
2026-03-19 02:09:09 +08:00
fawney19 56f3c95763 fix: 统一 input_context_expr 计算口径,移除按 api_format 分支的 CASE 逻辑
input_context_expr() 原先按 OpenAI/Gemini 和 Claude 分支计算输入上下文,
现统一为 input_tokens + cache_read_input_tokens,与 usage 表展示口径一致。
同步更新 admin 和 user_me 路由中的缓存命中率注释,并新增单元测试。
2026-03-19 01:35:00 +08:00
fawney19 8a6a961900 feat: 增强 OAuth 标识展示与 Codex 调试日志
- 优化 OAuth badge 支持 account_id/account_user_id,tooltip 展示完整身份信息
- 重构池管理额度进度条布局,倒计时独立展示并增加样式区分
- Codex 插件增加 token 解析快照日志,便于排查 OAuth 透传问题

Closes #245
Co-Authored-By: kayphoon <[email protected]>
2026-03-19 00:50:18 +08:00
kayphoon 6e55968487 feat: Codex account_name 透传并优化池列表 OAuth 标识与额度重置倒计时展示
(cherry picked from commit 1b9176fc4a)
2026-03-18 23:58:42 +08:00
fawney19 8d8cddcef6 fix(replay): rerun model mapping on replay target 2026-03-18 23:54:19 +08:00
RWDai b90d5095f1 Fix replay fallback when model mapping is missing 2026-03-18 23:54:19 +08:00
RWDai a4505b1281 Fix usage replay model remapping 2026-03-18 23:54:19 +08:00
fawney19 1d72a8f9c1 feat: 流式空闲超时、健康监控查询优化、限流桶内存上限与维护清理修复
Close #233

Co-authored-by: AAEE86 <[email protected]>

- cli_monitor_mixin: 引入 STREAM_IDLE_TIMEOUT_SECONDS(可通过环境变量配置),
  流传输开始后若超出空闲窗口无新 chunk 则提前取消并返回 504,避免长时间挂起
- stream_context: 新增 managed_recorded_bodies 上下文管理器,确保 chunks 在
  telemetry 完成后及时释放;stream_telemetry 使用该接口统一管理 response body 构建
- health endpoint: 将状态聚合改为 GROUP BY 直接统计,事件列表按 api_format
  单独查询,避免单次 limit 拉取大量记录导致的遗漏与性能问题;同时过滤不活跃
  provider/endpoint,与公开健康接口保持一致
- endpoint health service: 修正时间线数据按 endpoint_id 而非 key_id 聚合
- token_bucket: 引入 max_buckets/bucket_expiry 上限与定时清理,防止内存无限增长;
  修复 refill_rate=0 时 get_reset_time 除零异常;新增 _is_unlimited_rate_limit 判断
- maintenance_scheduler: 调整清理顺序(先删整行再按窗口清理),新增 newer_than
  边界参数,避免同一行在同一轮中被重复改写
- sync_execute: 新增 create_pending_usage 开关,允许已预创建记录的调用方跳过重复创建
- quota_reader / provider_ops balance: 小幅修复与健壮性提升
- Dockerfile: 添加 MALLOC_ARENA_MAX=2 环境变量以降低 gunicorn worker RSS
- 补充相关测试覆盖
2026-03-18 23:38:26 +08:00
fawney19 3d5b6141a5 feat(codex): 引入 upstream_headers hook 机制,为 Codex 注入 session/conversation/account headers
- 新增 upstream_headers.py:可注册 provider+endpoint 维度的 extra headers 构建 hook
- Codex openai:cli 注入 session_id + conversation_id(由 prompt_cache_key sha256 派生)
- Codex openai:compact 注入 chatgpt-account-id(来自 auth_config)+ session_id,不注入 conversation_id
- 修复 prompt_cache:compact 格式现统一为 codex 策略,不再跳过注入
- chat_handler_base / cli_request_mixin 均在 extra_headers 阶段调用 build_upstream_extra_headers
2026-03-18 20:43:14 +08:00
fawney19 203cd5a9d5 fix(redis): 按事件循环隔离 Redis 连接,防止子线程 asyncio.run 导致连接泄漏
RedisClientManager 新增 _redis_by_loop 字典,按 event loop id 维护独立连接,
避免 usage consumer 在 asyncio.to_thread + asyncio.run 场景下复用主循环连接。

同步重构 consumer_streams 的写库路径:_apply_record_event 统一走
record_usage_batch,_apply_streaming_event 改为 to_thread 执行同步 DB 操作,
移除冗余的 session 传递和手动 rollback 逻辑。
2026-03-18 17:56:48 +08:00
fawney19 696ec65175 refactor(normalizer): 移除 request key reorder 机制,保持自然插入顺序
移除 OpenAI/OpenAI CLI normalizer 中的 _reorder_request_prefix_keys 及
基类 _reorder_request_keys 死代码,request_from_internal 直接返回构建
顺序的 dict,测试同步更新为验证自然插入顺序。
2026-03-18 14:01:56 +08:00
fawney19 cbb66a5667 refactor(task): 引入 MutableRequestBodyState 替代 request_body_ref 字典容器
将请求体可变状态从 {"body": dict} 字典容器重构为独立的
MutableRequestBodyState 类,统一管理 original_body / current_body /
build_attempt_body / rectify 等语义,消除各层通过 ref["body"] 间接
读写的隐式约定。

- 新增 src/services/task/request_state.py 定义 Protocol 与实现
- handler/executor/mixin 层改用 request_state 参数传递
- error_handler/state_transition 通过 request_state 判断整流状态
- 新增 request_state 单元测试与 chat/cli 请求体隔离测试
2026-03-18 13:43:39 +08:00
fawney19 53ef35ec80 refactor(cli): 提取 _build_upstream_request 统一流式/非流式的上游请求构建逻辑
将 cli_stream_mixin 和 cli_sync_mixin 中重复的上游请求构建代码
(provider behavior / stream policy / envelope / auth / RequestBuilder / URL 构建)
提取到 cli_request_mixin._build_upstream_request,返回 CliUpstreamRequestResult dataclass。
2026-03-18 13:13:47 +08:00
fawney19 1af3067303 refactor(codex): 移除 envelope/request_patching 层,用 context var 统一 compact 状态判断
- 删除 CodexOAuthEnvelope 和 request_patching 模块,Codex 不再需要 envelope 层
- 移除 _aether_compact 请求体内部标记,改用 is_codex_compact_request() 集中查询
- 简化 OpenAI CLI adapter,移除 Codex 专用的 get_cli_extra_headers/build_test_request_body 逻辑
- 移除 Codex behavior variant 注册(same_format/cross_format)
- normalizer patch_same_format_request 对 codex 变为 no-op
- 更新相关测试适配新的架构
2026-03-18 12:35:57 +08:00
fawney19 d026398bab refactor(body-rules): 移除 protected_body_keys 机制,允许 body_rules 自由修改所有请求体字段
删除 get_cache_sensitive_protected_body_keys 函数及相关常量、_is_protected_path
辅助函数,从 apply_body_rules、RequestBuilder、ProviderRequestResult 等处移除
protected_body_keys 参数,更新所有调用点和测试用例。
2026-03-18 10:52:07 +08:00
fawney19 684689a82b fix(usage): session touch 独立提交避免行锁阻塞 & 管理员页面顺序加载降低并发压力
后端: 将 session touch 的 commit 从请求事务中分离,防止管理员 usage
页面的长查询持有 user_sessions 行锁阻塞后续请求。touch_session 改为
返回 bool 以支持按需提交。

前端: 管理员 Usage 页面将并行 API 调用改为顺序加载,优先显示记录表格,
统计面板在后台异步刷新,避免瞬时并发打满后端 worker。loadRecords 支持
传入 dateRange 参数确保时间范围一致性。
2026-03-18 00:13:28 +08:00
github-actions[bot] eeb5f41bad chore(proxy): update download links for proxy-v0.2.5 2026-03-17 14:23:27 +00:00
fawney19 7180eaea88 chore: bump aether-proxy version to 0.2.5 2026-03-17 22:16:47 +08:00
fawney19 7cb204f18a chore: bump aether-hub version to 0.2.0 2026-03-17 22:15:07 +08:00
fawney19 0342f609d0 feat(tunnel): 请求体流式传输 & OpenAI CLI 请求 key 排序优化
Hub 端:
- open_local_stream 不再接收 body 参数,改为通过 push_local_request_body 分块推送
- 请求体按 32KB 分帧发送,避免大请求一次性压缩和传输
- local_relay 改为流式解析 envelope 和转发请求体

Proxy 端:
- stream_handler 改为流式传输请求体到上游,不再预先收集完整 body
- upstream_client 请求体类型从 Full<Bytes> 改为 UnsyncBoxBody 以支持流式传输
- dispatcher 将 StreamEnd/StreamError 事件转发给 stream handler

Python 端:
- hub_transport relay envelope 改为异步生成器流式发送
- 提取 reorder_openai_cli_request_prefix_keys 为公共函数
- Codex passthrough 路径也应用稳定的前缀 key 排序
2026-03-17 22:07:09 +08:00
fawney19 59840fa419 feat(docker): 支持通过 GITHUB_MIRROR 参数加速 hub 二进制下载
Dockerfile.app.local 新增 GITHUB_MIRROR 构建参数,deploy.sh
新增 --mirror 选项,国内服务器可指定镜像代理地址加速下载。
2026-03-17 20:43:14 +08:00
fawney19 d390d46ee8 feat(docker): Dockerfile.app.local 支持本地 hub 二进制文件
将 aether-hub tar.gz 放到项目根目录即可跳过 GitHub 下载,
解决国内服务器构建时 GitHub 访问慢的问题。
2026-03-17 20:37:09 +08:00
fawney19 37eada9682 chore: bump aether-hub version to 0.1.9 2026-03-17 20:26:29 +08:00
fawney19 73a5325a38 refactor(hub): 用本地 HTTP relay 替代 Worker WebSocket 长连接
Hub 数据面改为 /local/relay/{node_id} HTTP 端点,Worker 通过本机
HTTP 请求转发 tunnel 帧,不再维护 /worker WebSocket 长连接。

Hub 侧:
- 新增 control_plane.rs: Hub 通过 HTTP 回调 Aether app 处理心跳 ACK 和节点状态变更
- 新增 local_relay.rs: 接收本地 HTTP 请求,在 Hub 内部打开 LocalStream 并透传到 proxy
- 移除 worker_conn.rs 及 Worker WebSocket 处理逻辑
- 简化 protocol.rs: 移除 NODE_STATUS 帧类型,抽取通用 encode_frame/decode_payload

Python 侧:
- 删除 tunnel_manager.py 及其 WebSocket 连接管理器 (HubConnectionManager)
- 简化 hub_transport.py 为 HTTP relay 调用
- 新增 src/api/internal/hub.py 接收 Hub 控制面回调 (heartbeat/node-status)
- hub_config.py 移除 WebSocket 相关配置,改为 HTTP relay URL
- service.py 新增 update_tunnel_status 方法
- 删除 src/api/admin/proxy_tunnel.py (旧管理接口)
- proxy_node 缓存 TTL 从 15s 降至 3s 加速状态感知
2026-03-17 20:22:12 +08:00
fawney19 460eb5434d fix(proxy-resolver): 将 resolve_ops_proxy_config 改为异步调用避免阻塞事件循环
在 anyrouter/nekocode/sub2api/yescode 架构及 service.py 中,
将同步的 resolve_ops_proxy_config 替换为 resolve_ops_proxy_config_async,
通过 asyncio.to_thread 包装避免同步 DB 查询阻塞事件循环。
2026-03-17 18:06:15 +08:00
fawney19 8a21cb9a55 chore: bump aether-hub version to 0.1.8 2026-03-17 17:32:01 +08:00
fawney19 d0df52ce35 fix(hub): worker 断连时向 proxy 端发送 STREAM_ERROR 防止流阻塞
将 unregister_worker 中的 stream 清理逻辑提取为 cancel_streams_for_worker
方法,与 cancel_streams_for_proxy 形成对称设计。worker 断连时主动通知
proxy 端终止相关 stream,防止 proxy 的 stream handler 永久阻塞等待请求体。
添加单元测试覆盖该场景。
2026-03-17 17:29:12 +08:00
fawney19 c1ed42fd3a feat(body-rules): 支持按 provider_type 覆盖 cache-sensitive 保护字段集合
Codex 通过 OpenAI CLI/Compact 格式转发时,endpoint body_rules 需要能
修改 instructions/input/tools 等 prompt 字段。新增 provider_type 维度的
保护集合映射,Codex 场景仅保护 prompt_cache_key,其余 prompt 字段交由
body_rules 自由调整。
2026-03-17 17:09:39 +08:00
fawney19andLewisPen c4bb6b8161 feat(auth): 重构认证系统,引入 session 会话管理
- 新增 user_sessions 数据库表及 Alembic 迁移
- 实现 SessionService 会话生命周期管理(创建/刷新/撤销/清理)
- 认证流程改用 refresh token cookie + access token 双令牌模式
- 前端实现自动静默刷新、跨标签页同步及设备指纹
- 用户设置页新增会话管理和密码修改功能
- 管理员用户管理新增强制登出和会话查看
- 密码策略增强,支持强度校验和泄露检测
- OAuth 登录流程适配新会话机制
- 新增完整的单元测试和 API 测试覆盖

Closes #232

Co-authored-by: LewisPen <[email protected]>
2026-03-17 16:34:09 +08:00
fawney19 d480aa11f3 fix(request-body): 使用 deepcopy 防止请求体在处理流程中被意外修改
handler 基类和格式转换 registry 中,原始请求体通过浅拷贝或直接引用传递,
导致下游处理(模型映射、格式转换、重试整流)可能修改原始数据,
影响后续重试或并发请求的正确性。统一改用 copy.deepcopy 隔离副本。
2026-03-17 13:23:02 +08:00
fawney19 d63d5eff85 fix(nginx): 统一所有 location 块的 CF 头剥离,补充 proxy_hide_header 防止响应泄露 2026-03-17 11:33:43 +08:00
fawney19 5dae2a4792 fix(normalizer): 移除 Claude system 中的 billing header 2026-03-17 11:10:04 +08:00
fawney19 dcbd7dc219 feat(normalizer): 稳定请求体字段顺序以提升 prompt cache 命中率
在 FormatNormalizer 基类新增 _reorder_request_keys 方法,OpenAI/OpenAI CLI
normalizer 各自定义前缀字段顺序(model, tools, messages / model, instructions,
tools, input),确保静态字段前置、动态内容后置。同时调整 OpenAI CLI 中
instructions 字段的构建位置使其在 input 之前。
2026-03-17 10:26:57 +08:00
fawney19 2dcf8b8414 refactor(prompt-cache): 移除 client_family 命名空间拆分,统一缓存 key 提升复用率
不再按 User-Agent 客户端类型拆分 prompt cache namespace,
所有客户端共享同一缓存 key,版本升级至 v3。
2026-03-17 09:49:35 +08:00
fawney19 438f16094f refactor(stream): 将不完整流 token 估算逻辑收敛到 StreamContext
- 新增 has_partial_response / ensure_estimated_output_tokens / should_estimate_incomplete_tokens 方法
- CancelledError 路径在归因前即补充 output_tokens,确保日志包含估算值
- CLI Handler 和 Chat Handler 的兜底估算统一使用 should_estimate_incomplete_tokens
- 移除 cli_monitor_mixin 和 stream_telemetry 中重复的条件判断
- 新增对应单元测试
2026-03-17 03:37:06 +08:00
fawney19 40e0b82fa0 fix(chat-handler): 修复 _prepare_provider_request 缺少 original_headers 参数导致的 NameError 2026-03-17 03:21:19 +08:00
fawney19 b9b0a75fe4 feat(cache-fingerprint): 新增字段级指纹,支持逐字段 sha256 和字节数追踪
在请求缓存指纹中为每个 cache_relevant 字段单独计算 sha256 和字节数,
便于精确定位哪些字段发生了变化。fingerprint version 升级为 2。
2026-03-17 02:57:36 +08:00
fawney19 b6cc0bc3a7 refactor(usage): 统一 cache token 提取逻辑,新增请求缓存指纹记录
- 新增 extract_cache_read_tokens() 兼容 OpenAI/Claude/Gemini 多种字段命名
- parsers/stream_processor/cli_event_mixin 统一使用提取函数替换内联逻辑
- 同时兼容 prompt_tokens/completion_tokens (OpenAI) 和 input_tokens/output_tokens (Claude)
- 新增 cache_fingerprint 模块,在 telemetry 记录时自动计算并附带请求缓存指纹
- 新增对应单元测试
2026-03-17 02:34:32 +08:00
fawney19 d2f1431269 refactor(prompt-cache): 将 prompt_cache_key 生成从 Codex 专用模块提取为通用服务,支持 OpenAI 官方 API 和 Codex 端点
- 新增 prompt_cache.py 统一管理 prompt cache key 的生成逻辑
- 基于 User-Agent 识别客户端家族(openai_python/openai_node/codex_desktop 等),不同客户端生成不同 cache key
- 在 chat_handler_base/cli_stream_mixin/cli_sync_mixin 统一调用 maybe_patch_request_with_prompt_cache_key
- Codex request_patching 不再负责 prompt cache key 注入,仅保留内部标记清理
- 新增 is_official_openai_api_url 工具函数区分 OpenAI 官方 API 与兼容端点
2026-03-17 01:55:24 +08:00
fawney19 4ecaefbade refactor(conversion): body_rules 保护 cache-sensitive 字段,normalizer 保真优化与诊断日志
- RequestBuilder 新增 protected_body_keys 机制,按 provider API 格式阻止 body_rules
  改写 prompt cache 相关的顶层请求字段(messages/tools/system 等)
- Claude/Gemini normalizer 优先复用原始 raw tool_choice,避免 round-trip 丢失信息
- Claude normalizer 修复 content blocks 输出顺序(flush_text_parts),
  _coerce_claude_message_sequence 返回结构化诊断
- OpenAI normalizer 保留 raw tool call arguments 字符串与原始 tool 定义 extra 字段
- schema_utils allOf 合并保持 required 字段插入顺序
- 各转换环节增加结构化 debug 日志用于调试
2026-03-17 01:25:29 +08:00
fawney19 c97c9332eb feat(codex): 基于用户 API key 生成稳定的 prompt_cache_key,实现跨 provider key 的 prompt 缓存复用
- prepare_context 传递 user_api_key_id 到 CodexRequestContext
- wrap_request 中调用 patch_openai_cli_request_for_codex 注入 prompt_cache_key
- 客户端已提供 prompt_cache_key 时不覆盖
- 补充对应单元测试
2026-03-16 17:25:52 +08:00
fawney19 c070e5a9f6 fix(conversion): OpenAI CLI 流式转换 tool 调用与文本输出 block index 不再复用
引入统一的 block index 分配器(_allocate_block_index / _ensure_text_block_index),
避免工具调用后紧跟文本输出时 Claude block index 冲突。
2026-03-16 16:19:27 +08:00
fawney19 8cd3a69803 fix(oauth): ACCOUNT_BLOCK 标记不再禁用 key,token 刷新成功自动解除所有 invalid 标记
- error_handler 标记 ACCOUNT_BLOCK 时保持 is_active=True,oauth_invalid 标记已
  足够阻止调度,配额刷新仍可覆盖该 key
- 移除 is_account_level_block 守卫:token 刷新成功即证明账号可用,清除所有
  oauth_invalid 标记(含 ACCOUNT_BLOCK)
- health_policy 401 瞬时失败不再设置 60s cooldown,仅清除 token 缓存后立即重试
- key_quota_service 全量刷新时纳入 ACCOUNT_BLOCK 的 key,账号恢复后可自动解除
- 各 refresher 刷新成功时显式恢复 is_active=True
2026-03-16 15:50:41 +08:00
fawney19 791c9c98dc refactor(conversion): OpenAI Chat/Responses API 跨格式字段双向转换统一化
- 将工具/tool_choice/web_search/custom tool 的双向转换函数提取到 constants.py,
  openai.py 和 openai_cli.py 共享,消除两端逻辑不一致
- 新增 Chat <-> Responses 的 passthrough 字段白名单,支持 metadata/user/
  service_tier/prompt_cache_key 等字段跨格式透传
- 支持 text config (response_format + verbosity) 在 Chat/Responses 间互转
- 修复 Gemini json_schema 解包:OpenAI 的 {name, schema, strict} 包装层
  不再被整体传入 Gemini responseSchema
- 修复 reasoning_effort 优先级:显式 effort 优先于 budget_tokens 反推,
  避免 Claude output_config.effort 被覆盖
- Gemini 格式输出增加 web_search_options -> googleSearch 工具映射
- Claude schema validator 增加 web_search 类型工具的宽松校验
- 新增覆盖测试:custom tool/tool_choice、allowed_tools、web_search 双向转换、
  text config 映射、passthrough 字段保留、跨格式 schema 校验
2026-03-16 14:21:14 +08:00
fawney19 025e979935 fix(codex): 配额刷新 401/403 不再自动禁用 key,区分软性请求失败与账户封禁
- 新增 OAUTH_REQUEST_FAILED_PREFIX 标记非 token 失效的 403 请求失败
- 引入 _merge_invalid_reason 合并逻辑,避免低优先级原因覆盖高优先级状态
- account_state 中 REFRESH_FAILED/REQUEST_FAILED 前缀不再触发封禁判定
- 401/403 返回 auto_disabled=False,不再直接设置 is_active=False
2026-03-16 01:12:35 +08:00
fawney19 131471a13f refactor(maintenance): body 压缩改为逐条独立事务,降低内存占用与锁粒度
将 _cleanup_body_fields 从批量加载完整记录改为先查询 ID 列表,
再逐条独立会话处理压缩,避免大批量事务导致的内存和锁问题。
批次大小上限从 100 降至 25,新增排序保证处理顺序确定性。
2026-03-15 23:46:01 +08:00
fawney19 7ff63077c3 feat(providers): provider 摘要列表排序增加启用状态优先
后端查询和前端展示均按 is_active 降序、priority 升序、created_at 升序排列,
确保已启用的 provider 始终排在前面。
2026-03-15 23:13:11 +08:00
fawney19 faba0cbd07 fix(openai-cli): item_id 与 call_id 不一致时 tool delta 映射错误
Responses API 中 function_call 的 item.id 和 call_id 可能不同,
增加别名注册和解析机制,确保后续 delta/done 事件统一使用 call_id。
2026-03-15 22:50:50 +08:00
fawney19 75f17935f9 fix(openai-cli): 补全 function_call 的 arguments 快照同步,修复无 delta 时参数丢失
在 output_item.done 和 function_call_arguments.done 事件中,通过 args snapshot
对比已发送的增量,补发缺失的 tool arguments delta,确保客户端收到完整参数。
2026-03-15 22:20:54 +08:00
fawney19 60842fbbb5 fix(openai): tool_call delta 重复携带 function.name,增强严格客户端兼容性
在 ToolCallStart 时记录 block_index 到 tool_name 的映射,
ToolCallDelta 时同步输出 function.name,避免严格客户端丢弃无 name 的 delta。
提取 _ss_dict 辅助方法统一 stream state 中 dict 字段的初始化逻辑。
2026-03-15 22:04:21 +08:00
fawney19 d58c27d22d fix(openai): tool_call delta 重复携带 id/type,修复严格客户端兼容性
部分 OpenAI 兼容客户端要求每个 tool_call delta chunk 都包含 id 和 type
字段,否则会将后续 delta 视为无效。通过 block_to_tool_id 映射在
ContentBlockStart 时记录 tool_id,确保后续 ToolCallDelta 能正确回填。
2026-03-15 21:51:44 +08:00
fawney19 65550159bb refactor(frontend): 优化条件编辑器 UI,限制嵌套深度为两层
- 顶层组合条件增加 ListFilter 图标按钮,可转回单条件
- 子组隐藏"+ 子组"按钮,防止无意义的深层嵌套
- 嵌套叶子节点隐藏转组合按钮,保持最多两层结构
2026-03-15 21:02:05 +08:00
fawney19 3c7ad81d62 feat(rules): 条件系统增强,支持 all/any 组合条件和 original/current 数据源切换
- evaluate_condition 支持递归 all/any 组合节点和 source 字段
- header_rules 支持 condition 条件触发,HeaderBuilder.apply_rules 透传 body/original_body
- 提取 EndpointConditionEditor 组件统一请求头/请求体规则的条件编辑 UI
- header_rules 新增服务端结构校验(action/key/from/to/condition)
- 新增组合条件、source 切换、fail-closed 等测试用例
2026-03-15 20:27:53 +08:00
fawney19 6b23c9b3ce fix(frontend): 使用记录表格"命中率"列名改为"缓存命中率" 2026-03-15 18:32:42 +08:00
fawney19 900e54d740 feat(conversion): 支持 Claude output_config.effort 跨格式转换,新增 xhigh 档位
- 新增 Claude output_config.effort 与标准化 reasoning_effort 的双向映射
- 新增 xhigh 档位(budget_tokens=8192),对应 Claude effort=max
- reasoning_effort 独立于 thinking 存入 extra,支持无 thinking 场景的跨格式传递
- OpenAI/Responses API 输出时 xhigh 自动降级为 high
2026-03-15 17:57:12 +08:00
fawney19 8cc70934da fix(openai,oauth): 修复 resp_ ID 前缀转换和 token 失效误判为账号级 block
- OpenAI normalizer: resp_ 前缀 ID 规范化为 chatcmpl- 前缀,流式 tool_call index 增加 block_index 回落
- OAuth token: 提取 token 失效关键词为共享常量,token invalidated 不再被判定为账号级 block
- Codex refresher: 403 + token invalidated 标记为 OAUTH_EXPIRED,允许 refresh_token 恢复
2026-03-15 16:56:58 +08:00
fawney19andLewisPen f92b0943b5 feat(rate-limit): 实现分层 RPM 限速,支持系统默认/用户/独立Key三级配置
- 新增用户级 rate_limit 字段,支持系统默认/用户自定义/不限制三种模式
- 独立 Key 的 rate_limit 语义调整:null=跟随系统默认,0=不限制,>0=自定义
- 实现 UserRpmLimiter 基于 Redis sliding window 的 RPM 限速引擎
- Pipeline 请求流程集成用户级 RPM 检查
- 管理后台和用户面板新增 RPM 限速配置与实时状态查看
- 系统设置新增全局默认 RPM 配置项
- 迁移脚本回填现有 API Key 的 rate_limit 默认值
- 新增用户/Key RPM 状态监控 API 和前端展示

Closes #231

Co-authored-by: LewisPen <[email protected]>
2026-03-15 14:22:59 +08:00
fawney19 920a383136 refactor(claude-code): 移除 TLS 指纹模拟配置,简化上下文构建逻辑
移除 enable_tls_fingerprint 配置项、TLS_PROFILE_CLAUDE_CODE 常量及
resolve_claude_code_tls_profile 函数,TLS profile 改为仅由 fingerprint
的 impersonate 字段决定,简化 build_and_set_claude_code_request_context
返回值为单一 context 对象。
2026-03-15 00:23:19 +08:00
fawney19 693e37d2df fix(frontend): 统计表格缓存列仅显示缓存读取token,与命中率计算口径一致
缓存列移除cache_creation_tokens,只保留cache_read_tokens,
避免缓存值大于输入值的反直觉展示。
2026-03-14 16:39:48 +08:00
fawney19 9c6036a103 fix(frontend): 统计表格token列防止换行,缓存token合并为总值显示 2026-03-14 16:26:07 +08:00
fawney19 751a4d9111 feat(usage): 统计表格新增输入/输出与缓存创建token细分展示
- 后端接口(admin/user_me/query)新增 output_tokens、cache_creation_tokens、total_input_context 字段返回
- 前端统计表格(模型/提供商/API格式)合并 Tokens 列为"输入/输出"+"缓存读取/创建"两行紧凑布局
- 合并"缓存Token"和"缓存命中率"为单一"命中率"列,减少表格宽度
- 修正 UsageRecordsTable 缩进及模板格式
2026-03-14 16:11:56 +08:00
fawney19 aafd332198 revert(frontend): 使用记录表格恢复格式与类型为独立列 2026-03-14 15:01:38 +08:00
fawney19 337cd0c505 fix(frontend): 修复模型缓存价格设为0时不生效的问题
缓存价格更新函数使用 numValue > 0 判断是否手动设置,
导致输入 0 时被错误地当作清空处理。改为检查原始输入
是否为空字符串/null/undefined 来区分清空与有效输入。
2026-03-14 14:32:58 +08:00
fawney19 b15ce9977a feat(frontend): 用量页面 UI 优化 - 合并格式/类型列、分析面板可折叠
- UsageRecordsTable: 合并 API 格式与类型为单列,缩减列宽使表格更紧凑
- MainLayout: 添加 header 级 Teleport 插入点供子页面注入操作按钮
- Usage: 用量分析面板改为可折叠,折叠状态持久化到 localStorage
2026-03-14 13:28:56 +08:00
fawney19andAAEE86 e0286aebe3 refactor: 共享请求管道、按需懒加载、流式内存护栏与连接池治理
- 抽取 ApiRequestPipeline 单例,44 个路由文件共享同一实例
- Handler/Adapter 模块级 __getattr__ 延迟导入,减少启动时间
- 新增 ensure_stream_buffer_limit() 流式内存护栏(16MB 单行 / 32MB 总量)
- HTTP 空闲连接清理与 curl_cffi LRU 会话池
- ensure_providers_bootstrapped 按需引导指定 provider_types
- Usage 事件序列化迁移至 msgpack,Redis codec 隔离
- 启动预热任务(/readyz 就绪门控)与优雅关闭
- 通知邮件模块独立开关与 SMTP 配置校验
- CryptoService DCL 线程安全修复
- 通知模块开关 DB 查询 30s 内存缓存
- /readyz 对 unknown 状态返回 503
- 预热关闭 5s 超时保护
- 预热适配器逐个 try-except 容错
- FormatConversionRegistry 哨兵模式防并发重复物化
- 流式缓冲检查无条件执行

Closes #230

Co-authored-by: AAEE86 <[email protected]>
2026-03-14 11:59:07 +08:00
fawney19 45985f1c04 Merge pull request #228 from NyaDoo/fix/config-timezone-and-key-dialog
fix: APP_TIMEZONE 循环导入 & 独立密钥额度编辑空白
2026-03-14 01:42:28 +08:00
fawney19andEntropy-Xu 776dd2f8ea feat(cleanup): 解耦 request_candidates 与 provider_api_keys 生命周期
- 移除 request_candidates.key_id 对 provider_api_keys 的外键约束(含迁移脚本)
- 删除 Key 时不再级联删除候选记录,改为独立按保留天数定时清理
- 新增 request_candidates_retention_days / request_candidates_cleanup_batch_size 配置项
- batch_delete_task 增加 lock_timeout 及超时自动降批重试机制
- cleanup_key_references 提取阶段化清理流程,移除 RequestCandidate 联动删除
- 前端 CleanupPolicySection 新增候选记录保留天数和清理批次配置

Closes #227

Co-authored-by: Entropy-Xu <[email protected]>
2026-03-14 01:33:08 +08:00
fawney19andAAEE86 bdfe4adc98 feat(usage): 修复缓存命中率计算并新增用户端 API 格式统计
- 新增 input_context_expr() 按 api_format 区分 input_tokens 语义
  (OpenAI/Gemini input_tokens 已含 cache_read,Claude 需额外加上)
- 缓存命中率统一改为基于归一化后的 total_input_context 计算
- 用户 /me/usage 接口新增 summary_by_api_format 后端聚合字段
- 前端 API 格式统计改用后端聚合数据,移除前端逐条记录手动统计
- 提取 formatHitRate 到 utils/format.ts 消除三处重复定义
- 移除 PoolManager 中未使用的 select_key 方法

Co-Authored-By: AAEE86 <[email protected]>
2026-03-14 00:33:19 +08:00
LewisPen 00a0371997 fix(frontend): 独立密钥编辑模式额度区域空白问题
v-else-if 条件遗漏导致编辑模式下非 unlimited 的密钥
额度区域既不显示 Input 也不显示提示文本,对齐 UserFormDialog
的 v-else 逻辑。
2026-03-13 10:58:34 +08:00
LewisPen ded6b5b081 fix(config): 统一 APP_TIMEZONE 至 Config 类,修复 wallet 循环导入
将散落在 scheduler / stats_aggregator / daily_usage_ledger / routes
中的 os.getenv("APP_TIMEZONE") 收归 Config.app_timezone,消除
wallet → system → maintenance_scheduler → user.preference → wallet
的循环导入链。
2026-03-13 10:54:21 +08:00
fawney19 e9678ea899 fix(admin): 优先级脏检查、base_url 校验、usage detail 延迟加载及导入数据验证
- 前端优先级管理: 保存时对比原始快照,仅提交实际变更的 provider/key 优先级,
  并限制并发请求数(SAVE_CONCURRENCY=6),避免无效 API 调用
- handler_adapter_base: _normalize_test_base_url 改为 _validate_test_base_url,
  移除对 dict 类型 base_url 的兼容,严格要求字符串输入
- provider_query: 新增 _require_test_endpoint_base_url,在测试链路提前校验
  endpoint.base_url 类型和非空
- system.py: 导入 endpoint 时通过 ProviderEndpointCreate 模型校验数据,
  拒绝非法 base_url 类型(如 dict)
- usage detail: 使用 defer() 延迟加载 body 列,通过 SQL CASE 表达式在
  数据库端计算 has_*_body 标记,减少不必要的大字段传输
- provider routes: 新建 provider 时 priority=0 边界处理,clamp 并 shift
2026-03-12 17:11:14 +08:00
fawney19 8d69f72e2a fix(stability): 健康监控 DB 操作异步化,防止 worker 超时崩溃
- executor: record_success 通过 asyncio.to_thread offload 到线程池
- error_handler: 3 处 record_failure 同样 offload 到线程池
- sync_execute: 设置 expire_on_commit=False 防止 commit 后 ORM 懒加载
- handler_adapter_base: 归一化 check_endpoint 的 base_url 输入
2026-03-12 16:17:22 +08:00
fawney19 127b4e11de ci(hub): 移除 build-hub workflow 中的 Docker 构建和推送步骤 2026-03-12 15:34:09 +08:00
fawney19 22093bed4d chore: bump aether-hub version to 0.1.7 2026-03-12 15:21:47 +08:00
fawney19 ebd53ad679 fix(hub): 修复 worker 连接清理逻辑,按完成顺序有序取消对端任务 2026-03-12 15:17:00 +08:00
fawney19 280c604327 移除deplpy.sh中每次自动拉取最新代码, 以便于回退版本 2026-03-12 14:57:15 +08:00
fawney19 ad31cdbf85 perf(stability): batch committer 异步化及降级冷却参数调优
- 将 batch_committer 的 DB commit 操作通过 asyncio.to_thread 移至线程池,避免阻塞事件循环
- 将 f-string 日志替换为 loguru 惰性格式化
- 降低事件循环延迟降级的冷却时间和乘数,加快降级恢复
2026-03-12 14:22:55 +08:00
fawney19 0112ab752b refactor(proxy): 将 proxy resolver 同步阻塞操作异步化,避免阻塞事件循环
- 为 resolve_proxy_info、resolve_delegate_config、build_proxy_url、
  get_system_proxy_config、build_post_kwargs、build_stream_kwargs 新增
  _async 异步版本,通过 asyncio.to_thread 在工作线程中执行同步 DB 查询
- 为 _proxy_node_cache 和 _system_proxy_cache 添加 threading.Lock 保护
  多线程并发读写安全
- 大 payload 的 gzip 压缩超过 64KB 阈值时走线程池,小 payload 仍在事件
  循环中同步执行以避免不必要的线程调度开销
- hub_transport 的 frame 压缩同样增加异步版本
- 删除已无调用者的同步方法 create_client_with_proxy,将其逻辑内联至
  get_upstream_client 并改为异步
- 更新所有 handler/executor/failover 调用点使用新的异步 API
- 补充 async 版本的单元测试
2026-03-12 13:59:21 +08:00
fawney19 66fec80e79 fix(frontend): 修复用户表单密码框无法输入及编辑模式额度框不显示
- 密码框: masked 属性改为始终开启,避免聚焦时 DOM 重建导致焦点丢失
- 额度框: 编辑模式下 unlimited=false 时显示"按钱包余额限制"而非空白
2026-03-12 13:31:43 +08:00
fawney19 fddfaecf5e chore: bump aether-hub version to 0.1.6 2026-03-12 12:21:52 +08:00
fawney19 71ae1a2307 feat(hub,stability): bounded outbound queue、worker liveness 检测、事件循环 watchdog 及 DB 操作异步化
- aether-hub: unbounded channel 改为 bounded channel (BoundedOutbound),队列满时标记拥塞并主动关闭连接,防止内存无限增长
- aether-hub: worker idle timeout 从命令行参数改为基于心跳的 liveness 检测,默认 60 秒
- aether-hub: 新增 ConnConfig 统一管理连接配置,新增 outbound_queue_capacity 参数
- hub_transport: 新增事件循环 watchdog,检测 lag 超过阈值时临时降级暂停新流
- gunicorn_conf: 启用 faulthandler,worker abort 时自动 dump 全部线程栈用于诊断
- health/endpoint_checker/recording: 同步 DB 操作移至 asyncio.to_thread,避免阻塞事件循环
- Dockerfile: 移除 --worker-idle-timeout 0 命令行参数,改由环境变量和默认值控制
2026-03-12 12:19:04 +08:00
fawney19 4d338ebd3d fix(priority): 允许号池聚合项在全局 Key 优先级管理中拖拽和编辑排序
之前号池(Pool)类型的 Key 在优先级管理对话框中被禁止拖拽和编辑优先级,
用户只能去号池高级设置中单独修改 global_priority,体验不一致。

现在号池聚合项可以像普通 Key 一样参与拖拽排序和点击编辑优先级,
变更会同步到 provider.pool_advanced.global_priority 并在保存时持久化。
2026-03-12 11:44:19 +08:00
fawney19 dc440f1507 fix(ci): 修正 deploy-pages action 版本号 (deploy-pages@v4, upload-pages-artifact@v3) 2026-03-12 10:41:26 +08:00
fawney19 5c732f844a chore(ci): 升级 GitHub Actions 版本,解决 Node.js 20 弃用警告
- actions/checkout v4 → v5
- actions/upload-artifact v4 → v5
- actions/download-artifact v4 → v5
- actions/setup-node v4 → v5
- actions/configure-pages v4 → v5
- actions/upload-pages-artifact v3 → v4
- actions/deploy-pages v4 → v5
- docker/build-push-action v5 → v6
- Node.js 版本 20 → 22
2026-03-12 10:34:29 +08:00
fawney19 c4044ba0b1 chore: bump aether-hub version to 0.1.5 2026-03-12 10:20:42 +08:00
fawney19 a8159b7bda fix(tunnel): 禁用 idle timeout 默认值,防止空闲 tunnel 连接被断开
Hub 端 proxy_idle_timeout 和 worker_idle_timeout 默认改为 0(禁用),
直连模式 proxy_tunnel.py 同步调整。连接存活检测依赖 PING/PONG 心跳。
2026-03-12 10:07:41 +08:00
fawney19 0ab20be667 fix(proxy): 抑制 tunnel/hub 高频重复日志,增加快速断开退避机制
- proxy_tunnel: idle timeout 日志首次 warning,后续降级为 debug/info 汇总
- hub_transport: 连续断开日志降级,追踪快速断开并调整重连初始 delay
- tunnel_manager: connect/disconnect 日志按 reconnect 计数分级输出
2026-03-12 09:56:34 +08:00
fawney19 3f048d373f refactor: 将路由层同步 DB 操作移至线程池执行,统一认证工具函数
- 管理端和用户端路由中的同步数据库操作提取为独立函数,通过 run_in_threadpool
  在线程池中执行,避免阻塞事件循环(涉及 api_keys、payments、users、wallets、
  provider_oauth、system、user_me、wallet 等模块)
- 抽取 authenticate_user_from_bearer_token 统一 token 验证逻辑,支持
  ManagementToken 和 JWT 两种认证方式,消除多处重复代码
- key_command_service 的 CRUD 操作改为线程池执行
- maintenance_scheduler 定时任务中的数据库操作改用 asyncio.to_thread
- Dockerfile 中 aether-hub 增加 --worker-idle-timeout 0 防止空闲断连
- 新增 test_api_auth_conventions 和 test_auth_utils 单元测试
2026-03-12 09:33:24 +08:00
fawney19 8b49a3d264 fix(frontend): 优化用户表单密码输入框,使用 masked 和 disable-autofill 属性简化实现 2026-03-12 01:46:41 +08:00
fawney19 6e51a3f45d feat: Provider 异步删除、可配置密码策略、Hub 超时优化及多项改进
- 新增 Provider 异步删除任务系统,后台分阶段删除子资源并清理残留引用
- 新增可配置密码策略等级(weak/medium/strong),支持系统设置面板调整
- aether-hub 升级至 0.1.4,idle timeout 支持禁用(设为 0),worker 默认超时调整为 120s
- OAuth 手动续期增加 Redis 分布式锁,防止并发刷新冲突
- ProxyNode 心跳检测改为 asyncio.to_thread,避免阻塞事件循环
- 删除 ModelMultiSelect 和 useInvalidModels,MultiSelect 组件通用化
- 明确 allowed_providers/allowed_api_formats 的 NULL 与空数组语义
- 前端 StandaloneKeyFormDialog、UserFormDialog 等多处 UI 优化
- 新增 Alembic 迁移脚本清理 Provider 删除后的残留引用
- 补充相关测试用例
2026-03-12 01:11:35 +08:00
fawney19 0d770d1c4d feat(oauth): 新增 Codex account_user_id 和 organizations 字段采集、展示与判重
- 从 Codex id_token claims 和 token_response 中提取 account_user_id 和 organizations
- OAuth 判重逻辑改为优先按 account_user_id 匹配,支持同用户不同 Team 不误判
- 号池和 Provider 详情页展示组织标签、account ID 和 account_user_id
- 前端重复的 OAuth identity 工具函数提取到 utils/oauthIdentity.ts
- 后端重复的 normalize_oauth_organizations 提取到 core/provider_oauth_utils.py
2026-03-11 21:37:08 +08:00
fawney19 b45f021bba feat(pool): 批量操作对话框改为服务端分页筛选,新增凭据导出功能
- 批量操作对话框从全量加载改为服务端分页+筛选,支持搜索和快捷选择器的服务端过滤
- 新增 resolve-selection API,支持"全选筛选结果"时解析完整匹配列表
- 新增批量导出凭据功能(仅 OAuth 账号),并发下载后导出为 JSON 文件
- 将快捷选择器和全文搜索的匹配逻辑从前端迁移到后端,统一复用
- 提取 pool key 序列化与过滤的公共函数,消除 AdminListPoolKeysAdapter 中的重复代码
- entrypoint.sh 增加 PostgreSQL 就绪等待,避免数据库未启动时迁移失败
2026-03-11 19:32:19 +08:00
fawney19 1a1bce3e8c Merge pull request #224 from AAEE86/fix/pool-last-used-at-display
fix(usage,pool): 修复号池最后使用时间不更新
2026-03-11 16:33:36 +08:00
fawney19 6aeb5d40ab Merge pull request #222 from AAEE86/fix/provider-mapping
fix(provider-mapping): 修复详情页映射延迟刷新并补齐 mapping-preview 缓存治理
2026-03-11 16:17:48 +08:00
AAEE86 31ef2d134e fix(usage,pool): 修复号池最后使用时间不更新
统一在 ProviderAPIKey 累计更新时刷新 last_used_at/updated_at,即使 token/cost 增量为 0。

补充回归测试:覆盖 zero delta 场景和 provider_api_key_id 为空时跳过更新。
2026-03-11 16:12:36 +08:00
fawney19 85b50e67e1 Merge pull request #221 from AAEE86/fix/frontend
fix(frontend): 修复健康监控中 OpenAI Compact 百分比换行
2026-03-11 16:02:38 +08:00
fawney19 9353f89af0 Merge pull request #220 from AAEE86/master
fix(pool): recent_refresh 按 provider_type 解析 reset_seconds 并锁定 Codex 周重置语义
2026-03-11 16:02:14 +08:00
fawney19andAAEE86 380d69e096 feat(pool): 记录并展示 Provider Key 累计 Token 与费用
Closes #219

Co-authored-by: AAEE86 <[email protected]>
2026-03-11 15:56:58 +08:00
fawney19 02e2f4f500 fix: 修复钱包迁移脚本 _to_decimal 遇到 NaN/Infinity 值时的 InvalidOperation 异常 2026-03-11 15:22:27 +08:00
fawney19 0dbfefa834 Merge branch 'feat/wallet-billing-state-machine-and-daily-usage' 2026-03-11 15:15:01 +08:00
fawney19andLewisPen 04ab4bd9f2 feat: 强化用量计费状态机,新增钱包每日消费汇总分类账
- 将 usage.billing_status 默认值从 settled 改为 pending,完善
  pending -> settled/void 的状态转换逻辑,确保终态不可逆
- 新增 WalletDailyUsageLedger 模型和聚合服务,按账单日汇总
  每个钱包的消费金额、请求数和 token 用量
- 前端钱包中心页面集成每日消费流水展示,支持与充值记录混合
  排序和分页
- 新增两个数据库迁移:修复历史数据状态一致性、创建每日汇总表
- 补充计费状态机单元测试

Closes #218

Co-authored-by: LewisPen <[email protected]>
2026-03-11 15:11:33 +08:00
AAEE86 4955166b85 fix(provider-mapping): 修复详情页映射延迟刷新并补齐 mapping-preview 缓存治理
- 前端:Provider 详情页在 key/模型关联/模型保存/映射保存后并行刷新 endpoints 与 mapping-preview
- 前端:为 mapping-preview 请求增加 requestId 保护,避免旧响应回写覆盖新状态
- 后端:Key.allowed_models 变更时按 provider 精准失效 mapping-preview 缓存
- 后端:GlobalModel 变更(含创建)时全量失效 mapping-preview 缓存
- 监控:在 Redis 缓存分类中新增 provider_mapping_preview(admin:providers:mapping-preview:*)

验证:
- frontend `npm run type-check`
- backend `python -m compileall`(相关文件)
- pytest 定向用例通过
2026-03-11 09:58:25 +08:00
fawney19 6235c772ac fix: 修正 PR #217 合并后的两个细节问题
- user_me usage 接口文档注释补充 summary_by_provider 字段说明
- 导出配置中 internal_priority 为 NULL 时映射为 inf,保持与原 SQL NULLS LAST 一致
2026-03-10 23:33:02 +08:00
fawney19 e2ec3f7942 Merge pull request #217 from AAEE86/1233
perf: 优化请求鉴权链路并批量化统计/调度查询
2026-03-10 23:31:28 +08:00
fawney19andEntropy.Xu a816235efb feat: 新增 Gemini CLI provider adapter
Closes #216

Co-authored-by: Entropy.Xu <[email protected]>
2026-03-10 23:15:05 +08:00
fawney19andEntropy.Xu 1e39ab3c2e feat: 新增 Gemini CLI provider adapter
- 新增 gemini_cli adapter 包(client/constants/envelope/plugin/quota)
- 实现 v1internal 协议封装、OAuth enrichment、loadCodeAssist/onboardUser 流程
- 实现配额耗尽检测与冷却元数据管理(RESOURCE_EXHAUSTED 解析)
- 新增 GeminiCliQuotaReader 支持按模型粒度的配额展示
- endpoint check 支持流式回退和 v1internal 响应解包
- health_policy 429 处理增加 Google 配额冷却解析
- error_handler 在限流时同步 Gemini CLI 配额状态
- 前端添加 Gemini CLI provider 类型选项和 OAuth 图标
- 新增 preset models(gemini-2.5-pro/flash, gemini-3-pro/flash, gemini-3.1-pro)
- 新增 gemini_cli quota 单元测试

Closes #216

Co-authored-by: Entropy.Xu <[email protected]>
2026-03-10 23:14:05 +08:00
fawney19 85aa66c76d refactor: 优化 pending 清理和压缩任务的内存使用
- 将 pending 请求清理改为分批处理,使用轻量列查询代替全量 ORM 加载
- 限制 pending 清理和历史压缩的批次大小上限,防止单次查询占用过多内存
- 移除 processed_ids 集合,改用 synchronize_session=False 避免内存累积
- 更新 README 中部署方式描述
2026-03-10 18:11:08 +08:00
AAEE86 9a4817faf8 fix(frontend): 修复健康监控中 OpenAI Compact 百分比换行
- 调整 HealthMonitorCard 左侧信息区宽度(sm:w-44 -> sm:w-52)
- 为 API 格式与成功率 Badge 添加 whitespace-nowrap,避免文本断行
2026-03-10 16:50:26 +08:00
fawney19 7b0c80a0c4 refactor: 增加内存保护措施,精简启动任务
- 流式响应不再存储完整文本,仅记录长度用于 token 估算
- complete_response 文本累积增加 64KB 上限保护
- 通知缓冲区(email/webhook)增加溢出保护,超限丢弃旧通知
- 熔断器淘汰策略改进,全部 open/half-open 时按最旧失败时间淘汰
- 移除启动时清理任务和统计聚合回填,减少启动负担
- token 估算简化为基于长度的方法,避免持有完整文本
2026-03-10 16:08:11 +08:00
fawney19 6ec8df97e8 refactor: 限制流式文本收集内存增长,降低默认连接池和缓存上限
- StreamContext.append_text 增加 16KB 上限,超出后仅计数不存储,
  避免长流式响应导致内存持续增长;token 估算改用 collected_text_length
- 降低 DB 连接池上限 (30->15) 和 HTTP 连接池上限 (200->100)
- tiktoken 编码器缓存从 32 缩减到 4(实际编码种类只有几种)
- dev.sh 添加开发环境低配连接池默认值,uvicorn 热重载仅监视 src 目录
2026-03-10 15:33:46 +08:00
AAEE86 f82964217e fix(pool): recent_refresh 按 provider_type 解析 reset_seconds 并锁定 Codex 周重置语义
- 为 `extract_reset_seconds` 增加 `provider_type` 解析链路(显式参数 -> key.provider_type -> key.provider.provider_type -> metadata 推断)。
- Codex 场景固定读取 `primary_reset_seconds`(周限额),避免误用 `secondary_reset_seconds`(5 小时窗口);Kiro/Antigravity 继续走统一 quota reader。
- 在 `RecentRefreshDimension` 和 `PoolManager` 的策略上下文中透传 `provider_type`,并在缺失时从 key provider 回退推断。
- 新增/补强测试覆盖:
  - multi_score `recent_refresh` 在 Codex 下按 weekly reset 排序;
  - quota reader helper 在 Codex provider 下返回 primary reset;
  - Codex 付费计划(plus/enterprise)在 headers 与 WHAM 响应中的窗口映射与主次窗口对齐(10080/300 分钟)。

降低 recent_refresh 评分偏差风险,并为 Codex 付费窗口解析提供回归保护。
2026-03-10 14:55:48 +08:00
fawney19 cfa5535f6e refactor: 引入 safe_create_task 防止后台任务被 GC 回收,降低默认连接池和 worker 数量
- 新增 safe_create_task 统一替代裸 asyncio.create_task,通过全局集合持有 task 引用
- 默认 worker 数量从 4 降为 1,HTTP 连接池总预算从 800 降为 200
- 为 health_cache 和 affinity_manager 内存缓存增加上限淘汰机制
- MemoryCachePlugin 支持延迟启动清理任务
- gunicorn when_ready 增加 gc.collect() 并记录 post_worker_init RSS
2026-03-10 14:42:33 +08:00
fawney19 2d846b2c58 refactor: 移除启动缓存预热功能
删除 cache_warmup.py 及相关配置项和测试,简化启动流程
2026-03-10 14:11:41 +08:00
fawney19 f40e8037dd feat: 添加启动任务开关,修复统计聚合内存泄漏
- 新增 CACHE_WARMUP_ENABLED 和 MAINTENANCE_STARTUP_TASKS_ENABLED 环境变量,
  允许禁用缓存预热和维护调度器启动任务
- 在统计聚合批量处理循环中添加 db.expunge_all(),释放 Session identity map,
  防止 ORM 对象累积导致内存暴涨
- 添加启动任务开关的单元测试
2026-03-10 13:42:38 +08:00
fawney19 2b21a75982 refactor: 优化模型获取调度器并发模型和缓存格式
- 用固定 worker 数的队列消费模式替换 Semaphore+gather,避免大号池一次性创建大量协程
- 分批扫描 Key ID(keyset pagination),避免一次性加载全量 ID
- 上游模型缓存从 per-api_format 去重改为 per-model-id 聚合 api_formats 列表,减少 Redis 占用
- 启动阶段支持环境变量控制(MODEL_FETCH_STARTUP_ENABLED / MODEL_FETCH_STARTUP_DELAY_SECONDS)
- 新增单元测试覆盖聚合逻辑和并发上限验证
2026-03-10 12:47:50 +08:00
fawney19 9ee27308db refactor: defer ProviderAPIKey 大 JSON 字段,减少调度器和模型获取的内存占用
- candidate_builder: defer adjustment_history/utilization_samples/upstream_metadata,
  号池路径在释放 DB 连接前预计算 _pool_account_state 避免后续 N+1 查询
- pool/manager: 优先读取预计算的 _pool_account_state,fallback 到实时解析
- fetch_scheduler: 三处 ProviderAPIKey 查询改用 defer/load_only 排除无关字段
2026-03-10 12:33:15 +08:00
fawney19 1518223de6 refactor: 优化调度器内存占用,增加 session 全局清理
- main: 启动阶段 Provider 查询改用聚合查询,避免全量加载 ORM 对象
- envelope: Claude Code session 增加全局定时清理,防止 scope_key 无界增长
- health_cache: 增量更新时清理已移除的 stale key 条目
- aware_scheduler: 日志中用 key.name 替代 api_key 后四位,避免泄露凭证片段
- candidate_builder: defer api_key/auth_config 等冷字段,减少调度热路径内存占用
2026-03-10 11:40:29 +08:00
fawney19 3063938a82 refactor: 优化调度器并发与内存占用,修复循环依赖
- fetch_scheduler: 模型获取从串行改为 Semaphore 限并发并行
- pool_quota_probe_scheduler: 取消启动时立即探测避免阻塞,
  改为逐 provider 独立查询 key 避免一次性加载全部到内存,
  移除未使用的 _ProviderProbeTask 数据类
- proxy_node/__init__: 移除 health_scheduler 导入避免循环依赖
2026-03-10 11:22:11 +08:00
fawney19 86449cae52 refactor: 迁移文件内联 helpers,删除共享 alembic/helpers.py
将 _SchemaCache、replace_fk_if_needed、batch_alter_type 等辅助函数
内联到各迁移文件中,使每个迁移自包含、不依赖外部模块。
同时修正 a3f1b7c9d2e4 的 down_revision 为 d7649c1f8e21。
2026-03-10 10:53:41 +08:00
fawney19 7c580e843f fix: 治理 Prometheus 指标基数爆炸和内存缓存无界增长
- 移除 token/latency Prometheus 指标的 model 标签,避免 provider x model 笛卡尔积
- HealthMonitor 滑动窗口从 DB JSON 迁移至进程内存,减少写放大
- ModelCostService 三层缓存增加 500 条上限,超限时清空
- StickyPriority 粘性缓存和健康状态字典增加容量淘汰
- AffinityManager 请求锁字典增加 500 条上限,淘汰空闲锁
- 配额刷新/探测查询使用 defer/load_only 避免加载大 JSON 列
- Alembic 迁移清理 DB 中遗留的 request_results_window 数据
- 同步更新测试适配 batch_get_cooldowns 返回值和批量删除异步化
2026-03-10 10:40:50 +08:00
fawney19 c9f0685b40 fix: 治理 Prometheus 指标基数爆炸和内存缓存无界增长
- 删除高基数标签 key_id/model,移除未使用的指标 concurrency_slots_in_use/streaming_request_duration_seconds
- HTTP 客户端池:命名客户端添加 LRU 淘汰上限,tunnel 客户端添加 LRU 淘汰和 last_used_time 追踪
- ResilienceManager:last_errors 改用 deque(maxlen),error_stats/circuit_breakers 添加上限淘汰
- HealthMonitor:_circuit_history 改用 deque(maxlen)
- ProviderHealthTracker:清理已无记录的过期 key
- PrometheusPlugin:动态指标数量添加上限,超限后拒绝创建
- 中间件使用路由模板替代实际路径,防止动态路径段导致标签爆炸
2026-03-10 09:52:21 +08:00
fawney19 afd0dcf2ff feat: OAuth 导入完成后自动触发配额刷新,统一配额常量定义
- 单个导入、批量导入、Kiro device flow 完成后后台异步刷新配额
- 将 CODEX_WHAM_USAGE_URL 和 QUOTA_REFRESH_PROVIDER_TYPES 统一到 key_quota_service.py,消除 3 处重复硬编码
2026-03-10 09:36:47 +08:00
AAEE86 68d4df71d8 perf(frontend): 优化图表更新链路与缓存监控倒计时开销
- 收敛 LineChart 的配置构建逻辑,复用 options 生成函数
- 将 LineChart 的 data/options 监听从深监听改为引用监听
- 统一使用 chart.update('none'),减少不必要的动画与重绘

- 为 ScatterChart 新增 prepareRenderData 流程,合并间隙压缩与点位转换
- 消除 createChart/updateChart 中重复的数据预处理逻辑
- 将散点图更新改为监听 data、compressGaps、gapThreshold、compressedGapSize
- 移除 compressGaps 切换时的 destroy + recreate 路径,改为原图更新
- 将散点图 options 更新改为无动画刷新,降低全量重算成本

- 为缓存监控页新增 nextExpireAt 状态,跟踪最近过期时间
- 在拉取 affinity 列表后立即按当前时间裁剪已过期数据
- 将每秒全表 filter 改为按最近过期时间触发清理
- 页面恢复可见时先补执行过期清理,再恢复倒计时
- 保留每秒 currentTime 更新,仅用于倒计时显示,降低常驻扫描开销
2026-03-10 09:09:06 +08:00
AAEE86 596227659a perf(admin-usage): 精简管理员 Usage 列表的 request_metadata 下发
- 从 request_metadata JSON 中直接查询 model_version
- 管理员 Usage 列表不再返回完整 request_metadata
- 前端表格、类型和 mock 改为使用顶层 model_version
- 新增轻量响应回归测试
2026-03-10 00:02:49 +08:00
fawney19 3b0dbadb1e fix: 单个提供商更新时不再清空全部余额缓存
loadBalances 新增 fullReload 参数,单个提供商刷新时传入 false,
避免更新一个提供商的余额数据时清除其他提供商的已加载余额。
2026-03-09 23:35:51 +08:00
fawney19 8a8bc999d2 fix: 链路追踪中提供商名称原样显示,不再自动去除"反代"后缀 2026-03-09 23:17:18 +08:00
AAEE86 57c7cca556 perf: 优化请求鉴权链路并批量化统计/调度查询
- 为 Pipeline/Context 增加按需读取请求体能力,支持 async 懒加载 JSON body
- 为 chat/cli/video/claude/openai-cli 适配器关闭默认预读,减少无效 body 读取与超时风险
- 将本地登录、JWT 用户加载、API Key 鉴权迁移到线程池隔离会话执行,避免阻塞事件循环
- 为 API Key 鉴权返回结构化余额结果,并在主请求会话中重新绑定 user/api_key 后再校验状态、过期和锁定信息
- 为 management/user token 前缀认证引入独立会话与结果回绑,避免跨会话对象写入失效

- 为 Usage 余额检查补充结构化返回,统一透出 remaining 与欠费/不可用文案映射
- 为用户与管理端活跃请求查询增加 maintain_status 开关,避免轮询指定 id 时误触发状态修复
- 重写 user_me usage 汇总逻辑,支持 group_by=None 的粗粒度聚合
- 修正 provider 维度成功率与平均响应时间统计,基于 success_count 和成功响应耗时汇总计算
- 前端 Usage 轮询由 setInterval 改为串行 setTimeout,避免并发轮询叠加

- 为 StatsAggregator 增加按本地日期批量计算百分位能力,替代逐天 fan-out 查询
- 为混合统计查询合并连续实时日期区间,并批量读取 StatsDaily,减少逐日查询次数
- 为用户日统计增加批量聚合入口,替代逐用户循环聚合
- 为系统配置导出改用 selectinload 预加载 provider 关联数据,减少 N+1 查询
- 为管理员用户列表增加钱包批量查询,避免逐用户回表

- 为调度器增加 provider 轻量引用预过滤,先按 allowed_providers 缩小范围再加载完整 provider 图
- 为 CandidateBuilder 增加 provider refs/provider_ids 查询能力,保留分页顺序
- 为模型缓存增加 provider_model_mappings 索引缓存与 model_mappings 规则缓存,减少重复全量扫描
- 为请求候选中间态改为 flush/batch commit,降低 pending/streaming 状态切换的事务往返
- 为钱包访问结果补充 balance_snapshot,并抽取余额快照复用逻辑

- 补充 pipeline、auth、admin users、user_me usage、stats aggregator、model cache、
  scheduler、wallet、request candidate 等回归与契约测试
2026-03-09 22:57:23 +08:00
fawney19 9e6578a71b fix: 对 actual_total_cost_usd 显式转换为 float,防止类型不匹配导致累加异常 2026-03-09 21:30:51 +08:00
fawney19 bf818e3b61 fix: 钱包迁移回填 clamp 超范围值,防止 NUMERIC(20,8) 溢出 2026-03-09 20:45:13 +08:00
fawney19 9e7f291aaf Merge pull request #213 from AAEE86/123
fix: 处理客户端断开连接情况,以防止出现虚假的系统错误报告
2026-03-09 20:23:56 +08:00
fawney19andAoaoMH 46bab1b97f feat: 为下拉选择组件添加搜索过滤功能,支持拼音匹配
- 新增 search.ts 搜索工具,支持中文拼音(全拼/首字母)模糊匹配
- 新增 select-search-context.ts,通过 provide/inject 为 radix-vue Select 组件注入搜索能力
- MultiSelect、ModelMultiSelect 组件集成搜索框,选项超过阈值时自动显示
- select-content/select-item 组件支持搜索过滤与空状态提示
- StandaloneKeyFormDialog、UserFormDialog 中手写下拉框替换为复用 MultiSelect 组件
- 引入 pinyin-pro 依赖,按需懒加载

Closes #210

Co-authored-by: AoaoMH <[email protected]>
2026-03-09 19:53:54 +08:00
fawney19andAAEE86 0258d01ee6 refactor: 将 adapter 层的计费/模型抓取/行为变体能力下沉到 core.api_format 注册表
- 新增 core/api_format/capabilities.py,统一注册计费模板、模型抓取、
  total_input_context 计算和 provider behavior variant
- 新增 core/usage_tokens.py,抽取 cache token 解析逻辑到 core 层
- handler adapter 移除各自的 compute_total_input_context / fetch_models /
  BILLING_TEMPLATE 覆盖,改为委托 core 注册表解析
- provider/behavior.py 改为薄封装,底层委托 core registry
- 新增 tests/test_architecture_import_rules.py 架构导入约束测试
- 新增 tests/services/api_format/test_capabilities.py 能力注册表测试

Closes #207

Co-authored-by: AAEE86 <[email protected]>
2026-03-09 18:26:53 +08:00
fawney19 4999a1a0a8 fix: batch_alter_type 转换 NUMERIC 类型前 clamp 超范围值,防止溢出 2026-03-09 16:56:41 +08:00
fawney19 8cb8666456 fix: alembic helpers 限定 schema 查询范围,修复跨 schema 误匹配
- information_schema 查询增加 current_schema() 过滤条件
- 新增 _fk_exists() 通过 pg_constraint 直接检查外键是否存在
- replace_fk_if_needed 在缓存未命中时回退到 pg_constraint 查找
- index_exists 增加 schemaname 过滤
2026-03-09 16:46:38 +08:00
fawney19 48f3f481db fix: 号池调度 label 默认维度数改为动态计算 2026-03-09 15:44:04 +08:00
fawney19 e60462e068 fix: 号池调度 label 在无配置时错误显示为 LRU + 粘性
pool_advanced 为 null 时,poolSchedulingLabel 的 fallback 逻辑因
可选链默认值导致 lruEnabled=true、stickyEnabled=true,始终显示
"LRU + 粘性"。现在提前返回 "2 维度" 以匹配后端默认行为
(cache_affinity + recent_refresh)。
2026-03-09 15:42:32 +08:00
fawney19 c4877e3b6a fix: 号池 legacy 配置默认回退到 cache_affinity,multi_score 模式补全 LRU 数据获取
_build_from_legacy_fields 在无 legacy 字段时默认返回 cache_affinity 而非 lru,
与 PoolConfig 默认值保持一致。PoolManager 在 scheduling_mode 为 multi_score 时
也获取和更新 LRU 数据,修复 cache_affinity/single_account 维度因缺少 lru_scores
导致排序失效的问题。
2026-03-09 15:22:17 +08:00
fawney19 afbb1b9a5d perf: Redis 操作优化,移除号池 key 列表的 Usage 聚合查询
- 号池 key 列表移除 Usage 表聚合查询(total_tokens/total_cost_usd),消除慢 SQL
- cooldown 计数从 SCAN 改为 SCARD (O(1)),通过 cooldown_idx SET 维护索引
- 读路径 Lua 脚本移除 ZREMRANGEBYSCORE,清理移至写路径减少开销
- affinity 清理从 KEYS 改为 SCAN 分批删除,invalidate_all_for_provider 改用 pipeline 批量 MGET+UNLINK
- 缓存监控页添加清除按钮 loading 状态,Redis SCAN 并发限制为 4
- 提取 redis_utils 模块统一 SCAN+批量删除逻辑
- 调度热路径跳过 cooldown TTL 查询,account_state 预计算移出循环
2026-03-09 14:54:45 +08:00
fawney19 9516619b92 feat: 号池调度默认选择缓存亲和,默认开启额度刷新优先 2026-03-09 14:07:20 +08:00
fawney19 e106a65c1d perf: 在 Redis 密集操作前释放 DB 连接,candidate_records 改为异步写入
- 号池排序涉及大量 Redis I/O,在调用前提前释放 DB 连接避免连接池压力
- 新增 create_candidate_records_async,通过 asyncio.to_thread 执行同步 DB 写入
- 同步执行、异步提交、TaskService 三条路径统一改用异步版本
2026-03-09 13:53:39 +08:00
fawney19 d84c9d4b71 feat: 配置导入导出支持 ProxyNode,号池 key 可用性检查延迟到排序后分页执行
配置导入导出:
- 导出/导入新增 ProxyNode(代理节点)数据
- 导入时自动建立 old_id -> new_id 映射,重映射 Provider/Endpoint/Key 中的 node_id
- 前端预览和结果展示新增代理节点统计

号池调度优化:
- CandidateBuilder 不再逐 key 调用 _check_key_availability,直接收集全部 active key
- 将可用性检查参数打包到 PoolCandidate._deferred_check_params
- PoolManager.select_pool_keys 排序后分页调用 availability_checker,找到足够可用 key 即停止
- 减少大号池场景下不必要的可用性检查开销
2026-03-09 13:26:54 +08:00
fawney19 0046123e22 feat: Provider 摘要 API 改为服务端分页,支持搜索和筛选
- 后端 /summary 接口新增 page/page_size/search/status/api_format/model_id 参数
- 新增 ProviderSummaryPageResponse 分页响应模型
- 前端 useProviderFilters 从客户端筛选改为构建服务端查询参数
- ProviderManagement 通过 watch queryParams 实现分页/筛选联动,搜索 debounce 300ms
- PriorityManagementDialog 改为对话框打开时自行加载全量 providers
- 其他使用方(StandaloneKeyFormDialog/UserFormDialog/ReplayDialog/ModelManagement)适配新接口
2026-03-09 12:45:54 +08:00
fawney19 40736a8334 refactor: 移除迁移脚本中不必要的 backfill SQL
快照字段(username/api_key_name)已在业务层写入时填充,无需迁移时回填历史数据
2026-03-09 12:17:45 +08:00
fawney19 0a256adc94 fix: 修复迁移helper 2026-03-09 11:52:16 +08:00
fawney19 0bddc7965b perf: 同步 DB 操作迁移到 asyncio.to_thread,避免阻塞事件循环
failover/stream_telemetry/recording/stream 中的同步 DB 操作(commit/execute/query)
会阻塞 asyncio 事件循环,导致 Hub PING 心跳无法发送、worker idle timeout 断连。
将这些操作包装到 asyncio.to_thread() 中执行。

同时提取 Alembic 迁移脚本中重复的幂等性辅助函数到 alembic/helpers.py,
用批量查询缓存替代逐条 information_schema 查询,backfill SQL 合并为 LEFT JOIN。

新增 failover 中客户端断连的快速终止路径,避免继续无意义的重试。
2026-03-09 11:49:03 +08:00
AAEE86 258be3b640 fix: 处理客户端断开连接情况,以防止出现虚假的系统错误报告
捕获 starlette.requests.ClientDisconnect 异常,避免客户端主动断开连接时被当作系统未知错误(500)处理。
- 在读取请求体阶段捕获 ClientDisconnect,返回 499 状态码
- 在 adapter.handle 阶段捕获 ClientDisconnect,记录审计日志并返回 499
- 日志级别从 ERROR 降为 WARNING,减少误报告警噪音
2026-03-09 10:38:14 +08:00
fawney19 4dbfeb87b8 fix: 迁移脚本循环 2026-03-09 03:56:20 +08:00
fawney19 fa69287449 perf: 依赖数据库 CASCADE/SET NULL 替代手动清理关联表,缩短删除事务
- 批量删除移除 cleanup_key_references 手动清理,改为依赖 FK CASCADE/SET NULL
- video_tasks.key_id FK 增加 ondelete="SET NULL",附带幂等迁移脚本
- _sync_delete 增加 statement_timeout 和任务级超时保护
- 批量导入在每次 await 前释放闲置 DB 连接,按批次提交写入避免长事务
- 前端轮询改为先查后等,首次查询不再多等一个间隔
2026-03-09 03:48:20 +08:00
fawney19 654ce89541 fix: 批量删除任务完成前等待所有进度更新的异步回调完成
收集 run_coroutine_threadsafe 返回的 Future,在标记任务完成前
通过 asyncio.gather 等待所有进度更新回调执行完毕,避免任务
状态提前跳到 completed 而进度数据尚未写入 Redis 的竞态问题。
2026-03-09 02:52:38 +08:00
fawney19 fd32597015 perf: 优化批量删除策略,缩小事务粒度并增强容错
- 批量删除分批大小从 500 降至 50,减少单事务锁持有时间
- 单批失败时跳过并继续,不再中断整个删除任务
- 进度上报改为按批次计数(每 5 批或末批),替代时间限频
- 关联表清理简化为直接 DELETE WHERE IN,移除逐行分批删除
2026-03-09 02:35:07 +08:00
fawney19 84cf07b7a2 refactor: 批量删除任务状态存储从内存字典迁移到 Redis
- BatchDeleteTask 改为 BatchDeleteTaskInfo,状态序列化存入 Redis
- submit_batch_delete / get_batch_delete_task 改为 async 函数
- 删除进度通过限频回调写入 Redis,替代直接修改内存属性
- 移除内存任务注册表和过期清理逻辑,依赖 Redis TTL 自动过期
- 路由层对应调整为 await 调用
2026-03-09 02:09:24 +08:00
fawney19 e4476d0bc6 perf: Pool 批量删除改为异步任务模式,避免大批量删除阻塞请求
- 新增 batch_delete_task 模块,提交删除后立即返回 task_id,后台线程分批执行
- 新增查询任务进度的 API 端点,前端轮询展示实时进度
- RequestCandidate 大表清理改为按行数分批删除,防止单条语句超时
2026-03-09 01:35:01 +08:00
fawney19 0379f01ce8 perf: 删除 Key 前显式清理关联表,避免 CASCADE 级联删除超时
在 ProviderAPIKey 删除前,先批量删除 RequestCandidate、GeminiFileMapping、
VideoTask 等关联表记录,替代依赖数据库 CASCADE 级联删除,防止大量关联
记录导致删除操作超时。Pool 批量删除、封禁清理、Endpoint 批量删除三处
统一使用 cleanup_key_references。
2026-03-09 00:07:20 +08:00
fawney19 95e72594ea perf: Pool 账号删除后改用乐观更新替代全量重载
- PoolAccountBatchDialog: 批量删除全部成功时直接从本地列表移除,其余操作保留原重载逻辑
- PoolManagement: 单个删除后本地移除条目并更新 total,当前页为空时自动跳转前一页
2026-03-08 23:46:04 +08:00
fawney19 f9ffb1cae5 chore: bump aether-hub version to 0.1.4 2026-03-08 23:06:06 +08:00
fawney19 91b6e0a382 fix: Hub 连接清理顺序修复、OAuth 类型常量提取、disassociate 逻辑优化
- Hub proxy/worker 连接关闭时先 unregister 再延迟 abort writer,确保缓冲消息排空
- 提取 OAUTH_AUTH_TYPES 常量,替代各处硬编码的 OAuth 类型列表
- auto-disassociate 跳过 OAuth Key,避免其动态 allowed_models 干扰判定
- 删除 Key 时传入 skip_disassociate=True,跳过不必要的解关联检查
- ModelMapper 缓存命中时将 ORM 实例脱离 Session,修复 DetachedInstanceError
2026-03-08 23:03:56 +08:00
fawney19 f5f7a23bb0 fix: 错误响应读取移至连接关闭前,ModelMapper 缓存改为模块级共享
1. chat_handler_base/cli_stream_mixin: 将 _extract_error_text 提前到
   response_ctx.__aexit__ 之前执行,避免连接关闭后无法读取错误响应体
2. ModelMapperMiddleware: 实例级缓存改为模块级共享缓存,消除多实例
   间缓存不一致问题;缓存失效服务改为直接调用静态方法
2026-03-08 22:18:10 +08:00
fawney19 8be9601963 fix: recording.py 中所有 float() 改为 Decimal(),修复与数据库 Numeric 字段相加的类型错误 2026-03-08 21:58:55 +08:00
fawney19 eaad1579e6 fix: dashboard 统计字段补充 float() 转换,修复 Decimal 与 float 混合运算的类型错误 2026-03-08 21:54:20 +08:00
fawney19 1d8bf56efd fix: defer() 改为链式调用,修复 SQLAlchemy 多参数报错 2026-03-08 21:40:36 +08:00
fawney19 73db997e92 fix: 统计聚合中成本字段 float 改 Decimal,修复与数据库 Numeric
字段相加的类型错误
2026-03-08 21:33:05 +08:00
fawney19 2c5654d694 fix: 单个迁移文件自行提交 2026-03-08 18:53:36 +08:00
fawney19 f490c3a5cd fix: 优化迁移脚本效率 2026-03-08 18:42:27 +08:00
fawney19 2d9158b321 fix: 把同一张表的多个列放在一条 ALTER TABLE 里 2026-03-08 17:56:03 +08:00
fawney19 48d13762d9 refactor(cost,perf): 成本字段 Float 改 Numeric 并优化多处查询性能
- 数据库所有 cost/price 字段从 Float 改为 Numeric(20,8),解决浮点精度问题
- API 响应中 Decimal 值统一用 float() 转换确保 JSON 序列化
- 路由中手动 commit 后标记 tx_committed_by_route 防止中间件重复提交
- 多处查询优化:SQL 聚合替代 Python 遍历、load_only/defer 减少字段加载、
  批量 DELETE 替代逐条 ORM 删除、N+1 查询消除、UNION ALL 合并多表日期查询
- 新增 provider_api_keys (provider_id, is_active) 复合索引
- 候选构建热路径 defer 冷字段,钱包扣费合并解析与加锁查询
2026-03-08 16:44:16 +08:00
fawney19 bd3f73c2fc feat(retention): 删除用户/Key 时保留历史记录,外键改 SET NULL 并添加名称快照
- Usage/RequestCandidate/VideoTask/Stats 等表的 user_id/api_key_id 外键从
  CASCADE 改为 SET NULL,删除用户或 Key 后历史记录不再丢失
- 各表添加 username/api_key_name 快照字段,删除后仍可追溯归属
- 新增 bulk_cleanup 模块,分批置空大表外键避免长事务锁
- 删除用户/Key 流程集成预清理步骤,先置空再删除
- 精简 candidate_builder 冗余 debug 日志
- 修复 proxy_nodes 启动日志 format 占位符错误({} -> %s)
- 前端批量操作请求增加 5 分钟超时配置
2026-03-08 14:31:15 +08:00
fawney19 25c33846be perf(pool): 添加性能计时日志,优化批量删除分批与模型解除关联查询
- 前端批量操作对话框添加 loadAllKeys/executeAction 计时日志
- 后端池账号列表接口和批量删除接口添加分阶段耗时日志
- 批量删除按数据库类型自动选择分批大小,统一前端 batch size 为 2000
- 优化 auto_disassociate 查询:先检查 unlimited key 提前返回,使用 load_only 减少字段加载
- 新增批量操作路由和自动解除关联的单元测试
2026-03-08 03:58:10 +08:00
fawney19 124c4ca403 fix(pool): 优化批量删除性能,使用 SQL 批量删除替代逐条 ORM 删除
- 后端批量删除改用 sa_delete 直接执行 SQL,避免逐条加载和删除
- 前端删除操作批次大小从 2000 减小到 50,防止大批量删除超时
- 增加前端批量操作失败时的错误日志输出
2026-03-08 03:08:27 +08:00
fawney19 ef0f8dd4d0 feat(pool,provider_ops): 扩展批量操作支持代理设置,优化签到缓存与余额刷新
- pool batch-action 新增 clear_proxy/set_proxy 操作,前端统一使用
  batch-action API 替代逐个调用,批量上限从 100 提升至 2000
- batch-action 删除操作后执行 key 删除副作用(run_delete_key_side_effects)
- BalanceAction 签到增加 6 小时缓存冷却,同一 host 避免重复签到
- 异步余额刷新增加 per-provider 防重入保护
2026-03-08 02:29:51 +08:00
fawney19 d0eca509d4 fix(http): 支持默认HTTP客户端原子重建以恢复HTTP/2流容量
feat(db): 为外键列补充缺失的数据库索引并添加迁移脚本

- HTTPClientPool._reset_default_client 原子替换共享客户端,旧客户端在宽限期后异步关闭
- reset_upstream_client 对无代理路由不再直接跳过,改为调用 _reset_default_client
- 为 Usage, WalletTransaction, PaymentCallback, RefundRequest, VideoTask, RequestCandidate 等表的外键列添加 index=True
2026-03-08 01:27:50 +08:00
fawney19 90663793a2 Merge branch 'feat/wallet-system' into master
feat(wallet): 钱包系统替代配额系统,新增支付与退款机制

Closes #204
2026-03-08 00:06:05 +08:00
LewisPen 783f654953 feat(wallet): 钱包系统替代配额系统,新增支付与退款机制
- 新增钱包余额管理、充值、扣费、退款完整流程
- 新增支付网关抽象层(支持手动/支付宝/微信)
- 用量计费从配额系统迁移到钱包余额扣费
- 新增管理员钱包管理与支付订单管理页面
- 新增用户钱包中心页面
- 移除独立 Key 锁定机制,统一由钱包余额控制
- 新增相关 API 路由、序列化器与数据库迁移
- 新增钱包、支付、退款相关测试
2026-03-08 00:05:48 +08:00
fawney19 9cdcce1b5f fix(providers): 允许 max_probe_interval_minutes 设为 0 并修复零值被覆盖的问题
- 将 max_probe_interval_minutes 校验范围从 2-32 改为 0-32
- 修复 routes.py 中 max_probe_interval_minutes 和 cache_ttl_minutes
  使用 `or` 短路导致零值被默认值覆盖的问题,改用 is not None 判断
- 前端表单同步调整最小值约束和提示文案
- 新增零值校验的单元测试
2026-03-07 18:43:28 +08:00
fawney19 06b483f79d fix(task): 修复 SyncTaskExecutionService 调用 dispatch 缺少 user_id 参数及返回值解包数量不匹配 2026-03-07 17:32:17 +08:00
fawney19 a00e137ffc fix(providers): 修复 Key 表单状态同步与 streaming 候选状态码处理
前端:
- KeyAllowedModelsEditDialog 同时监听 open 和 apiKey 变化,避免切换 key 时状态未刷新
- KeyFormDialog 新增 api_formats 过滤与默认值逻辑,可用格式变化时自动同步表单
- ProviderDetailDrawer 合并 provider 和 endpoint 的 api_formats 传递给 Key 表单,
  数据刷新后同步 currentEndpoint 和 editingKey 引用

后端:
- mark_candidate_streaming 移除 status_code 参数,streaming 阶段不再提前写入状态码
- 简化 active_requests 中 streaming 请求的完成判断,不再依赖 status_code 条件
2026-03-07 17:11:19 +08:00
fawney19andAAEE86 239238fe47 refactor(task): 拆分 TaskService 并重构任务生命周期
- 将 task 公共协议/上下文/异常/schema 下沉到 core,并迁移 polling 目录
- 新增 execute/submit/video 子模块,拆分同步执行、异步提交流程、错误处理与视频任务操作
- 收敛 TaskService 为门面编排,内部委派到 SyncTaskExecutionService、AsyncTaskSubmitService、VideoTaskOperationsService
- 重构 main 生命周期管理:引入 LifecycleState,拆分启动与关闭流程
- 删除未使用的 TaskExecuteFacadeService 与 TaskSubmitFacadeService

Closes #201

Co-authored-by: AAEE86 <[email protected]>
2026-03-07 15:33:29 +08:00
github-actions[bot] 4cd6e0d10f chore(proxy): update download links for proxy-v0.2.4 2026-03-06 20:06:59 +00:00
fawney19 9b29a65c68 chore: bump aether-hub version to 0.1.3 2026-03-07 04:02:02 +08:00
fawney19 df4a49a9fb chore: bump aether-proxy version to 0.2.4 2026-03-07 04:00:27 +08:00
fawney19 bb268310e2 feat(keys): 新增 Provider Keys 批量删除 API
- 后端新增 POST /keys/batch-delete 接口,支持一次删除最多 100 个 Key
- 按 provider_id 聚合执行副作用,避免逐个删除导致的重复 Redis 操作
- 前端批量操作对话框中删除操作改用批量 API,按 100 个一批分批调用
2026-03-07 03:49:44 +08:00
fawney19 7b0908cd87 fix(task): TaskPoller Redis 分布式锁操作增加异常捕获,避免 Redis 故障时任务轮询崩溃 2026-03-07 03:14:48 +08:00
fawney19 95bc742057 fix(gunicorn): worker timeout 从 120s 放宽到 300s,收敛到配置文件统一管理
- 移除 Dockerfile 命令行中硬编码的 --timeout 120,由 gunicorn_conf.py 统一管理
- timeout 默认值从 120 调整为 300,减少异步 worker 偶发心跳延迟导致的误杀
- timeout 和 graceful_timeout 均支持环境变量覆盖(GUNICORN_TIMEOUT / GUNICORN_GRACEFUL_TIMEOUT)
2026-03-07 03:13:44 +08:00
fawney19 e40c890a8e fix(usage): 超时请求清理增强,支持恢复已成功的 streaming 请求
重构 pending/streaming 请求清理逻辑:
- 提取 _find_completed_request_ids 和 _sync_candidate_status_to_success 公用方法
- 清理超时请求时检查 RequestCandidate,已成功的恢复为 completed 而非标记 failed
- 同步更新 candidate 状态,保持 Usage 与 RequestCandidate 一致
- 启动时主动执行一次 pending 清理
2026-03-07 03:04:09 +08:00
fawney19 357c4fd61f fix(build): 构建时 GitHub API 请求支持可选 GITHUB_TOKEN 避免限流
未认证 GitHub API 限流 60 次/小时/IP,频繁本地构建容易触发。
添加可选的 GITHUB_TOKEN 支持(认证后 5000 次/小时),不传 token 时行为不变。
2026-03-07 03:03:55 +08:00
fawney19 d28fea80df fix(trace): 请求追踪详情默认展示全部候选记录 2026-03-07 02:30:48 +08:00
fawney19 fb1aeb789a feat(test,quota,failover): 模型并发测试、统一配额读取器与故障转移取消支持
- 新增 QuotaReader 抽象层,统一 Codex/Kiro/Antigravity 配额解析逻辑,
  替换 pool/routes.py 中分散的配额构建函数
- 模型测试支持并发执行多候选,前端新增 useModelTest composable 统一
  ModelsTab 和 ModelMappingTab 的测试逻辑
- ModelTestDialog 增加结果概览摘要、超长结果折叠、端点列和新状态支持,
  删除已合并的 TestResultDialog
- FailoverEngine 新增客户端断开检测,支持取消剩余候选并标记记录
- 刷新配额改为分批执行,直连测试候选按可用性排序
- 修复 error 判断从 "error" in dict 改为 dict.get("error") 避免误判
2026-03-07 02:16:03 +08:00
fawney19 1f3693d3a2 fix(migration): body_rules 回填限定 codex 提供商类型
原迁移仅按 api_format 过滤,会误回填所有 openai:cli 端点;
现 JOIN providers 表增加 provider_type = 'codex' 条件。
2026-03-06 23:44:37 +08:00
fawney19 90760da499 feat(test,export,headers): 模型测试复用统一运行时、实时进度展示、导出增强与请求头大小写保留
- 模型测试 failover 从手动 FailoverEngine 改为 TaskService.execute_sync_candidates 统一运行时
- 前端新增实时 trace 轮询进度展示(候选状态、测试账号、进度条)
- 用户导出/导入支持明文 Key 优先(版本升至 1.2),新增 email_verified 字段
- SENSITIVE_CREDENTIAL_FIELDS 统一到 provider_ops/types.py,补充 refresh_token
- 请求头大小写保留机制(resolve_header_name_case + HeaderBuilder.add 语义修改)
- Codex envelope 移除合成头部,保留客户端原始请求头
- endpoint_checker 支持自定义超时透传
- 新增 x-forwarded-scheme 到上游丢弃头部列表
2026-03-06 21:06:45 +08:00
fawney19 7950ba7dc5 fix(usage): 请求详情默认选中轻量 tab,避免自动加载大 body
首次打开请求详情时优先选择 request-headers / response-headers / metadata
等轻量 tab,而非直接落到 body tab 触发按需加载。
2026-03-06 18:19:57 +08:00
fawney19 a7088ee538 fix: 修复 http缺少cls 2026-03-06 18:01:28 +08:00
fawney19 050cba9563 feat(proxy,failover,transport): Hyper 上游客户端精细计时、连续失败退避与连接泄漏修复
Proxy:
- 将上游 HTTP 客户端从 reqwest 替换为 hyper,新增 InstrumentedConnector
  实现 TCP 连接/TLS 握手级别的独立计时,上报 connection_reused 等指标
- 前端展示细粒度代理计时(连接复用、等待响应头等)

Failover:
- 引入连续失败退避机制,每 10 次失败递增退避间隔
- 检测 H2 max outbound streams 错误并触发上游客户端重建
- 新增 HTTPClientPool.reset_upstream_client 支持按需重建缓存客户端

连接泄漏修复:
- Handler 异常路径确保 response_ctx 被正确关闭
- HubResponseStream 迭代结束后在 finally 块中清理 stream_id
- HubTunnelTransport.handle_request 捕获所有异常并清理流状态
2026-03-06 17:55:14 +08:00
fawney19 2269617a9f fix(nginx): 剥离 Cloudflare 请求头,防止泄露给上游 AI 提供商
在反向代理配置中将 CF-Connecting-IP、CF-IPCountry、CF-Ray、
CF-Visitor、CDN-Loop、True-Client-IP、CF-Worker、CF-EW-Via
等头置空,避免用户真实 IP 及 CF 元数据被透传到上游。
2026-03-06 15:54:34 +08:00
fawney19 bdccfa6e78 feat(usage,pool,codex): 请求详情 body 按需加载、配额选择器重构与 Codex body rules 修正
- usage 详情 API 新增 include_bodies 参数,支持跳过 body 内容返回 has_*_body 标记
- 前端请求详情抽屉首次加载不含 body,切换到 body tab 时延迟加载并展示 skeleton
- Timeline 组件延迟 120ms 挂载,避免阻塞抽屉渲染
- 提取号池配额判断逻辑到 quota-selectors 工具模块并添加单测
- Codex 移除 openai:compact 的默认 body rules 注册
- ModelTestDialog/TestResultDialog 模板格式化
2026-03-06 15:40:11 +08:00
fawney19 d97ec3fde2 feat(admin,pool,billing): 端点级模型测试、Provider 自动置顶、缓存 TTL 分级计费展示与账号状态增强
- 模型测试支持指定端点:新增 ModelTestDialog 组件,多端点时弹窗选择,单端点直接测试;
  后端 test-model-failover 接口新增 endpoint_id 参数,支持 global/direct 模式下按端点过滤候选
- 创建 Provider 时优先级自动置顶(provider_priority 默认 None,后端取 min-1),
  显式指定优先级时 shift 已有行;前端创建时不发送 priority,更新时保留
- 缓存计费 UI 增强:RequestDetailDrawer 支持 5min/1h 缓存创建 token 分级展示,
  含按 TTL 匹配单价和分行成本计算;ModelDetailDrawer/ModelsTab 标签区分 5min/1h 缓存创建
- Pool 批量操作额度筛选拆分为「无5H限额」和「无周限额」,按 | 分隔 segment 匹配
- KeyFormDialog 优化非 vertex_ai 时布局,API 密钥输入内联到 grid 右列
- Codex refresher 结构化错误标记:401/402/403 使用 [OAUTH_EXPIRED]/[ACCOUNT_BLOCK] 前缀,
  新增 deactivated_workspace 识别与分类
- 前后端 accountBlock 关键词同步:新增 token invalidated、deactivated_workspace 识别,
  OAuth 失效提示清理 block 前缀后展示
- PoolConfig 新增 batch_concurrency 配置(默认 8,上限 32)
- 预设模型新增 gpt-5.4;TestResultDialog 响应式布局与 key 脱敏优化
2026-03-06 13:13:01 +08:00
github-actions[bot] d17472f09e chore(proxy): update download links for proxy-v0.2.3 2026-03-05 18:01:47 +00:00
fawney19 f8b7cd2925 chore: bump aether-proxy version to 0.2.3 2026-03-06 01:48:50 +08:00
fawney19 e9c3ac94c6 feat(proxy,oauth,pool): H2 头过滤、OAuth 过期分级标记与批量操作进度条
- proxy: 屏蔽 host/content-length 头转发,避免 H2 PROTOCOL_ERROR
- oauth: 区分 [REFRESH_FAILED] 与 [OAUTH_EXPIRED] 标记,token 过期
  自动阻止调度但不停用账号,便于管理员恢复
- pool/account_state: 识别新增的 OAUTH_EXPIRED/REFRESH_FAILED 前缀
- 前端: 批量操作显示实时进度条,倍率编辑 Escape/blur 竞态修复,
  OAuth 刷新失败后自动刷新列表
2026-03-06 01:47:09 +08:00
fawney19 fa71cddb60 feat(proxy,billing,pool): 增强隧道/流中断诊断日志、修复缓存 TTL 差异化计价
- stream_processor: 上游流中断时记录完整异常链与已传输 token 统计
- hub_transport: 断连影响 in-flight 流、STREAM_ERROR、超时场景补充 warning 日志;
  未启用时跳过重连循环,重连失败日志降频避免刷屏
- tunnel_manager: 流超时/错误/全部取消/STREAM_ERROR 增加诊断日志与字节统计
- billing_integration: 计费时自动补全 cache_ttl_minutes(从 provider key 查询
  或从 5m/1h 细分 token 回推),修复缓存 TTL 差异化计价缺失
- PoolManagement.vue: 移除重复的账号告警 Badge
- 新增 billing integration 单元测试
2026-03-06 00:14:10 +08:00
fawney19 228cbc8f87 feat(pool,admin): 拆分调度维度、重构调度 UI 与增强配置导入导出
调度维度:
- 新增 cache_affinity/free_first/team_first/plus_first/load_balance 五个独立维度
- 将 free_team_first 标记为 hidden,保留后向兼容但不再显示
- Registry 新增 hidden 属性,_helpers 新增 plus_only 优先级评分

前端调度对话框:
- 分配模式(互斥组)独立为按钮组选择,策略调度保留拖拽排序
- 默认调度从 LRU 轮转改为缓存亲和
- 提取 buildPresetListItem/insertMissingByPreferredOrder 消除重复代码

配置导入导出:
- 导出时 api_formats 支持规范化、去重与 None 回退到 Provider 端点
- 导入时兼容 supported_endpoints 别名与历史 None 语义
- 新增 test_admin_system_key_formats 单元测试
2026-03-05 23:22:25 +08:00
fawney19 da915208a8 feat(provider,adapter): Codex 默认 body_rules 按 provider_type 维度注册,adapter 全链路传递 provider_type
- 新增 register_provider_default_body_rules 注册机制,将 Codex 特有的
  body_rules 从 EndpointDefinition 全局默认移至 codex plugin 按
  (provider_type, endpoint_sig) 维度注册
- handler adapter 的 build_endpoint_url/build_request_body/get_cli_extra_headers
  增加 provider_type 参数,Codex 判断优先使用 provider_type 而非 URL 匹配
- 前端 getDefaultBodyRules API 支持 provider_type 参数,缓存 key 区分不同
  provider 类型;Codex 路径判断同样优先使用 provider_type
- ProviderDetailDrawer 将 mapping-preview 拆为独立加载,不阻塞首屏渲染
- PoolManagement 补全 KeyFormDialog 缺失的 endpoint/available-api-formats props
- 固定类型 Provider 创建时自动填充 provider-scoped 默认 body_rules
2026-03-05 19:04:21 +08:00
fawney19 694167f78f feat(pool,trace): 账号封禁原因细分与请求追踪 attempted_only 过滤
- 将账号封禁原因从笼统的"账号异常"细分为封禁/停用/需要验证三类,
  前后端关键词组同步拆分,号池管理页面展示对应分类标签
- trace API 新增 attempted_only 参数,支持仅返回实际尝试过的候选,
  前端时间线组件默认启用过滤,排除 available/unused/skipped 记录
2026-03-05 17:32:21 +08:00
fawney19 a7697032a4 feat(pool): 账号停用检测、OAuth refresh 失效标记与号池管理性能优化
- 新增 account_deactivated 关键词检测,覆盖 error_handler / health_policy / account_state / 前端
- 401 健康策略分级冷却:账号永久停用 1h,临时认证失败 60s
- OAuth refresh token 失败时标记 oauth_invalid(不停用 key),成功时清除非账号级标记
- 号池管理 API 使用 load_only + SQL 聚合替代全量拉取,减少查询开销
- Redis 冷却统计改用 batch_count_provider_cooldowns (SCAN) 替代逐 key 检查
- 前端时间线移除 executableTimeline 过滤层,Provider 选择改为非阻塞加载
2026-03-05 16:23:58 +08:00
fawney19 e86c8edd4b feat(provider): 新增 keep_priority_on_conversion 字段,支持格式转换时保持调度优先级 2026-03-05 15:28:03 +08:00
fawney19 b1be413dc0 feat(pool): 号池额度主动探测、封禁自动清除、调度硬优先级与前端重构
- 新增 PoolQuotaProbeScheduler,按 probing_interval_minutes 主动探测静默 Key 额度
- pool_advanced 增加 probing_enabled / auto_remove_banned_keys 配置项
- error_handler 和 quota_service 支持封禁 Key 自动删除及缓存清理
- multi_score 策略从加权混合重构为硬优先级排序,引入 mutex_group 互斥组
- 指纹注入从 handler 层下移至 ClaudeCode envelope 层
- OAuth 批量导入支持 concurrency 并发参数
- 前端号池管理拆分高级设置/账号批量/代理设置为独立组件
- 号池总览接口精简,仅返回已启用调度的 Provider
2026-03-05 15:15:26 +08:00
fawney19 fdb50a065b refactor(fingerprint): 移除请求头构建路径中的 per-key 指纹定制化注入
headers.py 删除 build_browser_fingerprint_headers / build_anthropic_extra_headers
函数及相关辅助,改为直接使用固定常量;PassthroughRequestBuilder 移除
Claude 格式的指纹注入步骤;Antigravity constants 移除 per-key 指纹
覆盖逻辑,统一使用进程级固定值。
2026-03-05 10:24:52 +08:00
fawney19 1ac59d4894 feat(pool): 新增调度维度、互斥组机制、健康策略扩展与配额刷新增强
- 新增 priority_first/health_first/latency_first/cost_first 四个调度维度
- 引入 mutex_group 互斥组机制,lru 与 single_account 归入 distribution_mode
- 维度 compute_metric 签名扩展 context 参数,支持获取 cost_totals 等上下文
- 各维度增加 evidence_hint 字段描述评分依据
- 健康策略扩展 408/409/423/425/5xx 瞬态状态码冷却,403 按 body 分级冷却
- Codex 配额刷新增强 401/402/403 错误处理,402 生成 fallback 元数据
- 前端号池管理支持账号优先级内联编辑与互斥维度切换 UI
- 列表排序改为 internal_priority + created_at,移除 sticky_counts 查询
2026-03-05 09:53:23 +08:00
fawney19 32ccf61baa feat(fingerprint): 引入 per-key 请求指纹系统,替代全局 TLS 指纹开关
为每个 ProviderAPIKey 生成并持久化独立的请求指纹配置,涵盖 TLS impersonate
profile、浏览器 UA、Stainless SDK 头部、Node/Chrome/Electron 版本等维度。
指纹基于 key ID 确定性生成,支持手动编辑和批量重新生成。

- 新增 fingerprint 模块:生成、加载、校验、懒持久化
- 数据库迁移:provider_api_keys 新增 fingerprint JSON 列
- 请求链路注入:handler 基类设置上下文指纹,request_builder 和 envelope 消费
- HTTP Client 支持动态 impersonate profile 选择
- Antigravity 适配器使用指纹覆盖 UA/session/Node 版本
- 前端移除手动 TLS 指纹开关,新增批量 regenerate_fingerprint 操作
- 号池管理 UI 优化:token 缩写格式、blocked 行样式、时间显示改为日期格式
2026-03-05 02:07:34 +08:00
fawney19 1d04c41ae7 refactor(pool): 移除调度多维评分详情,简化为 reasons 摘要
移除 scheduling_score、candidate_eligible、scheduling_dimensions 等字段,
前后端统一使用 scheduling_reasons 展示调度状态,精简号池列表信息密度。
2026-03-04 22:51:02 +08:00
fawney19 b2dcf82ca8 feat(pool): 引入多维评分调度策略与账号状态检测
- 新增 multi_score 调度模式,支持 LRU/延迟/健康度/剩余额度多维加权评分
- 新增调度预设维度系统(free_team_first, quota_balanced, recent_refresh, single_account),支持有序对象列表配置格式并兼容旧字符串列表
- 新增 account_state 模块,统一账号封禁/受限检测逻辑,替代分散在 routes 中的判断代码
- 新增 health_cache 模块和 latency 采样(redis_ops.record_latency / batch_get_latency_avgs)
- RequestDispatcher 返回 ttfb_ms,PoolManager.on_request_success 记录延迟样本
- 前端:PoolConfigDialog 替换为 PoolSchedulingDialog,支持预设维度可视化配置;号池管理页增加调度模式标签与账号异常 Badge 显示
- 提取前端 accountBlock 工具函数,ProviderDetailDrawer 复用统一判断
- scheduling_dimensions 增加 account_state 和 latency 维度评估
- 补充 account_state、health_cache、multi_score 策略、preset 维度、redis latency 等测试
2026-03-04 22:06:19 +08:00
fawney19 57b86034cf feat(pool,priority): 号池聚合显示、API 格式归一化与优先级管理重构
- 优先级管理对话框按 family 分组显示 API 格式,号池 key 聚合为单条目展示
- 拖拽排序改用 key ID 替代数组索引,号池聚合项禁用拖拽/编辑/开关操作
- 后端 key 分组查询增加 API 格式键归一化,返回 provider_id
- 提取 OAuth auth_config 解密逻辑,新增 _derive_oauth_expires_at 从加密配置派生过期时间
- 号池管理移除会话列,调整 OAuth 过期信息与刷新按钮的布局顺序
2026-03-04 12:59:27 +08:00
fawney19 095e312ab3 fix(usage): 修正请求时间线 unused 候选过滤与分组排序逻辑
- 号池内 unused key 不再展示,非号池 unused 仅保留 retry_index=0
- buildProviderGroups 使用 candidate_index 替代数组下标作为分组索引
- 号池组与供应商组合并后统一按 startIndex 排序
2026-03-04 10:26:19 +08:00
fawney19 82a9fb3c39 feat(codex): 增强 OAuth 导入解析与账号信息提取,优化维护调度器线程模型
Codex OAuth:
- 导入解析支持附加账号字段(account_id/plan_type/user_id/email)
- enrich_codex 扩展从 access_token 和直接字段提取账号信息
- parse_codex_id_token 支持 JWT/JSON 字符串/dict 三种输入格式
- request patching 新增 openai:compact 格式支持
- codex_usage_parser 新增 credits_unlimited 字段解析

维护调度器:
- 同步 DB 操作迁移到线程池执行,避免阻塞事件循环
- 新增 request_candidates 定期清理任务
- 新增每周 VACUUM ANALYZE 数据库表维护任务
- 新增 enable_db_maintenance 配置项

前端:
- ElapsedTimeText 从 setInterval 改为 requestAnimationFrame
- 时间线过滤 available/unused 占位记录
- Usage 页面默认关闭全局自动刷新
- 移除号池管理中的配额更新时间显示
2026-03-04 10:00:38 +08:00
fawney19 e181329a81 fix(usage): 抽取 ElapsedTimeText 组件、缩短缓存 TTL 提升实时性,修复 _KeyCandidate slots
- 将 UsageRecordsTable 中的内联计时逻辑抽取为独立的 ElapsedTimeText 组件
- usage records 缓存 TTL 从 15s 降为 3s,全局自动刷新间隔从 5s 改为 3s 并默认开启
- 补充 PoolManager._KeyCandidate 缺失的 __slots__ 字段
2026-03-04 02:24:46 +08:00
fawney19 cdce817928 fix(announcement): 将分页计数移至排序之前执行 2026-03-03 22:35:08 +08:00
fawney19 97b0146ce9 perf: 全栈查询优化、前端缓存去重与页面可见性优化
后端:
- SQL count 查询统一改用 func.count() 子查询替代 query.count()
- Dashboard/Audit 等页面多次独立查询合并为单次聚合查询
- Provider summary 列表改为批量查询消除 N+1 问题
- DailyStats 逐天循环查询改为 CASE 分桶单次查询
- 使用 load_only() 减少不必要的列加载
- cache_decorator 支持嵌套属性路径解析(dotted vary_by)
- 多个管理/公共端点新增 @cache_result 缓存装饰

前端:
- cache.ts 新增 in-flight 请求复用、dedupedRequest、buildCacheKey
- 大量 API 调用添加前端缓存或去重
- 多个页面定时器在标签页隐藏时暂停、可见时恢复
- Auth 检查从 setInterval 改为 storage + visibilitychange 事件驱动
- 请求竞态防护(requestId 模式)

数据库:
- Usage 表新增 idx_usage_status_user_created 复合索引
2026-03-03 22:04:40 +08:00
fawney19 0a60492146 fix(proxy_nodes): 更新代理节点模块描述 2026-03-03 17:28:56 +08:00
fawney19 4ea187cfac feat(pool): 号池候选重构为 PoolCandidate 单候选模式与池内 key 故障转移
- 新增 PoolCandidate 子类,排序阶段作为单候选参与,执行阶段在 pool_keys 内部选择/切换 key
- FailoverEngine 新增 _execute_pool_candidate 方法,支持池内 key 级别故障转移与重试
- 提取 _execute_attempt / _attach_attempt_context / _classify_attempt_error 公共方法
- CandidateBuilder 对号池 Provider 构建单个 PoolCandidate(包含所有可用 key)
- CandidateSorter 支持 PoolCandidate 独立优先级分组(global_priority / pool_priority)
- CandidateResolver PRE_EXPAND 模式按 pool_keys 展开预创建记录,附加 pool_group_id
- TaskService._apply_pool_reorder 改为对 PoolCandidate 调用 select_pool_keys
- PoolManager 新增 select_pool_keys 方法,复用 reorder_candidates 逻辑
- 新增 global_priority 号池配置字段(前后端同步)
- 前端 Timeline 支持按 pool_group_id 分组显示多号池尝试
- 异步提交路径新增 _expand_pool_candidates_for_async_submit 展开逻辑
2026-03-03 17:24:22 +08:00
fawney19 dcba7c62a2 docs: 更新 README 部署说明与默认配置
- 调整 Docker Compose 部署方式的标题描述
- 恢复预构建镜像部署的正常命令流程
- 移除 deploy.sh 的 --hub-tag 选项说明
- Aether Proxy 标注为可选组件
- GUNICORN_WORKERS 默认值调整为 2
2026-03-03 16:38:22 +08:00
fawney19 dba99455a7 fix: 移除redis持久化功能 2026-03-03 12:44:48 +08:00
fawney19 7ebce161e8 feat(pool,ui,trace): 号池按页配额刷新、配额倒计时、Trace 全量候选与 UI 用语统一
- 号池管理支持按当前页 Key 刷新配额,后端 refresh-quota 接口支持 key_ids 参数筛选
- 配额进度条 tooltip 展示重置倒计时,前端解析后端重置时间并实时倒计时
- Provider 选择器在无号池提供商时禁用,切换/刷新后保持选中状态对齐
- 启停账号后立即更新调度标签并刷新列表
- Trace 监控展示全量候选记录,不再过滤 available/unused 状态
- 请求时间线号池节点与 Provider 节点去重
- 全局 UI 用语统一:已停用/已禁用 -> 停用/禁用
- 导航菜单调整模型管理与号池管理顺序
2026-03-03 11:32:41 +08:00
fawney19 022aec5720 fix: 修复迁移版本号重复问题 2026-03-03 09:28:33 +08:00
fawney19andAAEE86 11997c024e feat(pool,scheduling): 号池调度维度、配额冷却机制与管理后台重构
- 新增 scheduling_dimensions 模块,为每个 Key 计算多维调度状态(手动/冷却/熔断/成本/健康)
- 新增 quota_cooldown 模块,统一判定 Key 的有效冷却原因
- Pool 管理后台 API 扩展 Key 详情字段(调度状态/维度/配额/OAuth 信息)
- 前端 Pool 管理页面重写,支持调度状态展示、批量清理封禁 Key
- Handler 基类增加请求调度元数据采集,stream telemetry 增强
- 请求时间线组件增强,支持 attempted 候选展示
- Kiro OAuth 凭证导入解析改进
- 新增 usage 表 provider_key 索引迁移
- 补充调度维度、配额冷却、候选枚举等单元测试

Closes #197

Co-authored-by: AAEE86 <[email protected]>
2026-03-03 09:22:20 +08:00
fawney19 f787b1b02a fix(codex): 取消对 openai:compact 端点的 body_rules 修改
Codex 的请求体规则(drop max_output_tokens/temperature/top_p, set store=false 等)
只应作用于 openai:cli, 不应影响 openai:compact 端点。
2026-03-03 03:39:14 +08:00
fawney19 a03368a3fe fix(usage): metadata tab 也显示展开/收缩和复制按钮 2026-03-02 22:44:49 +08:00
fawney19 e26ed8481f fix(scheduling): 负载均衡模式与无亲和键场景统一走随机排序
重构 CandidateSorter.shuffle_keys_by_internal_priority 中同优先级
Key 的排序逻辑,将三分支简化为两分支:
- 随机排序:TTL=0 / 负载均衡模式 / 无 affinity_key
- 哈希确定性排序:缓存亲和模式且有 affinity_key

移除了原先"无 affinity_key 时按 ID 排序"的冗余分支,新增
对应单元测试覆盖三种场景。
2026-03-02 22:36:24 +08:00
fawney19 8e98eed5c8 refactor(failover): 用 provider failover_rules 替代硬编码 ErrorClassifier 判断
移除 submit_with_failover 中基于 ErrorClassifier 的客户端错误硬编码逻辑,
改为读取 provider.config.failover_rules 进行规则匹配:
- error_stop_patterns: 错误响应命中时终止 failover
- success_failover_patterns: 2xx 响应命中时继续尝试下一个候选
同步更新相关注释、异常描述及测试用例
2026-03-02 22:21:13 +08:00
fawney19 0bff15f964 feat(endpoint): 端点默认 body_rules 机制与 Codex 规则回填
- EndpointDefinition 新增 default_body_rules 字段,openai:cli/compact 配置 Codex 默认规则
- 创建端点时若未指定 body_rules 则自动填充对应格式的默认值
- 新增 GET /defaults/{api_format}/body-rules 接口查询默认规则
- 前端 EndpointFormDialog 增加"重置请求体"按钮,支持一键恢复默认
- Alembic 迁移回填已有 Codex 端点的默认 body_rules
- 新增 metadata 和 endpoint 创建默认值的单元测试
2026-03-02 22:08:39 +08:00
fawney19 3384c6d666 fix(alembic,pipeline,resilience): 数据库迁移与异常处理健壮性修复
- alembic env: 改用事务级 advisory lock (pg_advisory_xact_lock),事务结束自动释放
- 迁移脚本: 使用原生 SQL IF NOT EXISTS/IF EXISTS 替代运行时列检查
- pipeline: SQLAlchemy 异常后先回滚事务再写审计,防止 aborted 状态二次报错
- resilience: ProgrammingError 标记为不可恢复,缩减 DB 重试异常范围
- 前端: 端点规则支持拖拽排序
2026-03-02 18:02:21 +08:00
fawney19 5f1c74aca0 fix(cli): 统一 CLI handler pending usage 的 api_format 来源
CLI stream/sync mixin 中创建 pending usage 记录时,api_format 取值
来源不一致(部分用 FORMAT_ID,部分用 allowed_api_formats[0])。
新增 primary_api_format 属性并统一使用,确保 pending 记录的格式
与实际客户端请求格式一致。
2026-03-02 15:13:28 +08:00
fawney19 310355bcc1 fix(ci): 修复多架构Docker镜像构建,arm64覆盖amd64问题
分开push到同一tag会导致后者覆盖前者manifest。
改为先按digest分别push,再用imagetools合并为多架构manifest。
2026-03-02 13:59:25 +08:00
fawney19 ab07d83aaf fix(ci): 使用认证的gh CLI替代未认证curl调用GitHub API
download-hub步骤使用未认证的curl获取release列表,
频繁触发GitHub API速率限制导致构建失败。
改用gh CLI自带GITHUB_TOKEN认证避免此问题。
2026-03-02 13:41:05 +08:00
github-actions[bot] 11cdf52f7b chore(proxy): update download links for proxy-v0.2.2 2026-03-02 05:37:31 +00:00
fawney19 8a2c3596ce chore: bump aether-hub version to 0.1.2 2026-03-02 13:29:27 +08:00
fawney19 a61d16d120 chore: bump aether-proxy version to 0.2.2 2026-03-02 13:29:15 +08:00
fawney19 01df063cc1 feat(ci,alembic): Hub Docker镜像构建发布,数据库迁移并发安全加固
- build-hub.yml 新增 Docker job,构建多架构镜像推送至 GHCR 和 Docker Hub
- build-hub/build-proxy Release 名称简化为 tag 名
- alembic/env.py 使用 PostgreSQL advisory lock 防止多进程并发迁移竞态
- 迁移脚本改用 ADD/DROP COLUMN IF NOT EXISTS 替代 inspector 检查
2026-03-02 13:24:58 +08:00
fawney19 68bae686da feat(proxy): 支持远程推送升级与proxy元数据上报
- aether-proxy 注册和心跳时上报 proxy_metadata(含版本号)
- 心跳 ACK 支持 upgrade_to 字段,proxy 收到后自动执行升级
- 重构 upgrade 逻辑,新增 perform_upgrade 用于远程触发的自动升级
- stream_handler 延迟统计改为仅记录连接建立延迟(DNS+TCP/TLS+TTFB)
- 后端新增 proxy_metadata 数据库字段和批量升级 API
- 远程配置支持下发 upgrade_to 版本指令
- 前端展示节点版本号,支持单节点和批量升级操作
2026-03-02 12:58:41 +08:00
fawney19 f978888759 feat(heartbeat): 心跳可靠性增强,原子计数与去重优化
- Rust proxy: 引入 snapshot+ACK 确认机制,心跳未确认时保留快照重发,
  避免指标丢失;添加 heartbeat_session_id 防跨进程去重误判
- Hub transport: Redis SETNX 心跳去重,避免多 worker 重复写库;
  ACK 回显 heartbeat_id 供 Rust 端匹配
- ProxyNodeService.heartbeat: 改用 SQLAlchemy atomic update 原子累加
  指标,避免 ORM read-modify-write 的并发覆盖问题
- 启动顺序修正: tunnel 状态重置移到 Hub 连接建立之前,避免竞态
- OAuth 批量导入: 动态超时(默认30s,走代理60s),Kiro 适配器透传
- 提取 normalize_heartbeat_id 到 tunnel_protocol 共享模块,消除重复
2026-03-02 12:27:51 +08:00
fawney19 f3b9f42202 refactor(tunnel): 移除直连tunnel模式,统一使用Hub转发
- 删除 TunnelTransport 及 TunnelManager 相关引用,所有 tunnel 请求统一走 Hub
- 简化 health_scheduler,不再依赖进程内 TunnelManager 状态判断节点在线
- 简化 resolver,移除本地 tunnel miss 判断与多 worker 告警逻辑
- 简化 service 中 tunnel 连通性检测,统一以 DB 状态为准
- 启动/关闭流程移除 hub_enabled 分支,始终初始化 Hub 连接
- Rust 端 RequestMeta.timeout 增加浮点数反序列化支持,Python 端确保发送整数
2026-03-02 11:45:30 +08:00
fawney19 0564893c4f refactor: Hub二进制改为Docker构建时下载,优化部署流程与worker初始化
- Dockerfile: 移除COPY预编译二进制,改为构建时通过HUB_TAG从GitHub Release下载
- CI: 移除artifact上传/下载步骤,通过build-args传递Hub tag
- deploy.sh: 重构参数解析,支持--hub-tag指定版本,跟踪tag变化触发重建
- build.sh: 新增--image模式支持构建并推送Hub Docker镜像
- hub.rs: worker连接时同步所有节点在线状态,避免状态不一致
- proxy_nodes: 启动时主动建立Hub worker连接,消除懒连接窗口期
- docker-compose.yml: app镜像支持APP_IMAGE环境变量配置
- README: 更新部署文档,推荐本地构建方式
2026-03-02 11:24:14 +08:00
fawney19 c9dbe7936b fix: 补充tunnel端点nginx代理的路径前缀 2026-03-02 04:18:28 +08:00
fawney19 3f7a3d7600 fix: 修复tunnel端点nginx代理端口未替换的占位符问题 2026-03-02 04:14:48 +08:00
fawney19 5b9d9452c9 fix: 修改hub构建方式 2026-03-02 04:05:40 +08:00
fawney19 b7ef900181 fix(deploy): 恢复部署时自动拉取最新代码 2026-03-02 03:09:15 +08:00
fawney19 3eef673885 fix: 添加cargo清华镜像源,加速Hub本地构建 2026-03-02 03:04:56 +08:00
fawney19 039a18c243 feat(tunnel): 引入 aether-hub 帧路由器,支持多 worker 共享 tunnel 连接
新增 Rust 实现的 aether-hub 服务,作为 Docker 容器内部 WebSocket 帧路由器,
解决多 Gunicorn worker 进程间 tunnel 连接隔离问题。

主要改动:
- 新增 aether-hub Rust 项目,实现 proxy/worker 双向帧路由与 stream_id 重映射
- 新增 HubConnectionManager/HubTunnelTransport,worker 通过 Hub 转发 tunnel 帧
- 新增 create_tunnel_transport 工厂函数,按运行环境自动选择 Hub 或直连模式
- 新增 NODE_STATUS 广播机制,Hub 实时通知所有 worker 节点连接状态变化
- CI/CD 新增 build-hub job,Dockerfile 集成 Hub 二进制,deploy.sh 适配 Hub 构建
- 默认 GUNICORN_WORKERS 从 4 降为 2
2026-03-02 02:43:14 +08:00
fawney19 97d42703da feat(codex): 拆分openai:compact为独立端点,简化Codex请求为透传模式
- 新增openai:compact端点类型(EndpointKind.COMPACT),独立于openai:cli
- OpenAICompactAdapter继承OpenAICliAdapter,自动标记compact模式
- Codex请求补丁改为纯透传:仅清理内部标记,不再修改客户端payload
- stream_policy支持openai:compact独立策略,compact端点移除stream字段
- candidate_builder支持compact回退到cli端点
- auth_type: vertex_ai重命名为service_account,保持向后兼容
- Vertex Provider新增api_formats与auth_type组合校验
- KeyAllowedModels对话框改为从Provider获取模型,展示provider_model_name
- Dialog内Select组件自动禁用Portal,修复层级遮挡问题
- 新增Codex compact端点回填迁移脚本
2026-03-01 23:55:26 +08:00
fawney19 4bf3a453e7 feat(vertex-ai): 重构 Vertex AI 为插件化 adapter,支持 service_account 认证与动态路由
将 Vertex AI 从 transport.py 的硬编码逻辑重构为独立的 plugin adapter,
支持 service_account/oauth 认证类型、模型格式自动识别、区域路由和 URL 构建。
前端新增 Key 认证类型选择和 Service Account 配置表单。

Co-authored-by: NyaDoo <[email protected]>
Closes #194
2026-03-01 23:32:48 +08:00
fawney19 a137601728 feat(codex): 支持compact端点非流式请求,更新请求头与字段清理
- Codex适配器支持compact端点:非流式请求使用application/json Accept头,
  stream_policy根据compact上下文返回FORCE_NON_STREAM
- 更新Codex请求头格式:添加Version/Connection头,header key首字母大写
- 简化include列表处理:normalizer和request_patching统一强制为固定列表
- 清理Codex不支持的字段:truncation、context_management、user
- 默认instructions改为空字符串
- 前端用量页面:用户页面使用前端筛选后总数,避免不必要的后端分页请求
2026-03-01 18:20:42 +08:00
fawney19 d4840df447 feat(tunnel): 重连指数退避、多worker兼容与手动节点请求计数
- Proxy tunnel 重连策略从固定1s改为指数退避+jitter,首次重试立即执行,
  稳定连接30s后重置退避计数,上限3s保证快速恢复
- 多连接启动时增加错峰延迟,避免同时发起连接风暴
- 服务端 tunnel ping间隔和空闲超时支持环境变量配置
- 修复多worker启动时tunnel状态重置逻辑,仅leader执行重置避免覆盖其他worker连接
- Resolver增加本地tunnel缺失的限频告警和更短缓存TTL,加速多worker场景恢复
- 用量记录中统计手动代理节点的请求数和失败数
2026-03-01 17:37:11 +08:00
fawney19 2c1e51a490 refactor(compression): 请求压缩策略改为跟随客户端行为,响应支持gzip压缩
- 移除全局 ENABLE_REQUEST_COMPRESSION 配置,改为根据客户端 Content-Encoding
  决定是否对上游请求体进行 gzip 压缩
- 非流式响应根据客户端 Accept-Encoding 返回 gzip 压缩的 JSON
- ApiRequestContext 记录客户端编码偏好并透传至 handler 链路
- 新增 http_compression 模块统一处理压缩相关判断逻辑
- 上游请求头丢弃列表新增 content-encoding 防止客户端值泄露
- ensure_json_body 支持解压 gzip 编码的请求体
2026-03-01 15:52:32 +08:00
fawney19 2dcccc9820 fix(kiro): 修复工具schema兼容性、消息交替和重复内容问题
- 递归清理工具schema中Kiro不支持的additionalProperties和空required字段
- 将thinking prefix注入从history移到currentMessage,仅作用于当前轮次
- 修复连续assistant消息缺少user消息导致角色不交替的问题
- 增加流式content事件去重,跳过Kiro发送的重复内容
2026-03-01 14:25:38 +08:00
fawney19 8fa97ab8da feat(provider): 更新请求模型添加格式转换开关字段 2026-03-01 13:08:20 +08:00
fawney19 a66fa9792d fix(stream): 修复流式请求超时和取消归因问题
- 流式请求 timeout 改为 None,避免 provider.request_timeout 作为整条流总时长超时导致长响应被硬切断
- 重构 CancelledError 断连归因逻辑:提取探测方法并支持多次重试确认,降低误判率
- 新增 cancelled_unknown 状态处理断连检测不确定的场景,避免错误归因为 server_cancelled
- 在 upstream_response 中记录取消详情,便于排查
- 新增断连归因逻辑的单元测试
2026-03-01 13:02:26 +08:00
fawney19 dc2bc83c17 refactor(frontend): 改进故障转移规则对话框的状态码输入与校验
状态码输入从即时双向解析改为原始字符串绑定,保存时统一校验,
增加明确的错误提示;提取正则验证函数;优化小屏布局防止按钮挤压
2026-03-01 12:37:08 +08:00
fawney19 8b0e92e408 perf(frontend): 管理页面更新操作改为局部刷新,避免全量重载列表
Provider、API Key、Pool Key、Management Token 的更新操作
完成后直接替换本地列表中对应记录,仅创建操作保留全量刷新。
2026-03-01 11:57:09 +08:00
fawney19 b61fc4eb6b feat(test): 模型测试支持故障转移,展示每次尝试详情
- 后端新增 test-model-failover 接口,支持 global/direct 两种测试模式
- 利用 FailoverEngine 遍历候选并记录每次尝试的状态、延迟、错误等详情
- 前端 ModelsTab/ModelMappingTab 切换到新接口,移除格式选择下拉菜单
- 新增 TestResultDialog 组件,失败时展示候选尝试详情表格
- 简化组件 props 传递,移除不再需要的 endpoints/mappingPreview 依赖
2026-03-01 03:37:41 +08:00
fawney19 fea3d183bf Merge pull request #195 from AAEE86/keys
feat(pool): 新增账号配额展示并清理 Codex 旧限额字段
2026-03-01 00:59:49 +08:00
fawney19 20ed9cd123 feat(proxy): worker 退出时优雅关闭 tunnel 连接,提升 max-requests 默认值
- TunnelManager 新增 shutdown_all 方法,drain 飞行中请求后发送 GoAway
- shutdown 期间标记 draining 拒绝新请求进入
- gunicorn max-requests 默认值从 4000 提升到 50000,减少不必要的 worker 重启
2026-03-01 00:54:37 +08:00
AAEE86 d7a8a89aa7 feat(pool): 新增账号配额展示并清理 Codex 旧限额字段
- 后端池子 Key 列表新增 account_quota 字段,按 provider_type 生成配额摘要(codex/kiro/antigravity)
- 前端 PoolManagement 新增“配额”列与进度条展示,桌面端与移动端同步支持
- Provider 详情页移除 Codex code_review 限额展示,仅保留周限额与 5H 限额
- 清理 codex usage parser/realtime quota 中 code_review 相关解析与比较逻辑
- 同步更新调度与配额相关测试,改为忽略无关历史字段
2026-03-01 00:32:14 +08:00
fawney19 005cc3e388 feat(failover): 支持 Provider 级别故障转移规则,默认全部转移策略
- 新增 failover_rules 配置:支持 success_failover_patterns(成功响应匹配时转移)
  和 error_stop_patterns(错误响应匹配时终止),支持按 status_code 过滤
- 修改默认转移策略:ErrorClassifier 不再返回 RAISE,所有错误默认继续转移
- TaskService 中客户端错误不再直接抛出,改为 break 继续尝试下一个候选
- 修复 proxy tunnel 连接/断连竞态:引入 per-node 锁和事件时间戳排序
- 优化 ProxyNode 状态判定:OFFLINE 统一由心跳超时判定,兼容多 worker 场景
- has_tunnel 改为纯检查方法,避免在 finally 块中误清理新注册连接
- Redis stream NOGROUP 异常自愈处理
- OAuthAccountDialog 输入框焦点样式补全
2026-03-01 00:21:16 +08:00
fawney19 fbcb54a8a5 feat(oauth): 优化凭据导入,支持多文件选择与更多 JSON 格式
前端:简化导入界面状态管理,支持多文件拖拽/选择并自动合并内容,
移除冗余的 importFileName/manualPasteText 状态。
后端:_parse_tokens_input 新增支持 JSON 对象数组和单个 JSON 对象格式解析。
2026-02-28 22:06:21 +08:00
github-actions[bot] 85a126f48a chore(proxy): update download links for proxy-v0.2.1 2026-02-28 13:09:33 +00:00
fawney19 44cd35c10e chore(proxy): bump aether-proxy to 0.2.1 2026-02-28 21:02:21 +08:00
fawney19 a2d1cff3b0 perf(transport): 全链路传输压缩优化
- 上游请求启用 HTTP/2 (HPACK 头部压缩 + 多路复用),添加 h2 依赖
- 上游请求体超过阈值时自动 gzip 压缩,使用紧凑 JSON 序列化
- 添加 brotli 依赖,Accept-Encoding 支持 gzip/deflate/br
- 隧道帧压缩: Rust 端响应帧和 Python 端请求/响应帧均支持 gzip
- Rust 端压缩/解压逻辑统一提取到 protocol.rs
- Rust 端请求头构建改用 .headers() 替换 reqwest 默认值
- 新增 ENABLE_HTTP2、ENABLE_REQUEST_COMPRESSION 等环境变量配置
2026-02-28 21:02:21 +08:00
fawney19 3ff67fec2f Merge pull request #193 from AAEE86/keys
refactor(provider-keys): 拆分 keys 端点逻辑并补充配额刷新测试
2026-02-28 17:15:19 +08:00
AAEE86 6a8b5e6c8e fix(provider_keys): 提升 Codex 配额异步同步的可靠性与可观测性
- 为 flush 引入 FlushResult,统一返回更新数与重试批次
- 增加指数退避与失败日志限流,成功后重置 backoff
- 批量提交失败时回退到单条提交,降低整批失败风险
- 补充测试,覆盖 flush 重试与提交失败回退场景
2026-02-28 16:54:37 +08:00
AAEE86 92e9caf57e feat(usage): 支持 Codex 配额响应头实时异步同步
- 新增 parse_codex_usage_headers,统一解析响应头中的 Codex 配额信息
- 新增实时配额同步与异步调度器,按 provider_api_key_id 去重并后台落库
- 在 usage 记录与结算流程中接入配额同步投递,并在应用生命周期中启动/停止调度器
- 补充 codex_realtime_quota 与 codex_quota_sync_dispatcher 相关单元测试
2026-02-28 16:35:56 +08:00
AAEE86 10ccbf109c Merge branch 'fawney19:master' into keys 2026-02-28 16:13:30 +08:00
fawney19 f2dc51434d refactor(models): 删除拆分的模型子模块文件,统一回 database.py
- 删除 _base.py, auth.py, misc.py, model.py, provider.py, stats.py, usage.py, user.py
- Usage 表新增 cache_creation_input_tokens_5m/1h 字段支持按缓存 TTL 细分计费
- Model.global_model_id 改为 nullable=False,强制关联 GlobalModel
2026-02-28 16:12:32 +08:00
fawney19 9d751e3525 fix: 修复迁移脚本 2026-02-28 15:33:43 +08:00
fawney19 27047e5880 fix(migration): backfill 前清空 user_model_usage_counts 确保幂等性 2026-02-28 15:28:34 +08:00
fawney19 4fdebfc78e perf(usage): 优化 admin usage records 查询性能
- count 查询按需 JOIN,避免不必要的表关联
- 前端 onMounted 将 stats/heatmap/records/users 全部并行加载
- 调整缓存 TTL(聚合 30s->60s,列表 10s->15s)
- 为 request_candidates 添加复合索引优化 fallback/retry 查询
2026-02-28 15:09:21 +08:00
fawney19 11832edf49 Merge pull request #192 from AAEE86/master
refactor(usage): 移除时间线中的格式转换标签展示
2026-02-28 14:47:18 +08:00
fawney19 87d50052cc refactor(proxy): 移除 unhealthy 状态、前端展示失败率替换连接数、命令提示修正
- 移除 ProxyNodeStatus.UNHEALTHY 枚举值,仅保留 online/offline
- 迁移脚本将已有 unhealthy 数据迁移为 offline,upgrade/downgrade 均有幂等性保护
- 前端代理节点列表用失败率列替换连接数列,超过 5% 高亮显示
- 节点列表排序从 updated_at DESC 改为 name ASC
- Rust 端命令行提示从 aether-proxy 改为 ./aether-proxy
2026-02-28 14:44:44 +08:00
AAEE86 08b89b7ef8 refactor(provider-keys): 拆分 keys 端点逻辑并补充配额刷新测试
将 admin keys 接口的创建、更新、查询、删除、导出与配额刷新逻辑迁移到 provider_keys 服务层

新增 auth_type 归一化、重复校验、响应构建与 quota_refresh(codex/kiro/antigravity)模块,并补充对应单元测试
2026-02-28 14:12:08 +08:00
fawney19 4b02078b60 fix(proxy): 自动注册节点时按 ip+port 匹配而非 name 2026-02-28 14:10:44 +08:00
fawney19 1d644de500 fix: 修复迁移脚本 2026-02-28 14:03:42 +08:00
fawney19 54530faf03 feat(proxy): 节点状态简化、连接事件记录、可靠性指标与批量删除
- 移除 UNHEALTHY 中间状态,节点状态简化为 ONLINE/OFFLINE
- 新增 proxy_node_events 表记录 tunnel 连接/断开/错误事件
- 新增 failed_requests/dns_failures/stream_errors 可靠性指标(增量累加)
- tunnel 重连改为固定 1s 延迟,移除指数退避逻辑
- resolver/service 改为以 TunnelManager 内存状态判断节点可用性,避免 DB 竞态
- 修正 Claude cache_control 字段格式,使用 ttl 字段控制缓存时长
- 移除前端手动勾选 capability 的 UI,改为从价格配置自动推断
- 新增全局模型批量删除 API,替换前端并行单个删除
2026-02-28 13:52:32 +08:00
fawney19 ecb16d345a feat: 缓存计费细分、能力匹配优化、用户模型调用计数
1. 缓存创建 tokens 区分 5min/1h TTL,支持按缓存时长差异化计费
   - Usage 表新增 cache_creation_input_tokens_5m/1h 字段
   - Claude handler 解析新格式 (ephemeral_5m/1h, claude_cache_creation_5/1h)
   - 计费规则支持 cache_ttl_pricing 覆盖 cache_creation 价格

2. 能力匹配机制优化
   - COMPATIBLE 能力不再硬过滤,改为排序阶段通过 capability_miss_count 优先级处理
   - cache_1h 改为 COMPATIBLE + REQUEST_PARAM(自动检测请求体中的 ttl=1h)
   - gemini_files 改为 EXCLUSIVE + REQUEST_PARAM(自动检测 fileData.fileUri)
   - 移除前端模型偏好/能力配置 UI(不再需要用户手动配置)

3. 新增用户-模型维度调用次数计数器 (UserModelUsageCount)
   - 原子递增,避免从 Usage 表聚合查询
   - 前端模型目录和用户可用模型列表展示调用次数

4. 其他改进
   - global_model_id 改为必填(NOT NULL),清理孤立模型
   - 模型映射对话框支持从上游获取模型列表并分组折叠
   - 端点测试不再依赖端点启用状态
   - 异步任务页面对普通用户隐藏用户信息列
   - Dashboard 响应式布局断点调整 (sm -> lg)
   - 号池管理仅展示已启用号池的提供商
2026-02-28 11:45:04 +08:00
AAEE86 daecdb7676 feat: 统一供应商凭据标签并优化统计展示布局 2026-02-28 10:09:51 +08:00
AAEE86 423eb95f7a refactor(usage): 移除时间线中的格式转换标签展示 2026-02-28 09:42:33 +08:00
fawney19 82bbed2720 Merge pull request #191 from AAEE86/master
feat(providers): 资源统计按 provider_type 显示密钥/账号
2026-02-28 09:18:55 +08:00
AAEE86 90e804f608 feat(providers): 资源统计按 provider_type 显示密钥/账号 2026-02-28 08:58:12 +08:00
fawney19 e748277902 feat(proxy): 安全加固与架构优化
- 引入 SafeDnsResolver 消除 DNS rebinding TOCTTOU 漏洞,DNS 缓存改为多地址存储
- 扩展私有 IP 检测范围(CGNAT 100.64/10、基准测试 198.18/15、保留 240/4)
- 请求处理增加 hop-by-hop 头过滤、URL scheme 校验、超时范围限制
- 动态配置从 RwLock 切换到 ArcSwap 实现无锁读取
- 启动注册失败的服务器支持后台自动重试
- WebSocket 帧大小上限提升至 64MiB 匹配 Python 端
- 心跳支持动态间隔更新,新增 failed_requests/dns_failures/stream_errors 指标
- 配置启动校验、systemd UMask=0077、配置文件权限 600
- Python 端支持 per-connection max_streams(X-Tunnel-Max-Streams)
2026-02-28 01:32:28 +08:00
fawney19 2a0c684e88 feat: 新增全局模型批量管理功能,优化提供商更新接口
- 模型管理页面新增批量管理对话框,支持搜索、快捷筛选和批量删除
- 快捷筛选支持:无提供商、无活跃提供商、已禁用、未调用、无价格
- 提供商 PATCH 接口返回完整的 ProviderWithEndpointsSummary
- 修复系统级格式转换开关的 tooltip 文案和按钮高亮样式
2026-02-28 00:01:15 +08:00
fawney19 5bfaee47cc refactor(proxy): 端点检查统一使用 proxy_config 替代 proxy_param
将 endpoint_checker、handler_adapter_base、gemini adapter 及
provider_query 中的 proxy_param 参数替换为 proxy_config,
通过 build_proxy_client_kwargs 统一构建代理客户端参数,
以支持 tunnel 模式代理。同时清理未使用的 import。
2026-02-27 22:41:31 +08:00
fawney19 8de2f41924 fix(proxy-tunnel): 修复隧道连接池竞态与跨平台兼容问题
- 修复 TCP keepalive with_retries 在 Windows 上不可用的编译问题
- dispatcher 中 try_send 失败时记录警告日志而非静默丢弃
- 优化 handler_handles 清理策略,每 64 帧定期清理
- StreamState 记住原始连接引用,清理时避免连接池竞态
2026-02-27 21:41:43 +08:00
fawney19 ee356c5e6e feat(proxy-tunnel): 实现隧道连接池与 TCP 底层优化
Rust 端:
- 支持每个 server 多条并行 WebSocket 连接 (tunnel_connections 配置)
- 手动控制 TCP 连接: connect/handshake 超时、keepalive、NODELAY (socket2)
- 预构建共享 TLS ClientConfig 避免每次重连重新解析根证书
- 增加 stale timeout 检测无数据连接,智能 backoff 按连接存活时长重置
- 每条连接独立 reconnect 计数器,仅主连接 (conn_idx=0) 发送心跳
- dispatcher 的错误帧和 PONG 改用 try_send 避免阻塞读循环
- stream_handler 增加 frame 发送超时保护防止写阻塞

Python 端:
- TunnelManager 改为连接池,按 least-loaded 策略分配请求
- handle_incoming_frame 按连接实例路由,unregister 精确移除单条连接
- 心跳和 PONG 回复改为 fire-and-forget 避免阻塞主读循环
- send_frame 增加超时保护防止 TCP 写阻塞级联
- WebSocket 先 accept 再认证,auth 加超时
- 调整 idle timeout (90s) 和 ping 间隔 (15s)
2026-02-27 21:25:42 +08:00
fawney19 a174cf1b02 feat(proxy-tunnel): 增强隧道连接稳定性与恢复速度
- writer 增加 WebSocket Ping keepalive,防止中间代理空闲超时断开
- 服务端增加应用层 PING 循环(30s 间隔),空闲超时延长至 180s
- 重连基础延迟从 1000ms 降低到 500ms
- 节点缓存 TTL 缩短至 15s,不可用节点 TTL 缩短至 5s 加速恢复感知
- 心跳检测间隔从 30s 缩短到 15s
2026-02-27 19:51:38 +08:00
fawney19 934723f5f5 fix(frontend): 移除代理节点配置的号池模式限制
取消 Popover 组件的 v-if="provider.pool_advanced" 条件,
使代理节点配置在所有模式下均可使用。
2026-02-27 19:25:53 +08:00
fawney19 892c4235ec fix(frontend): 修正 system.ts 中 client 的导入方式为默认导入 2026-02-27 19:15:38 +08:00
fawney19 ddeb357c0e feat(proxy): 增加 0.1.x 到 0.2.0 配置自动迁移,放宽 max_retries 上限至 999
- proxy: 启动时检测旧版配置并自动迁移(delegate_* -> upstream_*、单服务器 -> [[servers]]),备份原文件为 .v1.bak
- proxy: 升级后 systemd 重启改为 best-effort,失败不中断升级流程
- provider: max_retries 上限从 10 放宽到 999(前端、后端模型同步调整)
2026-02-27 19:12:02 +08:00
github-actions[bot] 7cb007b445 chore(proxy): update download links for proxy-v0.2.0 2026-02-27 10:52:41 +00:00
fawney19 50262a2f02 chore(proxy): bump version to 0.2.0 2026-02-27 18:46:01 +08:00
fawney19 178fe4be4b Merge pull request #189 from AAEE86/fix
fix(frontend): 修复 Provider Key 列表只显示 100 条的问题
2026-02-27 18:37:16 +08:00
fawney19 4855e7cbca fix(task-poller): 捕获 CancelledError 避免关闭时产生 traceback
视频任务轮询在应用关闭时,Redis 异步操作被取消会抛出
CancelledError,由于其继承自 BaseException 而非 Exception,
原有的异常处理无法捕获,导致 APScheduler 打印完整 traceback。
拆分 poll_pending_tasks 为入口方法和 _do_poll 内部方法,
在入口层捕获 CancelledError 后静默返回。
2026-02-27 18:23:07 +08:00
fawney19 d2450cbdc1 feat(claude-code): 增加 TLS 指纹伪装、Cache TTL 统一、CLI 限制与流超时冷却
- 新增 curl_cffi Transport,支持真实浏览器 TLS 指纹伪装 (Chrome/Node.js)
- 增加 Cache TTL Override 功能,强制统一 cache_control 类型防止行为指纹差异
- 增加 CLI-only 客户端限制,支持仅允许 Claude Code CLI 访问
- 池健康策略增加 stream timeout 计数与自动冷却机制
- OAuth 账号 Region 选择改为从 AWS API 动态获取,支持搜索和自定义输入
- 前端 PoolConfigDialog 增加对应配置 UI
2026-02-27 18:10:27 +08:00
AAEE86 bfee1019ed fix(frontend): 修复 Provider Key 列表只显示 100 条的问题
- 在 getProviderKeys 中增加 skip/limit 自动分页拉取
  - 默认每页 1000,循环获取直到无更多数据
  - 避免后端默认 limit=100 导致账号展示被截断
2026-02-27 15:07:44 +08:00
fawney19 76a6d0ce8e feat(pool): 增加 OAuth 账号池管理功能
- 新增 pool manager / strategy / health_policy / cost_tracker / redis_ops 等核心模块
- 新增 pool admin API 路由与 schemas
- 新增 OAuth 账号类型解析 (oauth_plan)
- 前端增加 PoolManagement 页面、PoolConfigDialog、PoolImportDialog、PoolStatusCard 组件
- 补充 pool config / cost tracker / health policy / manager / strategy / trace 等测试
2026-02-27 13:58:58 +08:00
fawney19andAAEE86 579b5e4623 feat(provider): 增加 Claude Code 适配器、高级配置能力与 OAuth 账号类型统一解析
- 新增 Claude Code provider adapter (context, envelope, plugin, constants)
- 扩展 provider admin 路由,支持 Claude Code 高级配置 (CRUD)
- 统一 OAuth 账号类型解析逻辑,前后端对齐
- 重构 BatchAssignModelsDialog / ModelMappingDialog,简化组件逻辑
- handler 基类增强: request_builder 支持 Claude Code 信封格式
- CLI stream/sync mixin 适配 Claude Code 流式与同步模式
- 扩展 candidate builder / failover / scheduler 对 Claude Code 的支持
- 前端增加请求时间线可视化 (HorizontalRequestTimeline)
- 补充 Claude Code envelope / runtime controls / distributed sessions 等测试

Closes #183
Closes #185

Co-Authored-By: AAEE86 <[email protected]>
2026-02-27 13:54:46 +08:00
fawney19 f2f2a2dbc4 fix(proxy-node): 心跳和健康检查增加 tunnel 节点状态修正能力
- 心跳处理:收到心跳说明 tunnel 连通,若状态非 ONLINE 则修正
- 健康检查:移除 OFFLINE 过滤,允许检测 tunnel 重连后的状态恢复
- 前端配额刷新:就地更新 key metadata,避免重拉列表导致分页重置
2026-02-26 17:53:31 +08:00
fawney19 3fff147b9b Merge pull request #184 from rcdfrd/feat/new-api-checkin-apikey-support
feat(new-api): 签到支持 API Key 认证模式
2026-02-26 17:52:37 +08:00
rcdfrd a25f491f45 feat(new-api): 签到支持 API Key 认证模式
- 无 Cookie 时不再直接跳过,改为尝试以 API Key 发起签到请求
- 401/403 及认证失败类消息在 API Key 模式下静默跳过,不再误报 cookie_expired
- 将 message.lower() 预计算为 message_lower,避免重复调用
- 修正 already_indicators 与 auth_fail_indicators 中各条目也做 lower 处理,确保大小写不敏感比较一致
2026-02-26 15:42:42 +08:00
fawney19 8c2928c39d refactor(proxy-node): 移除非 tunnel 模式兼容,注册和心跳统一为 tunnel 模式
- 注册接口移除 tunnel_mode 参数,固定按 name upsert 并标记为 tunnel 模式
- 心跳接口拒绝非 tunnel 模式节点,返回升级提示
- 移除按 ip+port 查找的旧模式分支
2026-02-26 13:54:32 +08:00
fawney19 7e73e7fab2 fix(proxy-node): tunnel 模式节点注册不再覆盖由连接管理的状态
register_node 对 tunnel 模式已有节点不再写 status 字段,
状态完全由 _update_tunnel_status 和 health_scheduler 管理,
避免心跳注册将已连接的 ONLINE 状态覆盖为 UNHEALTHY。
2026-02-26 13:38:08 +08:00
fawney19 61535dfb7e feat(proxy-node): tunnel 节点支持连通性测试
通过 TunnelTransport 经 WebSocket tunnel 发送测试请求,
测量往返延迟并获取出口 IP,与手动代理节点测试语义一致。
2026-02-26 13:16:06 +08:00
fawney19 54332d31bc fix(proxy-node): 监控 writer task 退出以避免 tunnel 连接半关闭时 dispatcher 阻塞
当对端关闭连接导致 write half 退出但 read half 仍打开时,
dispatcher 会在 ws_stream.next() 上永久阻塞。通过在 tokio::select!
中监控 writer_handle,检测到 writer 退出后立即触发重连。
2026-02-26 13:07:42 +08:00
fawney19 67b8eb5c2f fix(proxy-node): tunnel 模式节点状态由连接管理,心跳不再覆盖
- resolver 中增加 tunnel 未连接节点的过滤,避免请求路由到不可达节点
- 心跳处理中 tunnel 模式节点仅更新指标,不改变在线状态
2026-02-26 12:50:08 +08:00
fawney19 27cef96789 refactor(proxy-node): 统一代理解析,全面支持 tunnel 模式
新增 resolve_ops_proxy_config 合并 proxy 和 tunnel_node_id 的解析,
避免各架构重复调用 _resolve_effective_node。所有架构连接器
(anyrouter/nekocode/sub2api/yescode) 和 HTTPClientPool 均适配
tunnel 模式,通过 TunnelTransport 替代传统代理。
2026-02-26 12:25:23 +08:00
fawney19 560345a889 feat(orchestration): 增加 AWS 账号被暂停 (suspended) 的错误识别和自动停用处理
- ErrorClassifier 新增 403 suspended 状态检测,归类为 ProviderAuthException
- ErrorHandlerService 新增 _is_account_suspended 静态方法,匹配多种 suspended 错误文本
- 403 suspended 的 OAuth key 自动标记为账号异常并停用
- 重构 _mark_oauth_key_blocked 支持自定义 reason 参数
- 将 _is_account_validation_required 调用改为静态方法调用
2026-02-26 11:32:31 +08:00
fawney19 f41bfecf0b feat(usage): 增强 cURL/Replay 对 Vertex AI、OAuth 和 envelope 提供商的支持
- _resolve_provider_auth 返回 decrypted_auth_config 用于 Vertex AI URL 构建
- _build_provider_url_safe 新增 API 格式默认路径回退及模板变量防护
- Replay 适配器基于 api_format 智能匹配端点和 Key
- Replay 适配器支持 envelope 包装(kiro/codex/antigravity 等特殊提供商)
- 复用 target_provider_obj 减少重复 Provider 查询
2026-02-26 10:32:29 +08:00
fawney19 80438e1d61 fix(proxy-node): 修复服务重启后 tunnel 连接状态不一致的问题
- 启动时重置 DB 中残留的 tunnel_connected=True 状态
- health_scheduler 以 TunnelManager 内存实际连接为准判断节点状态
- tunnel 模式注册时初始状态设为 UNHEALTHY,等 tunnel 连接后再上线
2026-02-26 09:15:00 +08:00
fawney19 19f1be03fa feat(nginx): 添加 WebSocket 隧道端点的 nginx 代理配置 2026-02-26 03:33:53 +08:00
fawney19andAAEE86 b6ca16e084 fix(headers): 兼容非 ASCII header 值,避免 httpx ASCII 编码报错
将 HeaderBuilder 中的 latin-1 透传方案替换为 ASCII 归一化策略:
- 对 x-codex-turn-metadata 做 JSON 重编码(ensure_ascii=true)保留语义
- 其他非 ASCII 头值仅转义非 ASCII 字符为 \uXXXX,保留 ASCII 字符原样
- 新增单测覆盖中文 header 与 codex turn metadata 场景

Close #180

Co-Authored-By: AAEE86 <[email protected]>
2026-02-26 02:56:54 +08:00
fawney19 c0b80c923a refactor(usage): 移除 RequestHeadersContent 中未使用的 diff props
移除 clientHeadersWithDiff 和 providerHeadersWithDiff 属性及相关计算逻辑
2026-02-26 02:36:37 +08:00
fawney19 61e0959a06 Merge branch 'feat/compare-view' 2026-02-26 02:32:29 +08:00
fawney19 0ecf8b703e feat(api-keys): 支持独立密钥额度重置(手动+定时自动)
- 新增手动重置接口 PATCH /api/admin/api-keys/{id}/reset-usage
- 前端 ApiKeys 页面增加重置按钮,支持确认后归零已使用额度
- 新增独立密钥额度定时自动重置任务,支持配置周期和执行时间
- 支持 all/selected 两种重置模式,selected 模式可指定密钥
- 移除已废弃的 CleanupScheduler 兼容别名
2026-02-26 02:16:48 +08:00
fawney19 1a1bb0a99c feat(admin): 新增数据管理模块,支持分类清空系统数据
添加 purge API 支持按类别清空配置、用户、使用记录、审计日志、请求体和统计数据,
前端系统设置页新增数据管理区块。
2026-02-26 01:20:24 +08:00
fawney19 5415057a5d fix(models): 模型创建对话框支持连续添加
创建模式下提交后保持对话框打开,方便批量添加模型;
编辑模式提交后仍正常关闭对话框。按钮文案调整为"添加"。
2026-02-25 23:33:16 +08:00
fawney19 a2493b4bc0 fix(compatibility): 同族格式透传不再依赖格式转换开关
将 data_format_id 相同的格式对(如 claude:chat / claude:cli)的透传判断
提前到三层开关检查之前,使其无需开关即可直接透传。
同时用 pytest.mark.parametrize 精简同族格式测试用例。
2026-02-25 23:06:25 +08:00
fawney19 fd9040b9aa refactor(proxy): 将 aether-proxy 从 HMAC 正向代理迁移到 WebSocket 隧道模式
移除 HMAC 认证、TLS 自签名证书、HTTP CONNECT 代理和代发(delegate)模式,
改为 aether-proxy 主动通过 WebSocket 连接 Aether 服务端建立隧道。

Aether 服务端新增:
- WebSocket 隧道端点 (proxy_tunnel.py)
- TunnelManager 管理隧道连接和请求分发
- TunnelTransport 作为 httpx 自定义 transport 层
- 基于二进制帧的隧道协议 (tunnel_protocol.py)

aether-proxy (Rust) 重构:
- 新增 tunnel 模块 (client/dispatcher/stream_handler/protocol)
- 支持多 Aether 服务端连接 ([[servers]] 配置)
- 移除 proxy/auth/delegate 模块和 hyper 依赖
- 改用 tokio-tungstenite 实现 WebSocket 客户端

同时:
- 添加浏览器指纹 Headers 绕过 Cloudflare 防护
- 删除节点时自动清理 Provider/Endpoint 的代理引用
- 数据库迁移: 新增 tunnel_mode/tunnel_connected/tunnel_connected_at 字段
2026-02-25 21:59:29 +08:00
LewisPen adaf9e4b93 feat(usage): 响应头 Tab 增加并排对比模式
泛化 RequestHeadersContent 组件支持任意 header 对,
响应头 Tab 复用同一组件实现客户端/提供商响应头 Diff 对比。
2026-02-24 15:25:59 +08:00
fawney19 39b036abd5 refactor(docs/home): 重构首页导航和文档结构
- README 新增架构图(明暗主题 SVG),更新 Proxy 描述
- 首页导航栏将"文档"链接从底部按钮区移至顶部导航
- 删除独立的 body-rules-spec 文档
- GuideLayout 路由切换时自动滚动到顶部
- 架构图组件移除网格背景
2026-02-24 12:15:14 +08:00
947 changed files with 139215 additions and 37271 deletions
+1
View File
@@ -11,6 +11,7 @@ ENV/
.uv/
*.egg-info/
dist/
!aether-hub/dist/aether-hub
build/
*.egg
+52 -8
View File
@@ -21,11 +21,9 @@ JWT_SECRET_KEY=change-this-to-a-secure-random-string
# 注意:更换此密钥后需要在管理面板重新配置所有 Provider API Key
ENCRYPTION_KEY=change-this-to-another-secure-random-string
# 代理节点 HMAC 密钥(用于 aether-proxy 认证)
# 可选:不设置时会从 ENCRYPTION_KEY 自动派生
# 显式设置时,aether-proxy.toml 的 hmac_key 配置相同值即可
# 可通过 python generate_keys.py 生成
# PROXY_HMAC_KEY=change-this-to-a-proxy-hmac-key
# 支付回调共享密钥(公开 /api/payment/callback/* 入口必须携带 x-payment-callback-token)
# 建议使用 32+ 位随机字符串
PAYMENT_CALLBACK_SECRET=change-this-to-a-secure-callback-secret
# 管理员账号(仅首次初始化时使用, 创建完成后可在系统内修改密码)
ADMIN_EMAIL=[email protected]
@@ -38,15 +36,53 @@ ADMIN_PASSWORD=admin123456
# 应用端口(默认 8084)
# APP_PORT=8084
# Gunicorn Worker 数量(默认 4)
# 建议最小设置为 2
# GUNICORN_WORKERS=4
# 生产部署镜像(deploy.sh 会读取)
# APP_IMAGE=ghcr.io/fawney19/aether:latest
# 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
@@ -58,6 +94,14 @@ ADMIN_PASSWORD=admin123456
# 默认: * (允许所有源)
# 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
+91
View File
@@ -0,0 +1,91 @@
name: Build aether-hub
on:
push:
tags: ['hub-v*']
workflow_dispatch:
permissions:
contents: write
jobs:
build:
name: ${{ matrix.name }}
runs-on: ubuntu-latest
strategy:
fail-fast: false
matrix:
include:
- name: linux-amd64
target: x86_64-unknown-linux-gnu
use_cross: true
- name: linux-arm64
target: aarch64-unknown-linux-gnu
use_cross: true
steps:
- uses: actions/checkout@v5
- name: Install Rust toolchain
uses: dtolnay/rust-toolchain@stable
with:
targets: ${{ matrix.target }}
- name: Rust cache
uses: Swatinem/rust-cache@v2
with:
workspaces: aether-hub -> target
key: ${{ matrix.target }}
- name: Install cross
if: matrix.use_cross
uses: taiki-e/install-action@cross
- name: Build
working-directory: aether-hub
shell: bash
run: |
if [ "${{ matrix.use_cross }}" = "true" ]; then
cross build --release --target ${{ matrix.target }}
else
cargo build --release --target ${{ matrix.target }}
fi
- name: Package
shell: bash
run: |
cd aether-hub/target/${{ matrix.target }}/release
chmod +x aether-hub
tar czf ../../../../aether-hub-${{ matrix.name }}.tar.gz aether-hub
- name: Upload artifact
uses: actions/upload-artifact@v5
with:
name: aether-hub-${{ matrix.name }}
path: aether-hub-*.tar.gz
if-no-files-found: error
release:
needs: build
runs-on: ubuntu-latest
if: startsWith(github.ref, 'refs/tags/')
steps:
- name: Download all artifacts
uses: actions/download-artifact@v5
with:
merge-multiple: true
path: artifacts
- name: Generate checksums
working-directory: artifacts
run: sha256sum aether-hub-* > SHA256SUMS.txt
- name: Create GitHub Release
uses: softprops/action-gh-release@v2
with:
name: "${{ github.ref_name }}"
generate_release_notes: true
files: |
artifacts/aether-hub-*
artifacts/SHA256SUMS.txt
fail_on_unmatched_files: true
+8 -8
View File
@@ -44,7 +44,7 @@ jobs:
use_cross: false
steps:
- uses: actions/checkout@v4
- uses: actions/checkout@v5
- name: Install Rust toolchain
uses: dtolnay/rust-toolchain@stable
@@ -87,7 +87,7 @@ jobs:
7z a ../../../../aether-proxy-${{ matrix.name }}.zip aether-proxy.exe
- name: Upload artifact
uses: actions/upload-artifact@v4
uses: actions/upload-artifact@v5
with:
name: aether-proxy-${{ matrix.name }}
path: |
@@ -101,7 +101,7 @@ jobs:
if: startsWith(github.ref, 'refs/tags/')
steps:
- name: Download all artifacts
uses: actions/download-artifact@v4
uses: actions/download-artifact@v5
with:
merge-multiple: true
path: artifacts
@@ -113,7 +113,7 @@ jobs:
- name: Create GitHub Release
uses: softprops/action-gh-release@v2
with:
name: "aether-proxy ${{ github.ref_name }}"
name: "${{ github.ref_name }}"
generate_release_notes: true
files: |
artifacts/aether-proxy-*
@@ -125,10 +125,10 @@ jobs:
runs-on: ubuntu-latest
if: startsWith(github.ref, 'refs/tags/')
steps:
- uses: actions/checkout@v4
- uses: actions/checkout@v5
- name: Download Linux artifacts
uses: actions/download-artifact@v4
uses: actions/download-artifact@v5
with:
pattern: aether-proxy-linux-*
merge-multiple: true
@@ -174,7 +174,7 @@ jobs:
latest=auto
- name: Build and push
uses: docker/build-push-action@v5
uses: docker/build-push-action@v6
with:
context: ./aether-proxy
push: true
@@ -187,7 +187,7 @@ jobs:
runs-on: ubuntu-latest
if: startsWith(github.ref, 'refs/tags/')
steps:
- uses: actions/checkout@v4
- uses: actions/checkout@v5
with:
ref: master
+4 -4
View File
@@ -18,12 +18,12 @@ jobs:
build:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
- uses: actions/checkout@v5
- name: Setup Node.js
uses: actions/setup-node@v4
uses: actions/setup-node@v5
with:
node-version: '20'
node-version: '22'
cache: 'npm'
cache-dependency-path: frontend/package-lock.json
@@ -41,7 +41,7 @@ jobs:
run: cp frontend/dist/index.html frontend/dist/404.html
- name: Setup Pages
uses: actions/configure-pages@v4
uses: actions/configure-pages@v5
- name: Upload artifact
uses: actions/upload-pages-artifact@v3
+83 -13
View File
@@ -15,6 +15,7 @@ env:
REGISTRY: ghcr.io
BASE_IMAGE_NAME: fawney19/aether-base
APP_IMAGE_NAME: fawney19/aether
GITHUB_REPO: fawney19/Aether
# Base image hash inputs:
# - Dockerfile.base
# - pyproject.toml (dependency fingerprint only; ignores tool/optional deps)
@@ -29,7 +30,7 @@ jobs:
outputs:
base_changed: ${{ steps.check.outputs.base_changed }}
steps:
- uses: actions/checkout@v4
- uses: actions/checkout@v5
- name: Log in to Container Registry
uses: docker/login-action@v3
@@ -107,7 +108,7 @@ jobs:
contents: read
packages: write
steps:
- uses: actions/checkout@v4
- uses: actions/checkout@v5
- name: Set up Docker Buildx
uses: docker/setup-buildx-action@v3
@@ -163,7 +164,7 @@ jobs:
org.opencontainers.image.base.hash=${{ steps.hash.outputs.hash }}
- name: Build and push base image
uses: docker/build-push-action@v5
uses: docker/build-push-action@v6
with:
context: .
file: ./Dockerfile.base
@@ -174,15 +175,36 @@ jobs:
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]
if: always() && (needs.build-base.result == 'success' || needs.build-base.result == 'skipped')
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@v4
- uses: actions/checkout@v5
- name: Set up Docker Buildx
uses: docker/setup-buildx-action@v3
@@ -243,15 +265,63 @@ jobs:
version_tuple = __version_tuple__
EOF
- name: Build and push app image
uses: docker/build-push-action@v5
- 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 }}
no-cache-filters: builder
cache-from: type=gha,scope=app
cache-to: type=gha,mode=min,scope=app
platforms: linux/amd64,linux/arm64
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
+6
View File
@@ -2,6 +2,7 @@
# Edit at https://www.toptal.com/developers/gitignore?templates=python
# AI Assistant Configuration
.codex/
.claude/
.serena/
.gemini*/
@@ -203,6 +204,7 @@ logs/
# Git backup
.git.backup/
.worktrees/
# Database backups
backups/
@@ -226,6 +228,10 @@ test.py
.deps-hash
.code-hash
.migration-hash
.hub-hash
# Hub prebuilt binaries
aether-hub/dist/
# Version file (auto-generated by hatch-vcs)
src/_version.py
+137 -5
View File
@@ -2,14 +2,22 @@
# 运行镜像:从 base 提取产物到精简运行时
# 构建命令: docker build -f Dockerfile.app -t aether-app:latest .
# 用于 GitHub Actions CI(官方源)
FROM aether-base:latest AS builder
WORKDIR /app
# 复制前端源码并构建(CI 通过 no-cache-filters=builder 确保每次重建)
COPY frontend/ ./frontend/
RUN cd frontend && npm run build
# ==================== 运行时镜像 ====================
FROM python:3.13-slim
WORKDIR /app
ARG HUB_RELEASE_REPO=fawney19/Aether
ARG HUB_TAG
ARG TARGETARCH
ARG GITHUB_TOKEN
# 运行时依赖(无 gcc/nodejs/npm,使用 BuildKit 缓存加速)
RUN --mount=type=cache,target=/var/cache/apt,sharing=locked \
--mount=type=cache,target=/var/lib/apt,sharing=locked \
@@ -17,13 +25,49 @@ RUN --mount=type=cache,target=/var/cache/apt,sharing=locked \
nginx \
supervisor \
libpq5 \
curl
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
@@ -46,6 +90,11 @@ RUN printf '%s\n' \
' "" $remote_addr;' \
'}' \
'' \
'map $http_upgrade $connection_upgrade {' \
' default upgrade;' \
' "" "";' \
'}' \
'' \
'server {' \
' listen 80;' \
' server_name _;' \
@@ -75,6 +124,39 @@ RUN printf '%s\n' \
' 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;' \
@@ -83,11 +165,31 @@ RUN printf '%s\n' \
' 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 Connection "";' \
' 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;' \
@@ -107,6 +209,25 @@ RUN printf '%s\n' \
' 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)' \
@@ -129,7 +250,7 @@ RUN printf '%s\n' \
'stderr_logfile=/var/log/nginx/error.log' \
'' \
'[program:app]' \
'command=/bin/bash -c "MAX_REQUESTS_JITTER=$((${MAX_REQUESTS:-4000}/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 --timeout 120 --max-requests ${MAX_REQUESTS:-4000} --max-requests-jitter $MAX_REQUESTS_JITTER --access-logfile - --error-logfile - --log-level info"' \
'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' \
@@ -137,7 +258,16 @@ RUN printf '%s\n' \
'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' > /etc/supervisor/conf.d/supervisord.conf
'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
# 入口脚本(启动前执行迁移)
@@ -149,8 +279,10 @@ ENV PYTHONUNBUFFERED=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=4 \
GUNICORN_WORKERS=2 \
MAX_REQUESTS=4000
EXPOSE 80
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
+148 -5
View File
@@ -2,6 +2,7 @@
# 运行镜像:从 base 提取产物到精简运行时(国内镜像源版本)
# 构建命令: docker build -f Dockerfile.app.local -t aether-app:latest .
# 用于本地/国内服务器部署
FROM aether-base:latest AS builder
WORKDIR /app
@@ -15,6 +16,16 @@ 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 \
@@ -23,7 +34,12 @@ RUN --mount=type=cache,target=/var/cache/apt,sharing=locked \
nginx \
supervisor \
libpq5 \
curl
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
@@ -33,6 +49,45 @@ 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; \
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
# 从 builder 阶段复制前端构建产物
COPY --from=builder /app/frontend/dist /usr/share/nginx/html
RUN chmod -R 755 /usr/share/nginx/html
@@ -57,6 +112,11 @@ RUN printf '%s\n' \
' "" $remote_addr;' \
'}' \
'' \
'map $http_upgrade $connection_upgrade {' \
' default upgrade;' \
' "" "";' \
'}' \
'' \
'server {' \
' listen 80;' \
' server_name _;' \
@@ -86,6 +146,39 @@ RUN printf '%s\n' \
' 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;' \
@@ -94,11 +187,31 @@ RUN printf '%s\n' \
' 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 Connection "";' \
' 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;' \
@@ -118,6 +231,25 @@ RUN printf '%s\n' \
' 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)' \
@@ -141,7 +273,7 @@ RUN printf '%s\n' \
'stderr_logfile=/var/log/nginx/error.log' \
'' \
'[program:app]' \
'command=/bin/bash -c "MAX_REQUESTS_JITTER=$((${MAX_REQUESTS:-4000}/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 --timeout 120 --max-requests ${MAX_REQUESTS:-4000} --max-requests-jitter $MAX_REQUESTS_JITTER --access-logfile - --error-logfile - --log-level info"' \
'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' \
@@ -149,7 +281,16 @@ RUN printf '%s\n' \
'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' > /etc/supervisor/conf.d/supervisord.conf
'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
@@ -164,8 +305,10 @@ ENV PYTHONUNBUFFERED=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=4 \
GUNICORN_WORKERS=2 \
MAX_REQUESTS=4000
EXPOSE 80
+13 -5
View File
@@ -22,6 +22,14 @@
Aether 是一个自托管的 AI API 网关,为团队和个人提供多租户管理、智能负载均衡、成本配额控制和健康监控能力。通过统一的 API 入口,可以无缝对接 Claude、OpenAI、Gemini 等主流 AI 服务及其 CLI 工具。
<p align="center">
<picture>
<source media="(prefers-color-scheme: dark)" srcset="docs/architecture/architecture-dark.svg">
<source media="(prefers-color-scheme: light)" srcset="docs/architecture/architecture-light.svg">
<img src="docs/architecture/architecture-light.svg" width="680" alt="Aether Architecture">
</picture>
</p>
页面预览: https://fawney19.github.io/Aether/
## 部署
@@ -40,7 +48,7 @@ python generate_keys.py # 生成密钥, 并将生成的密钥填入 .env
# 3. 部署 / 更新(自动执行数据库迁移)
docker compose pull && docker compose up -d
# 4. 升级前备份
# 4. 升级前备份 (可选)
docker compose exec postgres pg_dump -U postgres aether | gzip > backup_$(date +%Y%m%d_%H%M%S).sql.gz
```
@@ -56,6 +64,7 @@ cp .env.example .env
python generate_keys.py # 生成密钥, 并将生成的密钥填入 .env
# 3. 部署 / 更新(自动构建、启动、迁移)
git pull
./deploy.sh
```
@@ -73,9 +82,9 @@ uv sync
cd frontend && npm install && npm run dev
```
## Aether Proxy
## Aether Proxy (可选)
Aether Proxy 是配套的正向代理节点,部署在海外 VPS 上,为墙内的 Aether 实例中转 API 流量。支持 TUI 向导一键配置、systemd 服务管理、TLS 加密、DNS 缓存及连接池调优。
Aether Proxy 是配套的正向代理节点,部署在海外 VPS 上,为墙内的 Aether 实例中转 API 流量。或者部署在其他服务器为指定的提供商、账号、Key使用不同的节点访问。支持 TUI 向导一键配置、systemd 服务管理、TLS 加密、DNS 缓存及连接池调优。
- Docker Compose 部署或下载预编译二进制直接运行
- 通过 `aether-proxy setup` 完成交互式配置,自动注册为系统服务
@@ -102,7 +111,7 @@ Aether Proxy 是配套的正向代理节点,部署在海外 VPS 上,为墙
| `APP_PORT` | 8084 | 应用端口 |
| `API_KEY_PREFIX` | sk | API Key 前缀 |
| `LOG_LEVEL` | INFO | 日志级别 (DEBUG/INFO/WARNING/ERROR) |
| `GUNICORN_WORKERS` | 4 | Gunicorn 工作进程数 |
| `GUNICORN_WORKERS` | 2 | Gunicorn 工作进程数 |
| `DB_PORT` | 5432 | PostgreSQL 端口 |
| `REDIS_PORT` | 6379 | Redis 端口 |
@@ -172,4 +181,3 @@ docker compose up -d app
## Star History
[![Star History Chart](https://api.star-history.com/svg?repos=fawney19/Aether&type=Date)](https://star-history.com/#fawney19/Aether&Date)
+3
View File
@@ -0,0 +1,3 @@
target/
.git/
.DS_Store
+2012
View File
File diff suppressed because it is too large Load Diff
+27
View File
@@ -0,0 +1,27 @@
[package]
name = "aether-hub"
version = "0.2.0"
edition = "2021"
description = "Tunnel Hub for Aether - frame router between workers and proxies"
[dependencies]
tokio = { version = "1", features = ["full"] }
axum = { version = "0.8", features = ["ws"] }
serde = { version = "1", features = ["derive"] }
serde_json = "1"
tracing = "0.1"
tracing-subscriber = { version = "0.3", features = ["env-filter"] }
clap = { version = "4", features = ["derive", "env"] }
dashmap = "6"
parking_lot = "0.12"
flate2 = "1"
futures-util = "0.3"
bytes = "1"
async-stream = "0.3"
http-body-util = "0.1"
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls"] }
[profile.release]
lto = true
strip = true
codegen-units = 1
+34
View File
@@ -0,0 +1,34 @@
# syntax=docker/dockerfile:1
FROM rust:1.85-slim AS builder
WORKDIR /build/aether-hub
# 可选:配置国内 Cargo 镜像源(本地构建时传 --build-arg CARGO_MIRROR=1)
ARG CARGO_MIRROR
RUN if [ -n "$CARGO_MIRROR" ]; then \
printf '[source.crates-io]\nreplace-with = "tuna"\n\n[source.tuna]\nregistry = "sparse+https://mirrors.tuna.tsinghua.edu.cn/crates.io-index/"\n' \
> /usr/local/cargo/config.toml; \
fi
# 先构建依赖层,最大化后续代码变更时的缓存命中
COPY Cargo.toml Cargo.lock ./
RUN mkdir src && printf 'fn main() {}\n' > src/main.rs
RUN --mount=type=cache,target=/usr/local/cargo/registry,sharing=locked \
--mount=type=cache,target=/build/aether-hub/target,sharing=locked \
cargo build --release --locked
RUN rm -rf src
COPY src ./src
RUN --mount=type=cache,target=/usr/local/cargo/registry,sharing=locked \
--mount=type=cache,target=/build/aether-hub/target,sharing=locked \
cargo build --release --locked && \
cp target/release/aether-hub /tmp/aether-hub
FROM debian:bookworm-slim
RUN apt-get update && apt-get install -y --no-install-recommends ca-certificates && \
rm -rf /var/lib/apt/lists/*
COPY --from=builder /tmp/aether-hub /usr/local/bin/aether-hub
EXPOSE 8085
ENTRYPOINT ["/usr/local/bin/aether-hub"]
CMD ["--bind", "0.0.0.0:8085"]
+38
View File
@@ -0,0 +1,38 @@
# aether-hub
`aether-hub` 是 Tunnel Hub 服务,负责在 proxy 与 worker 之间路由帧。
已集成在Docker镜像中, 无需单独部署。
## 部署端指定 Hub 版本并构建
```bash
cd /path/to/Aether
./deploy.sh --hub-tag hub-v0.1.0
```
不指定 `--hub-tag` 时,`./deploy.sh` 会自动解析最新 `hub-v*` release,并在构建 app 镜像时从 GitHub Release 下载对应架构的 Hub 二进制。
## build.sh 模式说明
- 默认是 `binary` 模式(`cross` 构建二进制)。
- `--upload <hub-vX.Y.Z>` 会把构建产物上传到 GitHub Release。
- 加 `--image` 后进入镜像模式(`docker buildx`,可选)。
常用参数:
- `--tag <tag>`: 镜像 tag
- `--image-name <name>`: 镜像名(默认 `ghcr.io/fawney19/aether-hub`)
- `--platforms <list>`: 例如 `linux/amd64,linux/arm64`
- `--push`: 推送镜像
- `--load`: 加载到本地 Docker(单平台)
- `--latest`: 额外打 `latest` tag
## 运行时参数
- `TUNNEL_HUB_WORKER_IDLE_TIMEOUT`:worker 心跳空闲超时,默认 `60` 秒
- `TUNNEL_HUB_OUTBOUND_QUEUE_CAPACITY`:单连接出站队列容量,默认 `128`;队列打满时会把连接视为拥塞并主动关闭,避免 Hub 内存无限增长
## 与部署脚本关系
- `./deploy.sh`: 本地构建部署(会本地构建 app/base,并在构建 app 时从 GitHub Release 下载 Hub,可用 `--hub-tag` 固定版本)。
+272
View File
@@ -0,0 +1,272 @@
#!/bin/bash
# aether-hub 构建脚本
#
# 支持两种模式:
# 1) binary 模式(默认): 构建多架构二进制并可上传 GitHub Release
# 2) image 模式: 构建并推送/加载 Docker 镜像(推荐生产发布用)
#
# 示例:
# # binary 模式(兼容旧行为)
# ./build.sh
# ./build.sh amd64
# ./build.sh --upload hub-v0.1.0
#
# # image 模式(多架构推送)
# ./build.sh --image --tag v0.2.5 --push --latest
# ./build.sh --image --tag sha-abc123 --image-name ghcr.io/fawney19/aether-hub --push
# ./build.sh --image --tag local-test --platforms linux/amd64 --load
set -euo pipefail
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)"
PROJECT_DIR="$(cd "$SCRIPT_DIR/.." && pwd)"
DIST_DIR="$SCRIPT_DIR/dist"
# -------------------------------
# Defaults
# -------------------------------
MODE="binary" # binary | image
# binary mode options
UPLOAD=false
UPLOAD_TAG=""
BINARY_TARGETS=""
# image mode options
IMAGE_NAME="${IMAGE_NAME:-ghcr.io/fawney19/aether-hub}"
IMAGE_TAG=""
IMAGE_PLATFORMS="linux/amd64,linux/arm64"
IMAGE_PUSH=false
IMAGE_LOAD=false
IMAGE_LATEST=false
usage() {
cat <<'EOF'
用法:
./build.sh [binary-args]
./build.sh --image [image-args]
binary 模式(默认):
amd64|arm64 仅构建指定架构(可重复)
--upload <hub-vX.Y.Z> 上传到 GitHub Release(需要 gh CLI)
image 模式:
--image 启用镜像模式
--tag <tag> 镜像 tag(默认自动从 git describe 推导)
--image-name <name> 镜像名(默认 ghcr.io/fawney19/aether-hub)
--platforms <list> 平台列表,逗号分隔(默认 linux/amd64,linux/arm64)
--push 推送镜像到仓库
--load 加载到本地 Docker(仅单平台)
--latest 额外打 latest tag
通用:
-h, --help 显示帮助
EOF
}
while [ $# -gt 0 ]; do
case "$1" in
--image)
MODE="image"
shift
;;
--tag)
IMAGE_TAG="${2:-}"
shift 2
;;
--image-name)
IMAGE_NAME="${2:-}"
shift 2
;;
--platforms)
IMAGE_PLATFORMS="${2:-}"
shift 2
;;
--push)
IMAGE_PUSH=true
shift
;;
--load)
IMAGE_LOAD=true
shift
;;
--latest)
IMAGE_LATEST=true
shift
;;
--upload)
UPLOAD=true
UPLOAD_TAG="${2:-}"
shift 2
;;
amd64|arm64)
BINARY_TARGETS="$BINARY_TARGETS $1"
shift
;;
-h|--help)
usage
exit 0
;;
*)
echo "❌ 未知参数: $1"
usage
exit 1
;;
esac
done
build_binary() {
if [ -z "$BINARY_TARGETS" ]; then
BINARY_TARGETS="amd64 arm64"
fi
if ! command -v cross >/dev/null 2>&1; then
echo "❌ 需要安装 cross: cargo install cross --git https://github.com/cross-rs/cross"
exit 1
fi
mkdir -p "$DIST_DIR"
echo "🔨 开始构建 aether-hub 二进制..."
echo " 目标平台: $BINARY_TARGETS"
echo ""
ARTIFACTS=""
for arch in $BINARY_TARGETS; do
case "$arch" in
amd64) target="x86_64-unknown-linux-gnu" ;;
arm64) target="aarch64-unknown-linux-gnu" ;;
*) echo "❌ 未知架构: $arch"; exit 1 ;;
esac
echo ">>> 构建 $arch ($target)..."
cd "$SCRIPT_DIR"
cross build --release --target "$target" --locked
BIN="target/$target/release/aether-hub"
if [ ! -f "$BIN" ]; then
echo "❌ 未找到二进制文件: $BIN"
exit 1
fi
ARCHIVE="$DIST_DIR/aether-hub-linux-$arch.tar.gz"
tar czf "$ARCHIVE" -C "target/$target/release" aether-hub
ARTIFACTS="$ARTIFACTS $ARCHIVE"
SIZE=$(du -h "$ARCHIVE" | cut -f1)
echo "✅ $arch 构建完成: $ARCHIVE ($SIZE)"
echo ""
done
cd "$DIST_DIR"
shasum -a 256 aether-hub-*.tar.gz > SHA256SUMS.txt
echo "📋 SHA256 校验和:"
cat SHA256SUMS.txt
echo ""
if [ "$UPLOAD" = true ]; then
if [ -z "$UPLOAD_TAG" ]; then
echo "❌ --upload 需要指定 tag,例如: ./build.sh --upload hub-v0.1.0"
exit 1
fi
if ! command -v gh >/dev/null 2>&1; then
echo "❌ 需要安装 GitHub CLI: brew install gh"
exit 1
fi
echo "📦 上传到 GitHub Release: $UPLOAD_TAG"
cd "$PROJECT_DIR"
if ! git rev-parse "$UPLOAD_TAG" >/dev/null 2>&1; then
git tag "$UPLOAD_TAG"
git push origin "$UPLOAD_TAG"
fi
gh release create "$UPLOAD_TAG" \
--title "aether-hub ${UPLOAD_TAG#hub-}" \
--generate-notes \
$ARTIFACTS \
"$DIST_DIR/SHA256SUMS.txt"
echo "✅ 上传完成!"
fi
echo "🎉 binary 模式完成!"
}
build_image() {
if ! command -v docker >/dev/null 2>&1; then
echo "❌ 未找到 docker,请先安装 Docker"
exit 1
fi
if ! docker buildx version >/dev/null 2>&1; then
echo "❌ 未找到 docker buildx,请先启用 buildx"
exit 1
fi
if [ "$IMAGE_PUSH" = true ] && [ "$IMAGE_LOAD" = true ]; then
echo "❌ --push 与 --load 不能同时使用"
exit 1
fi
if [ "$IMAGE_PUSH" = false ] && [ "$IMAGE_LOAD" = false ]; then
# image 模式默认走 push,符合发布场景
IMAGE_PUSH=true
fi
if [ -z "$IMAGE_TAG" ]; then
IMAGE_TAG=$(git -C "$PROJECT_DIR" describe --tags --always 2>/dev/null | sed 's/^v//')
if [ -z "$IMAGE_TAG" ]; then
IMAGE_TAG=$(date +%Y%m%d%H%M%S)
fi
fi
if [ "$IMAGE_LOAD" = true ] && [[ "$IMAGE_PLATFORMS" == *,* ]]; then
echo "❌ --load 仅支持单平台,请用 --platforms linux/amd64(或 arm64)"
exit 1
fi
local ref="${IMAGE_NAME}:${IMAGE_TAG}"
local cmd=(docker buildx build
--platform "$IMAGE_PLATFORMS"
-f "$SCRIPT_DIR/Dockerfile"
-t "$ref"
)
if [ "$IMAGE_LATEST" = true ]; then
cmd+=(-t "${IMAGE_NAME}:latest")
fi
if [ "$IMAGE_PUSH" = true ]; then
cmd+=(--push)
else
cmd+=(--load)
fi
cmd+=("$SCRIPT_DIR")
echo "🔨 开始构建 aether-hub 镜像..."
echo " image: $ref"
echo " platforms: $IMAGE_PLATFORMS"
echo " mode: $([ "$IMAGE_PUSH" = true ] && echo push || echo load)"
echo ""
"${cmd[@]}"
if [ "$IMAGE_PUSH" = true ]; then
echo "✅ 镜像已推送: $ref"
if [ "$IMAGE_LATEST" = true ]; then
echo "✅ 镜像已推送: ${IMAGE_NAME}:latest"
fi
else
echo "✅ 镜像已加载到本地: $ref"
fi
echo "🎉 image 模式完成!"
}
if [ "$MODE" = "image" ]; then
build_image
else
build_binary
fi
+85
View File
@@ -0,0 +1,85 @@
use reqwest::Client;
#[derive(Clone)]
pub struct ControlPlaneClient {
client: Option<Client>,
base_url: String,
}
impl ControlPlaneClient {
pub fn new(base_url: String) -> Self {
let client = Client::builder()
.timeout(std::time::Duration::from_secs(10))
.build()
.ok();
Self { client, base_url }
}
pub fn disabled() -> Self {
Self {
client: None,
base_url: String::new(),
}
}
pub async fn heartbeat_ack(&self, payload: &[u8]) -> Result<Vec<u8>, String> {
let Some(client) = &self.client else {
return Ok(b"{}".to_vec());
};
let url = format!(
"{}/api/internal/hub/heartbeat",
self.base_url.trim_end_matches('/')
);
let response = client
.post(&url)
.header("content-type", "application/json")
.body(payload.to_vec())
.send()
.await
.map_err(|e| format!("heartbeat callback request failed: {e}"))?;
if !response.status().is_success() {
return Err(format!(
"heartbeat callback failed with status {}",
response.status()
));
}
response
.bytes()
.await
.map(|bytes| bytes.to_vec())
.map_err(|e| format!("heartbeat callback body read failed: {e}"))
}
pub async fn push_node_status(
&self,
node_id: &str,
connected: bool,
conn_count: usize,
) -> Result<(), String> {
let Some(client) = &self.client else {
return Ok(());
};
let url = format!(
"{}/api/internal/hub/node-status",
self.base_url.trim_end_matches('/')
);
let response = client
.post(&url)
.json(&serde_json::json!({
"node_id": node_id,
"connected": connected,
"conn_count": conn_count,
}))
.send()
.await
.map_err(|e| format!("node-status callback request failed: {e}"))?;
if response.status().is_success() {
Ok(())
} else {
Err(format!(
"node-status callback failed with status {}",
response.status()
))
}
}
}
+878
View File
@@ -0,0 +1,878 @@
use std::collections::HashMap;
use std::sync::atomic::{AtomicBool, AtomicU32, AtomicU64, AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::Duration;
use axum::extract::ws::Message;
use bytes::Bytes;
use dashmap::DashMap;
use parking_lot::{Mutex, RwLock};
use tokio::sync::mpsc;
use tokio::sync::mpsc::error::TrySendError;
use tokio::sync::{watch, Notify};
use tracing::{debug, info, warn};
use crate::control_plane::ControlPlaneClient;
use crate::protocol;
const MAX_REQUEST_BODY_FRAME_SIZE: usize = 32 * 1024;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SendStatus {
Queued,
Closed,
Congested,
}
#[derive(Debug, Clone, Copy)]
pub struct ConnConfig {
pub ping_interval: Duration,
pub idle_timeout: Duration,
pub outbound_queue_capacity: usize,
}
pub struct BoundedOutbound {
tx: mpsc::Sender<Message>,
close_tx: watch::Sender<bool>,
closing: AtomicBool,
}
impl BoundedOutbound {
pub fn new(tx: mpsc::Sender<Message>, close_tx: watch::Sender<bool>) -> Self {
Self {
tx,
close_tx,
closing: AtomicBool::new(false),
}
}
pub fn send(&self, msg: Message) -> SendStatus {
if self.is_closing() {
return SendStatus::Closed;
}
match self.tx.try_send(msg) {
Ok(()) => SendStatus::Queued,
Err(TrySendError::Closed(_)) => {
self.mark_closing();
SendStatus::Closed
}
Err(TrySendError::Full(_)) => {
self.mark_closing();
SendStatus::Congested
}
}
}
pub fn is_closing(&self) -> bool {
self.closing.load(Ordering::Acquire)
}
pub fn mark_closing(&self) -> bool {
if self.closing.swap(true, Ordering::AcqRel) {
return false;
}
let _ = self.close_tx.send(true);
true
}
}
pub struct ProxyConn {
pub id: u64,
pub node_id: String,
pub node_name: String,
pub outbound: BoundedOutbound,
next_stream_id: AtomicU32,
pub stream_count: AtomicUsize,
pub max_streams: usize,
}
impl ProxyConn {
pub fn new(
id: u64,
node_id: String,
node_name: String,
tx: mpsc::Sender<Message>,
close_tx: watch::Sender<bool>,
max_streams: usize,
) -> Self {
Self {
id,
node_id,
node_name,
outbound: BoundedOutbound::new(tx, close_tx),
next_stream_id: AtomicU32::new(2),
stream_count: AtomicUsize::new(0),
max_streams,
}
}
pub fn alloc_stream_id(&self) -> Option<u32> {
let mut current = self.stream_count.load(Ordering::Relaxed);
loop {
if current >= self.max_streams || !self.is_available() {
return None;
}
match self.stream_count.compare_exchange_weak(
current,
current + 1,
Ordering::AcqRel,
Ordering::Relaxed,
) {
Ok(_) => break,
Err(observed) => current = observed,
}
}
let sid = loop {
let current_sid = self.next_stream_id.load(Ordering::Relaxed);
let next_sid = if current_sid >= 0xFFFF_FFFE {
2
} else {
current_sid + 2
};
if self
.next_stream_id
.compare_exchange_weak(current_sid, next_sid, Ordering::AcqRel, Ordering::Relaxed)
.is_ok()
{
break current_sid;
}
};
Some(sid)
}
pub fn release_stream(&self) {
let mut current = self.stream_count.load(Ordering::Relaxed);
while current > 0 {
match self.stream_count.compare_exchange_weak(
current,
current - 1,
Ordering::AcqRel,
Ordering::Relaxed,
) {
Ok(_) => return,
Err(observed) => current = observed,
}
}
}
pub fn is_available(&self) -> bool {
!self.outbound.is_closing()
}
pub fn request_close(&self) {
self.outbound.mark_closing();
}
pub fn send(&self, msg: Message) -> SendStatus {
let was_closing = self.outbound.is_closing();
let status = self.outbound.send(msg);
if status == SendStatus::Congested && !was_closing {
warn!(
conn_id = self.id,
node_id = %self.node_id,
node_name = %self.node_name,
queued_streams = self.stream_count.load(Ordering::Relaxed),
"proxy outbound queue full, closing congested connection"
);
}
status
}
}
#[derive(Debug, Clone)]
pub struct LocalResponseHead {
pub status: u16,
pub headers: Vec<(String, String)>,
}
#[derive(Debug)]
pub enum LocalBodyEvent {
Chunk(Bytes),
End,
Error(String),
}
#[derive(Debug, Default)]
struct LocalWaitState {
response: Option<LocalResponseHead>,
error: Option<String>,
}
pub struct LocalStream {
pub id: u64,
proxy_conn_id: u64,
proxy_stream_id: u32,
wait_state: Mutex<LocalWaitState>,
headers_notify: Notify,
body_tx: mpsc::Sender<LocalBodyEvent>,
body_rx: Mutex<Option<mpsc::Receiver<LocalBodyEvent>>>,
terminal: AtomicBool,
}
impl LocalStream {
fn new(id: u64, proxy_conn_id: u64, proxy_stream_id: u32) -> Self {
let (body_tx, body_rx) = mpsc::channel(128);
Self {
id,
proxy_conn_id,
proxy_stream_id,
wait_state: Mutex::new(LocalWaitState::default()),
headers_notify: Notify::new(),
body_tx,
body_rx: Mutex::new(Some(body_rx)),
terminal: AtomicBool::new(false),
}
}
pub async fn wait_headers(&self, timeout: Duration) -> Result<LocalResponseHead, String> {
tokio::time::timeout(timeout, async {
loop {
let outcome = {
let state = self.wait_state.lock();
if let Some(response) = &state.response {
return Ok(response.clone());
}
state.error.clone()
};
if let Some(error) = outcome {
return Err(error);
}
self.headers_notify.notified().await;
}
})
.await
.map_err(|_| "timed out waiting for response headers".to_string())?
}
pub fn take_body_receiver(&self) -> Option<mpsc::Receiver<LocalBodyEvent>> {
self.body_rx.lock().take()
}
fn set_response_headers(&self, meta: protocol::ResponseMeta) {
let mut notify = false;
{
let mut state = self.wait_state.lock();
if state.response.is_none() && state.error.is_none() {
state.response = Some(LocalResponseHead {
status: meta.status,
headers: meta.headers,
});
notify = true;
}
}
if notify {
self.headers_notify.notify_waiters();
}
}
fn push_body_chunk(&self, payload: Bytes) -> bool {
if self.terminal.load(Ordering::Acquire) {
return false;
}
self.body_tx
.try_send(LocalBodyEvent::Chunk(payload))
.is_ok()
}
fn finish(&self) {
if self.terminal.swap(true, Ordering::AcqRel) {
return;
}
let mut notify = false;
{
let mut state = self.wait_state.lock();
if state.response.is_none() && state.error.is_none() {
state.error = Some("stream ended before response headers".to_string());
notify = true;
}
}
if notify {
self.headers_notify.notify_waiters();
}
let _ = self.body_tx.try_send(LocalBodyEvent::End);
}
fn fail(&self, error: impl Into<String>) {
if self.terminal.swap(true, Ordering::AcqRel) {
return;
}
let error = error.into();
let mut notify = false;
{
let mut state = self.wait_state.lock();
if state.response.is_none() && state.error.is_none() {
state.error = Some(error.clone());
notify = true;
}
}
if notify {
self.headers_notify.notify_waiters();
}
let _ = self.body_tx.try_send(LocalBodyEvent::Error(error));
}
}
pub struct HubRouter {
proxy_conns: RwLock<HashMap<String, Vec<Arc<ProxyConn>>>>,
proxy_conns_by_id: DashMap<u64, Arc<ProxyConn>>,
local_streams: DashMap<u64, Arc<LocalStream>>,
proxy_to_local: DashMap<(u64, u32), u64>,
next_conn_id: AtomicU64,
next_local_stream_id: AtomicU64,
control_plane: ControlPlaneClient,
}
impl HubRouter {
pub fn new(control_plane: ControlPlaneClient) -> Arc<Self> {
Arc::new(Self {
proxy_conns: RwLock::new(HashMap::new()),
proxy_conns_by_id: DashMap::new(),
local_streams: DashMap::new(),
proxy_to_local: DashMap::new(),
next_conn_id: AtomicU64::new(1),
next_local_stream_id: AtomicU64::new(1),
control_plane,
})
}
pub fn alloc_conn_id(&self) -> u64 {
self.next_conn_id.fetch_add(1, Ordering::Relaxed)
}
pub fn register_proxy(&self, conn: Arc<ProxyConn>) {
let node_id = conn.node_id.clone();
let node_name = conn.node_name.clone();
let conn_id = conn.id;
self.proxy_conns_by_id.insert(conn_id, conn.clone());
let pool_size = {
let mut map = self.proxy_conns.write();
map.entry(node_id.clone()).or_default().push(conn);
map.get(&node_id).map(|v| v.len()).unwrap_or(0)
};
info!(
node_id = %node_id,
node_name = %node_name,
conn_id = conn_id,
pool_size = pool_size,
"proxy connected"
);
self.notify_node_status(node_id, true, pool_size);
}
pub fn unregister_proxy(&self, conn_id: u64, node_id: &str) {
self.proxy_conns_by_id.remove(&conn_id);
let pool_size = {
let mut map = self.proxy_conns.write();
if let Some(conns) = map.get_mut(node_id) {
conns.retain(|c| c.id != conn_id);
if conns.is_empty() {
map.remove(node_id);
}
}
map.get(node_id).map(|v| v.len()).unwrap_or(0)
};
info!(
node_id = %node_id,
conn_id = conn_id,
remaining = pool_size,
"proxy disconnected"
);
self.cancel_streams_for_proxy(conn_id);
self.notify_node_status(node_id.to_string(), pool_size > 0, pool_size);
}
fn notify_node_status(&self, node_id: String, connected: bool, conn_count: usize) {
let control_plane = self.control_plane.clone();
tokio::spawn(async move {
if let Err(error) = control_plane
.push_node_status(&node_id, connected, conn_count)
.await
{
warn!(
node_id = %node_id,
connected = connected,
conn_count = conn_count,
error = %error,
"failed to push node status to app control plane"
);
}
});
}
fn get_proxy_conn(&self, node_id: &str) -> Option<Arc<ProxyConn>> {
let map = self.proxy_conns.read();
let conns = map.get(node_id)?;
conns
.iter()
.filter(|c| c.is_available())
.min_by_key(|c| c.stream_count.load(Ordering::Relaxed))
.cloned()
}
pub fn open_local_stream(
&self,
node_id: &str,
meta: &protocol::RequestMeta,
) -> Result<Arc<LocalStream>, String> {
let proxy_conn = self
.get_proxy_conn(node_id)
.ok_or_else(|| format!("no proxy connection for node {node_id}"))?;
let proxy_stream_id = proxy_conn
.alloc_stream_id()
.ok_or_else(|| format!("stream limit reached for node {node_id}"))?;
// Encode frames before registering the stream so that encoding failures
// (practically impossible but theoretically possible) don't leak a stream
// slot or orphan map entries.
let meta_json = match serde_json::to_vec(meta) {
Ok(json) => json,
Err(e) => {
proxy_conn.release_stream();
return Err(format!("failed to encode request metadata: {e}"));
}
};
let (meta_payload, meta_flags) = match protocol::compress_payload(&meta_json) {
Ok(result) => result,
Err(e) => {
proxy_conn.release_stream();
return Err(format!("failed to compress request metadata: {e}"));
}
};
let header_frame = protocol::encode_frame(
proxy_stream_id,
protocol::REQUEST_HEADERS,
meta_flags,
&meta_payload,
);
// Frames encoded successfully -- now register the stream.
let local_stream_id = self.next_local_stream_id.fetch_add(1, Ordering::Relaxed);
let local_stream = Arc::new(LocalStream::new(
local_stream_id,
proxy_conn.id,
proxy_stream_id,
));
self.local_streams
.insert(local_stream_id, local_stream.clone());
self.proxy_to_local
.insert((proxy_conn.id, proxy_stream_id), local_stream_id);
match proxy_conn.send(Message::Binary(header_frame.into())) {
SendStatus::Queued => Ok(local_stream),
SendStatus::Closed | SendStatus::Congested => {
self.cleanup_local_stream(local_stream_id);
proxy_conn.release_stream();
Err("proxy connection congested".to_string())
}
}
}
pub fn push_local_request_body(
&self,
local_stream_id: u64,
payload: Bytes,
end_stream: bool,
) -> Result<(), String> {
let stream = self
.local_streams
.get(&local_stream_id)
.map(|entry| entry.value().clone())
.ok_or_else(|| "local stream not found".to_string())?;
let proxy_conn = self
.proxy_conns_by_id
.get(&stream.proxy_conn_id)
.map(|entry| entry.value().clone())
.ok_or_else(|| "proxy connection unavailable".to_string())?;
let total_chunks = payload.len().div_ceil(MAX_REQUEST_BODY_FRAME_SIZE);
if total_chunks == 0 {
if end_stream {
self.send_request_body_frame(&proxy_conn, stream.proxy_stream_id, &[], true)?;
}
} else {
for (index, chunk) in payload.chunks(MAX_REQUEST_BODY_FRAME_SIZE).enumerate() {
let is_last_chunk = index + 1 == total_chunks;
self.send_request_body_frame(
&proxy_conn,
stream.proxy_stream_id,
chunk,
end_stream && is_last_chunk,
)?;
}
}
Ok(())
}
fn send_request_body_frame(
&self,
proxy_conn: &Arc<ProxyConn>,
proxy_stream_id: u32,
payload: &[u8],
end_stream: bool,
) -> Result<(), String> {
let (body_payload, body_flags) = protocol::compress_payload(payload)
.map_err(|e| format!("failed to compress request body: {e}"))?;
let body_frame = protocol::encode_frame(
proxy_stream_id,
protocol::REQUEST_BODY,
body_flags
| if end_stream {
protocol::FLAG_END_STREAM
} else {
0
},
&body_payload,
);
match proxy_conn.send(Message::Binary(body_frame.into())) {
SendStatus::Queued => Ok(()),
SendStatus::Closed | SendStatus::Congested => {
Err("proxy connection congested".to_string())
}
}
}
pub fn cancel_local_stream(&self, local_stream_id: u64, reason: &str) {
let Some((_, stream)) = self.local_streams.remove(&local_stream_id) else {
return;
};
self.proxy_to_local
.remove(&(stream.proxy_conn_id, stream.proxy_stream_id));
if let Some(pc) = self.proxy_conns_by_id.get(&stream.proxy_conn_id) {
pc.release_stream();
let frame = protocol::encode_stream_error(stream.proxy_stream_id, reason);
let _ = pc.send(Message::Binary(frame.into()));
}
stream.fail(reason.to_string());
}
fn cleanup_local_stream(&self, local_stream_id: u64) {
let Some((_, stream)) = self.local_streams.remove(&local_stream_id) else {
return;
};
self.proxy_to_local
.remove(&(stream.proxy_conn_id, stream.proxy_stream_id));
}
pub async fn handle_proxy_frame(&self, proxy_conn_id: u64, data: &mut [u8]) {
let header = match protocol::FrameHeader::parse(data) {
Some(h) => h,
None => return,
};
let expected_len = protocol::HEADER_SIZE + header.payload_len as usize;
if data.len() < expected_len {
return;
}
match header.msg_type {
protocol::RESPONSE_HEADERS => {
self.route_response_headers(proxy_conn_id, header, data);
}
protocol::RESPONSE_BODY => {
self.route_response_body(proxy_conn_id, header, data);
}
protocol::STREAM_END => {
self.finish_proxy_stream(proxy_conn_id, header.stream_id);
}
protocol::STREAM_ERROR => {
let message = protocol::decode_payload(data, &header)
.ok()
.and_then(|payload| String::from_utf8(payload).ok())
.unwrap_or_else(|| "stream error".to_string());
self.fail_proxy_stream(proxy_conn_id, header.stream_id, message);
}
protocol::HEARTBEAT_DATA => {
self.handle_heartbeat(proxy_conn_id, header.stream_id, data, &header)
.await;
}
protocol::PING => {
let payload = protocol::frame_payload_by_header(data, &header).unwrap_or(&[]);
let pong = protocol::encode_pong(payload);
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
let _ = pc.send(Message::Binary(pong.into()));
}
}
protocol::PONG => {}
protocol::GOAWAY => {
warn!(
proxy_conn_id = proxy_conn_id,
"received GOAWAY from proxy connection"
);
}
_ => {
debug!(
msg_type = header.msg_type,
proxy_conn_id = proxy_conn_id,
"unexpected frame type from proxy"
);
}
}
}
fn route_response_headers(
&self,
proxy_conn_id: u64,
header: protocol::FrameHeader,
data: &[u8],
) {
let Some(local_id) = self.lookup_local_stream(proxy_conn_id, header.stream_id) else {
return;
};
let Ok(payload) = protocol::decode_payload(data, &header) else {
self.fail_proxy_stream(
proxy_conn_id,
header.stream_id,
"failed to decode response headers",
);
return;
};
let Ok(meta) = serde_json::from_slice::<protocol::ResponseMeta>(&payload) else {
self.fail_proxy_stream(
proxy_conn_id,
header.stream_id,
"invalid response headers payload",
);
return;
};
if let Some(entry) = self.local_streams.get(&local_id) {
entry.value().set_response_headers(meta);
}
}
fn route_response_body(&self, proxy_conn_id: u64, header: protocol::FrameHeader, data: &[u8]) {
let Some(local_id) = self.lookup_local_stream(proxy_conn_id, header.stream_id) else {
return;
};
let Ok(payload) = protocol::decode_payload(data, &header) else {
self.fail_proxy_stream(
proxy_conn_id,
header.stream_id,
"failed to decode response body",
);
return;
};
let stream = match self.local_streams.get(&local_id) {
Some(entry) => entry.value().clone(),
None => return,
};
if !stream.push_body_chunk(Bytes::from(payload)) {
self.cancel_local_stream(local_id, "local relay response congested");
}
}
fn handle_stream_cleanup(
&self,
proxy_conn_id: u64,
proxy_stream_id: u32,
) -> Option<Arc<LocalStream>> {
let local_id = self
.proxy_to_local
.remove(&(proxy_conn_id, proxy_stream_id))
.map(|(_, local_id)| local_id)?;
let stream = self
.local_streams
.remove(&local_id)
.map(|(_, stream)| stream)?;
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
pc.release_stream();
}
Some(stream)
}
fn finish_proxy_stream(&self, proxy_conn_id: u64, proxy_stream_id: u32) {
if let Some(stream) = self.handle_stream_cleanup(proxy_conn_id, proxy_stream_id) {
stream.finish();
}
}
fn fail_proxy_stream(
&self,
proxy_conn_id: u64,
proxy_stream_id: u32,
error: impl Into<String>,
) {
if let Some(stream) = self.handle_stream_cleanup(proxy_conn_id, proxy_stream_id) {
stream.fail(error.into());
}
}
fn lookup_local_stream(&self, proxy_conn_id: u64, proxy_stream_id: u32) -> Option<u64> {
self.proxy_to_local
.get(&(proxy_conn_id, proxy_stream_id))
.map(|entry| *entry.value())
}
async fn handle_heartbeat(
&self,
proxy_conn_id: u64,
stream_id: u32,
data: &[u8],
header: &protocol::FrameHeader,
) {
let payload = match protocol::decode_payload(data, header) {
Ok(payload) => payload,
Err(error) => {
warn!(proxy_conn_id = proxy_conn_id, error = %error, "failed to decode heartbeat payload");
return;
}
};
let ack_payload = match self.control_plane.heartbeat_ack(&payload).await {
Ok(payload) => payload,
Err(error) => {
warn!(proxy_conn_id = proxy_conn_id, error = %error, "control-plane heartbeat callback failed");
b"{}".to_vec()
}
};
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
let frame = protocol::encode_frame(stream_id, protocol::HEARTBEAT_ACK, 0, &ack_payload);
let _ = pc.send(Message::Binary(frame.into()));
}
}
fn cancel_streams_for_proxy(&self, proxy_conn_id: u64) {
let mut cancelled = 0usize;
self.proxy_to_local.retain(|key, local_id| {
if key.0 != proxy_conn_id {
return true;
}
if let Some((_, stream)) = self.local_streams.remove(local_id) {
stream.fail("proxy disconnected".to_string());
}
cancelled += 1;
false
});
if cancelled > 0 {
warn!(
proxy_conn_id = proxy_conn_id,
streams_cancelled = cancelled,
"cancelled in-flight streams due to proxy disconnect"
);
}
}
pub fn stats(&self) -> HubStats {
let proxy_conns = self.proxy_conns.read();
let total_proxy = proxy_conns.values().map(|v| v.len()).sum();
let nodes = proxy_conns.len();
drop(proxy_conns);
HubStats {
proxy_connections: total_proxy,
nodes,
active_streams: self.local_streams.len(),
}
}
}
#[derive(serde::Serialize)]
pub struct HubStats {
pub proxy_connections: usize,
pub nodes: usize,
pub active_streams: usize,
}
#[cfg(test)]
mod tests {
use super::*;
fn build_meta() -> protocol::RequestMeta {
protocol::RequestMeta {
method: "GET".to_string(),
url: "https://example.com".to_string(),
headers: HashMap::new(),
timeout: 30,
}
}
#[tokio::test]
async fn cancel_local_stream_notifies_proxy() {
let hub = HubRouter::new(ControlPlaneClient::disabled());
let (proxy_tx, mut proxy_rx) = mpsc::channel(8);
let (proxy_close_tx, _) = watch::channel(false);
let proxy = Arc::new(ProxyConn::new(
100,
"node-1".to_string(),
"Node 1".to_string(),
proxy_tx,
proxy_close_tx,
16,
));
hub.register_proxy(proxy);
let stream = hub
.open_local_stream("node-1", &build_meta())
.expect("open local stream");
let _ = proxy_rx.try_recv().expect("headers frame");
hub.push_local_request_body(stream.id, Bytes::new(), true)
.expect("finish empty body");
let _ = proxy_rx.try_recv().expect("body frame");
hub.cancel_local_stream(stream.id, "client dropped");
let cancelled = proxy_rx.try_recv().expect("cancel frame");
let cancelled_data = match cancelled {
Message::Binary(data) => data.to_vec(),
other => panic!("unexpected message: {other:?}"),
};
let header = protocol::FrameHeader::parse(&cancelled_data).expect("cancel frame header");
assert_eq!(header.msg_type, protocol::STREAM_ERROR);
}
#[tokio::test]
async fn push_local_request_body_splits_large_payload_and_marks_end() {
let hub = HubRouter::new(ControlPlaneClient::disabled());
let (proxy_tx, mut proxy_rx) = mpsc::channel(8);
let (proxy_close_tx, _) = watch::channel(false);
let proxy = Arc::new(ProxyConn::new(
200,
"node-2".to_string(),
"Node 2".to_string(),
proxy_tx,
proxy_close_tx,
16,
));
hub.register_proxy(proxy);
let stream = hub
.open_local_stream("node-2", &build_meta())
.expect("open local stream");
let _ = proxy_rx.try_recv().expect("headers frame");
let payload = Bytes::from(vec![b'x'; MAX_REQUEST_BODY_FRAME_SIZE + 17]);
hub.push_local_request_body(stream.id, payload, true)
.expect("push request body");
let first = match proxy_rx.try_recv().expect("first body frame") {
Message::Binary(data) => data.to_vec(),
other => panic!("unexpected message: {other:?}"),
};
let first_header = protocol::FrameHeader::parse(&first).expect("first body header");
assert_eq!(first_header.msg_type, protocol::REQUEST_BODY);
assert_eq!(first_header.flags & protocol::FLAG_END_STREAM, 0);
let second = match proxy_rx.try_recv().expect("second body frame") {
Message::Binary(data) => data.to_vec(),
other => panic!("unexpected message: {other:?}"),
};
let second_header = protocol::FrameHeader::parse(&second).expect("second body header");
assert_eq!(second_header.msg_type, protocol::REQUEST_BODY);
assert_ne!(second_header.flags & protocol::FLAG_END_STREAM, 0);
}
}
+262
View File
@@ -0,0 +1,262 @@
use std::io;
use std::net::SocketAddr;
use std::time::Duration;
use async_stream::stream;
use axum::body::{Body, Bytes};
use axum::extract::{ConnectInfo, Path, Request, State};
use axum::http::{HeaderMap, HeaderName, HeaderValue, Response, StatusCode};
use axum::response::IntoResponse;
use bytes::BytesMut;
use futures_util::StreamExt;
use tracing::warn;
use crate::hub::{LocalBodyEvent, LocalStream};
use crate::protocol;
use crate::AppState;
pub const TUNNEL_ERROR_HEADER: &str = "x-aether-tunnel-error";
const MAX_RELAY_META_LEN: usize = 256 * 1024;
struct StreamGuard {
hub: std::sync::Arc<crate::hub::HubRouter>,
stream_id: u64,
finished: bool,
}
impl Drop for StreamGuard {
fn drop(&mut self) {
if !self.finished {
self.hub
.cancel_local_stream(self.stream_id, "local relay client dropped");
}
}
}
pub async fn relay_request(
Path(node_id): Path<String>,
State(state): State<AppState>,
ConnectInfo(addr): ConnectInfo<SocketAddr>,
request: Request,
) -> impl IntoResponse {
if !addr.ip().is_loopback() {
return tunnel_error_response(
StatusCode::FORBIDDEN,
"forbidden",
"local relay only accepts loopback requests",
);
}
let mut body_stream = request.into_body().into_data_stream();
let mut envelope_buf = BytesMut::new();
let mut meta: Option<protocol::RequestMeta> = None;
let mut stream: Option<std::sync::Arc<LocalStream>> = None;
while let Some(chunk_result) = body_stream.next().await {
let chunk = match chunk_result {
Ok(chunk) => chunk,
Err(error) => {
if let Some(active_stream) = &stream {
state
.hub
.cancel_local_stream(active_stream.id, "failed to read relay request body");
}
warn!(error = %error, "failed to read local relay request body");
return tunnel_error_response(
StatusCode::BAD_GATEWAY,
"relay",
"failed to read relay request body",
);
}
};
if stream.is_none() {
envelope_buf.extend_from_slice(&chunk);
let Some((parsed_meta, body_offset)) = (match try_decode_envelope_meta(&envelope_buf) {
Ok(result) => result,
Err(error) => {
return tunnel_error_response(StatusCode::BAD_REQUEST, "bad_request", &error);
}
}) else {
continue;
};
let opened_stream = match state.hub.open_local_stream(&node_id, &parsed_meta) {
Ok(stream) => stream,
Err(error) => {
return tunnel_error_response(
StatusCode::SERVICE_UNAVAILABLE,
"connect",
&error,
);
}
};
if envelope_buf.len() > body_offset {
let first_body_chunk = Bytes::copy_from_slice(&envelope_buf[body_offset..]);
if let Err(error) =
state
.hub
.push_local_request_body(opened_stream.id, first_body_chunk, false)
{
state.hub.cancel_local_stream(opened_stream.id, &error);
return tunnel_error_response(
StatusCode::SERVICE_UNAVAILABLE,
"connect",
&error,
);
}
}
envelope_buf.clear();
meta = Some(parsed_meta);
stream = Some(opened_stream);
continue;
}
let Some(active_stream) = &stream else {
continue;
};
if let Err(error) = state
.hub
.push_local_request_body(active_stream.id, chunk, false)
{
state.hub.cancel_local_stream(active_stream.id, &error);
return tunnel_error_response(StatusCode::SERVICE_UNAVAILABLE, "connect", &error);
}
}
let (meta, stream) = match (meta, stream) {
(Some(meta), Some(stream)) => (meta, stream),
_ => {
return tunnel_error_response(
StatusCode::BAD_REQUEST,
"bad_request",
"relay envelope metadata truncated",
);
}
};
if let Err(error) = state
.hub
.push_local_request_body(stream.id, Bytes::new(), true)
{
state.hub.cancel_local_stream(stream.id, &error);
return tunnel_error_response(StatusCode::SERVICE_UNAVAILABLE, "connect", &error);
}
let request_guard = StreamGuard {
hub: state.hub.clone(),
stream_id: stream.id,
finished: false,
};
let wait_timeout = Duration::from_secs(meta.timeout.clamp(5, 300));
let response_head = match stream.wait_headers(wait_timeout).await {
Ok(response) => response,
Err(error) => {
state.hub.cancel_local_stream(stream.id, &error);
return tunnel_error_response(StatusCode::GATEWAY_TIMEOUT, "timeout", &error);
}
};
let Some(mut body_rx) = stream.take_body_receiver() else {
state
.hub
.cancel_local_stream(stream.id, "missing relay response body receiver");
return tunnel_error_response(
StatusCode::BAD_GATEWAY,
"relay",
"missing relay response body receiver",
);
};
let hub = state.hub.clone();
let stream_id = stream.id;
let body_stream = stream! {
let mut guard = request_guard;
guard.hub = hub;
guard.stream_id = stream_id;
while let Some(event) = body_rx.recv().await {
match event {
LocalBodyEvent::Chunk(chunk) => yield Ok::<Bytes, io::Error>(chunk),
LocalBodyEvent::End => {
guard.finished = true;
break;
}
LocalBodyEvent::Error(error) => {
guard.finished = true;
yield Err(io::Error::other(error));
break;
}
}
}
guard.finished = true;
};
let mut builder = Response::builder().status(response_head.status);
if let Some(headers) = builder.headers_mut() {
append_headers(headers, &response_head.headers);
}
match builder.body(Body::from_stream(body_stream)) {
Ok(response) => response,
Err(error) => {
warn!(error = %error, "failed to build relay response");
tunnel_error_response(
StatusCode::BAD_GATEWAY,
"relay",
"failed to build relay response",
)
}
}
}
fn try_decode_envelope_meta(
buffer: &BytesMut,
) -> Result<Option<(protocol::RequestMeta, usize)>, String> {
if buffer.len() < 4 {
return Ok(None);
}
let meta_len = u32::from_be_bytes([buffer[0], buffer[1], buffer[2], buffer[3]]) as usize;
if meta_len > MAX_RELAY_META_LEN {
return Err("relay metadata too large".to_string());
}
let meta_end = 4usize
.checked_add(meta_len)
.ok_or_else(|| "relay envelope length overflow".to_string())?;
if buffer.len() < meta_end {
return Ok(None);
}
let meta = serde_json::from_slice::<protocol::RequestMeta>(&buffer[4..meta_end])
.map_err(|e| format!("invalid relay metadata: {e}"))?;
Ok(Some((meta, meta_end)))
}
fn append_headers(target: &mut HeaderMap, headers: &[(String, String)]) {
for (name, value) in headers {
let Ok(name) = HeaderName::from_bytes(name.as_bytes()) else {
continue;
};
let Ok(value) = HeaderValue::from_str(value) else {
continue;
};
target.append(name, value);
}
}
fn tunnel_error_response(status: StatusCode, kind: &str, message: &str) -> Response<Body> {
let mut builder = Response::builder().status(status);
if let Some(headers) = builder.headers_mut() {
headers.insert(
HeaderName::from_static(TUNNEL_ERROR_HEADER),
HeaderValue::from_str(kind).unwrap_or_else(|_| HeaderValue::from_static("relay")),
);
headers.insert(
axum::http::header::CONTENT_TYPE,
HeaderValue::from_static("text/plain; charset=utf-8"),
);
}
builder
.body(Body::from(message.to_string()))
.unwrap_or_else(|_| Response::new(Body::from("relay error")))
}
+167
View File
@@ -0,0 +1,167 @@
mod control_plane;
mod hub;
mod local_relay;
mod protocol;
mod proxy_conn;
use std::net::SocketAddr;
use std::time::Duration;
use axum::extract::ws::WebSocketUpgrade;
use axum::extract::State;
use axum::response::{IntoResponse, Json};
use axum::routing::{get, post};
use axum::Router;
use clap::Parser;
use tracing::{info, warn};
use crate::control_plane::ControlPlaneClient;
use crate::hub::{ConnConfig, HubRouter};
use crate::local_relay::relay_request;
#[derive(Parser, Debug)]
#[command(name = "aether-hub", about = "Tunnel Hub for Aether")]
struct Args {
/// Bind address
#[arg(long, default_value = "0.0.0.0:8085", env = "TUNNEL_HUB_BIND")]
bind: String,
/// Proxy-side idle timeout in seconds (0 to disable)
#[arg(long, default_value_t = 0, env = "TUNNEL_HUB_PROXY_IDLE_TIMEOUT")]
proxy_idle_timeout: u64,
/// Ping interval in seconds (for both sides)
#[arg(long, default_value_t = 15, env = "TUNNEL_HUB_PING_INTERVAL")]
ping_interval: u64,
/// Max concurrent streams per proxy connection
#[arg(long, default_value_t = 2048, env = "TUNNEL_HUB_MAX_STREAMS")]
max_streams: usize,
/// Per-connection outbound queue capacity before treating the socket as congested
#[arg(
long,
default_value_t = 128,
env = "TUNNEL_HUB_OUTBOUND_QUEUE_CAPACITY"
)]
outbound_queue_capacity: usize,
/// Local Aether app base URL for control-plane callbacks
#[arg(
long,
default_value = "http://127.0.0.1:8084",
env = "TUNNEL_HUB_APP_BASE_URL"
)]
app_base_url: String,
}
#[derive(Clone)]
pub struct AppState {
pub hub: std::sync::Arc<HubRouter>,
pub proxy_conn_cfg: ConnConfig,
pub max_streams: usize,
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
// Initialize tracing
tracing_subscriber::fmt()
.with_env_filter(
tracing_subscriber::EnvFilter::try_from_default_env()
.unwrap_or_else(|_| "aether_hub=info".into()),
)
.init();
let args = Args::parse();
let hub = HubRouter::new(ControlPlaneClient::new(args.app_base_url));
let outbound_queue_capacity = args.outbound_queue_capacity.clamp(8, 4096);
let ping_interval = Duration::from_secs(args.ping_interval);
let state = AppState {
hub,
proxy_conn_cfg: ConnConfig {
ping_interval,
idle_timeout: Duration::from_secs(args.proxy_idle_timeout),
outbound_queue_capacity,
},
max_streams: args.max_streams,
};
let app = Router::new()
.route("/health", get(health))
.route("/stats", get(stats))
.route("/proxy", get(ws_proxy))
.route("/local/relay/{node_id}", post(relay_request))
.with_state(state);
let listener = tokio::net::TcpListener::bind(&args.bind).await?;
info!(bind = %args.bind, "aether-hub started");
axum::serve(
listener,
app.into_make_service_with_connect_info::<SocketAddr>(),
)
.await?;
Ok(())
}
// ---------------------------------------------------------------------------
// HTTP endpoints
// ---------------------------------------------------------------------------
async fn health() -> impl IntoResponse {
Json(serde_json::json!({"status": "ok"}))
}
async fn stats(State(state): State<AppState>) -> impl IntoResponse {
Json(state.hub.stats())
}
// ---------------------------------------------------------------------------
// WebSocket endpoints
// ---------------------------------------------------------------------------
async fn ws_proxy(
ws: WebSocketUpgrade,
State(state): State<AppState>,
headers: axum::http::HeaderMap,
) -> impl IntoResponse {
let node_id = headers
.get("x-node-id")
.and_then(|v| v.to_str().ok())
.unwrap_or("")
.trim()
.to_string();
let node_name = headers
.get("x-node-name")
.and_then(|v| v.to_str().ok())
.unwrap_or(&node_id)
.trim()
.to_string();
let max_streams: usize = headers
.get("x-tunnel-max-streams")
.and_then(|v| v.to_str().ok())
.and_then(|v| v.parse().ok())
.unwrap_or(state.max_streams)
.clamp(64, 2048);
if node_id.is_empty() {
warn!("proxy connection rejected: missing X-Node-ID header");
return axum::http::StatusCode::BAD_REQUEST.into_response();
}
ws.max_frame_size(64 * 1024 * 1024)
.on_upgrade(move |socket| {
proxy_conn::handle_proxy_connection(
socket,
state.hub,
node_id,
node_name,
max_streams,
state.proxy_conn_cfg,
)
})
.into_response()
}
+176
View File
@@ -0,0 +1,176 @@
/// Tunnel binary frame protocol
///
/// Frame format (10-byte header + payload):
/// | stream_id (4B) | msg_type (1B) | flags (1B) | payload_len (4B) | payload (NB) |
use std::io::Read;
use flate2::read::GzDecoder;
use flate2::write::GzEncoder;
use flate2::Compression;
pub const HEADER_SIZE: usize = 10;
// Message types
pub const REQUEST_HEADERS: u8 = 0x01;
pub const REQUEST_BODY: u8 = 0x02;
pub const RESPONSE_HEADERS: u8 = 0x03;
pub const RESPONSE_BODY: u8 = 0x04;
pub const STREAM_END: u8 = 0x05;
pub const STREAM_ERROR: u8 = 0x06;
pub const PING: u8 = 0x10;
pub const PONG: u8 = 0x11;
pub const GOAWAY: u8 = 0x12;
pub const HEARTBEAT_DATA: u8 = 0x13;
pub const HEARTBEAT_ACK: u8 = 0x14;
// Flags
pub const FLAG_END_STREAM: u8 = 0x01;
pub const FLAG_GZIP_COMPRESSED: u8 = 0x02;
#[derive(Debug, Clone, Copy)]
pub struct FrameHeader {
pub stream_id: u32,
pub msg_type: u8,
pub flags: u8,
pub payload_len: u32,
}
impl FrameHeader {
/// Parse frame header from raw bytes (must be >= HEADER_SIZE)
#[inline]
pub fn parse(data: &[u8]) -> Option<Self> {
if data.len() < HEADER_SIZE {
return None;
}
Some(Self {
stream_id: u32::from_be_bytes([data[0], data[1], data[2], data[3]]),
msg_type: data[4],
flags: data[5],
payload_len: u32::from_be_bytes([data[6], data[7], data[8], data[9]]),
})
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct RequestMeta {
pub method: String,
pub url: String,
pub headers: std::collections::HashMap<String, String>,
#[serde(default = "default_timeout", deserialize_with = "deserialize_timeout")]
pub timeout: u64,
}
fn default_timeout() -> u64 {
60
}
fn deserialize_timeout<'de, D>(deserializer: D) -> Result<u64, D::Error>
where
D: serde::Deserializer<'de>,
{
#[derive(serde::Deserialize)]
#[serde(untagged)]
enum TimeoutValue {
Int(u64),
Float(f64),
}
match <TimeoutValue as serde::Deserialize>::deserialize(deserializer)? {
TimeoutValue::Int(v) => Ok(v),
TimeoutValue::Float(v) => {
if !v.is_finite() || v < 0.0 {
return Err(serde::de::Error::custom(
"timeout must be a non-negative finite number",
));
}
if v.fract() != 0.0 {
return Err(serde::de::Error::custom("timeout must be integer seconds"));
}
if v > (u64::MAX as f64) {
return Err(serde::de::Error::custom("timeout is too large"));
}
Ok(v as u64)
}
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct ResponseMeta {
pub status: u16,
pub headers: Vec<(String, String)>,
}
pub fn encode_frame(stream_id: u32, msg_type: u8, flags: u8, payload: &[u8]) -> Vec<u8> {
let mut buf = Vec::with_capacity(HEADER_SIZE + payload.len());
buf.extend_from_slice(&stream_id.to_be_bytes());
buf.push(msg_type);
buf.push(flags);
buf.extend_from_slice(&(payload.len() as u32).to_be_bytes());
buf.extend_from_slice(payload);
buf
}
/// Encode a STREAM_ERROR frame for a given stream_id with an error message
pub fn encode_stream_error(stream_id: u32, msg: &str) -> Vec<u8> {
encode_frame(stream_id, STREAM_ERROR, 0, msg.as_bytes())
}
/// Encode a PING frame (stream_id=0)
pub fn encode_ping() -> Vec<u8> {
encode_frame(0, PING, 0, &[])
}
/// Encode a PONG frame (stream_id=0, echo payload)
pub fn encode_pong(payload: &[u8]) -> Vec<u8> {
encode_frame(0, PONG, 0, payload)
}
/// Encode a GOAWAY frame (stream_id=0)
pub fn encode_goaway() -> Vec<u8> {
encode_frame(0, GOAWAY, 0, &[])
}
#[inline]
pub fn frame_payload_by_header<'a>(data: &'a [u8], header: &FrameHeader) -> Option<&'a [u8]> {
let payload_len = header.payload_len as usize;
let end = HEADER_SIZE.checked_add(payload_len)?;
if data.len() < end {
return None;
}
Some(&data[HEADER_SIZE..end])
}
pub fn decode_payload(data: &[u8], header: &FrameHeader) -> Result<Vec<u8>, String> {
let payload = frame_payload_by_header(data, header)
.ok_or_else(|| "incomplete frame payload".to_string())?;
if header.flags & FLAG_GZIP_COMPRESSED != 0 {
let mut decoder = GzDecoder::new(payload);
let mut decoded = Vec::new();
decoder
.read_to_end(&mut decoded)
.map_err(|e| format!("failed to decompress payload: {e}"))?;
Ok(decoded)
} else {
Ok(payload.to_vec())
}
}
pub fn compress_payload(payload: &[u8]) -> Result<(Vec<u8>, u8), std::io::Error> {
maybe_recompress_payload(payload, true)
}
fn maybe_recompress_payload(
payload: &[u8],
prefer_gzip: bool,
) -> Result<(Vec<u8>, u8), std::io::Error> {
if !prefer_gzip {
return Ok((payload.to_vec(), 0));
}
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
std::io::Write::write_all(&mut encoder, payload)?;
let compressed = encoder.finish()?;
if compressed.len() < payload.len() {
Ok((compressed, FLAG_GZIP_COMPRESSED))
} else {
Ok((payload.to_vec(), 0))
}
}
+156
View File
@@ -0,0 +1,156 @@
/// Proxy-side WebSocket connection handler
///
/// Handles the lifecycle of a single aether-proxy connection:
/// accept -> authenticate (headers) -> read loop -> cleanup
use std::sync::Arc;
use std::time::Duration;
use axum::extract::ws::{Message, WebSocket};
use futures_util::{SinkExt, StreamExt};
use tokio::sync::{mpsc, watch};
use tracing::{debug, info, warn};
use crate::hub::{ConnConfig, HubRouter, ProxyConn, SendStatus};
use crate::protocol;
/// Maximum single frame size: 64 MB
const MAX_FRAME_SIZE: usize = 64 * 1024 * 1024;
pub async fn handle_proxy_connection(
ws: WebSocket,
hub: Arc<HubRouter>,
node_id: String,
node_name: String,
max_streams: usize,
cfg: ConnConfig,
) {
let conn_id = hub.alloc_conn_id();
let (mut ws_tx, ws_rx) = ws.split();
let (tx, mut rx) = mpsc::channel::<Message>(cfg.outbound_queue_capacity);
let (close_tx, mut close_rx) = watch::channel(false);
let conn = Arc::new(ProxyConn::new(
conn_id,
node_id.clone(),
node_name.clone(),
tx,
close_tx,
max_streams,
));
hub.register_proxy(conn.clone());
let writer = tokio::spawn(async move {
loop {
tokio::select! {
msg = rx.recv() => match msg {
Some(msg) => {
if ws_tx.send(msg).await.is_err() {
break;
}
}
None => break,
},
changed = close_rx.changed() => {
if changed.is_err() || *close_rx.borrow() {
break;
}
}
}
}
let _ = ws_tx.close().await;
});
let ping_conn = conn.clone();
let ping_interval = cfg.ping_interval;
let ping_task = tokio::spawn(async move {
loop {
tokio::time::sleep(ping_interval).await;
let ping = protocol::encode_ping();
if !matches!(
ping_conn.send(Message::Binary(ping.into())),
SendStatus::Queued
) {
break;
}
}
});
let reader_hub = hub.clone();
let reader_conn = conn.clone();
let reader = tokio::spawn(async move {
run_proxy_reader(ws_rx, reader_hub, reader_conn, cfg.idle_timeout).await;
});
let _ = reader.await;
ping_task.abort();
conn.request_close();
hub.unregister_proxy(conn_id, &node_id);
drop(conn);
tokio::time::sleep(Duration::from_millis(100)).await;
writer.abort();
let _ = writer.await;
}
async fn run_proxy_reader(
mut ws_rx: futures_util::stream::SplitStream<WebSocket>,
hub: Arc<HubRouter>,
conn: Arc<ProxyConn>,
idle_timeout: Duration,
) {
let idle_enabled = !idle_timeout.is_zero();
let mut oversized_count = 0u32;
loop {
let msg = if idle_enabled {
tokio::select! {
msg = ws_rx.next() => msg,
_ = tokio::time::sleep(idle_timeout) => {
warn!(conn_id = conn.id, node_id = %conn.node_id, "proxy idle timeout");
let _ = conn.send(Message::Binary(protocol::encode_goaway().into()));
conn.request_close();
break;
}
}
} else {
ws_rx.next().await
};
match msg {
Some(Ok(Message::Binary(data))) => {
let mut data = data.to_vec();
if data.len() > MAX_FRAME_SIZE {
oversized_count += 1;
warn!(
conn_id = conn.id,
size = data.len(),
"oversized frame from proxy"
);
if oversized_count >= 5 {
warn!(conn_id = conn.id, "too many oversized frames, closing");
conn.request_close();
break;
}
continue;
}
oversized_count = 0;
if data.len() < protocol::HEADER_SIZE {
debug!(conn_id = conn.id, "frame too small, skipping");
continue;
}
hub.handle_proxy_frame(conn.id, &mut data).await;
}
Some(Ok(Message::Close(_))) | None => {
info!(conn_id = conn.id, node_id = %conn.node_id, "proxy WebSocket closed");
break;
}
Some(Err(e)) => {
warn!(conn_id = conn.id, error = %e, "proxy WebSocket error");
break;
}
_ => {}
}
}
}
-6
View File
@@ -4,11 +4,5 @@ AETHER_PROXY_AETHER_URL=https://aether.example.com
# Management Token (ae_xxx, must belong to an ADMIN user)
AETHER_PROXY_MANAGEMENT_TOKEN=ae_xxxxx
# HMAC key (must match Aether's PROXY_HMAC_KEY)
AETHER_PROXY_HMAC_KEY=
# Proxy listen port
AETHER_PROXY_LISTEN_PORT=18080
# Node identification
AETHER_PROXY_NODE_NAME=proxy-01
+120 -72
View File
@@ -10,9 +10,10 @@ checksum = "320119579fcad9c21884f5c4861d16174d0e06250625266f50fe6898340abefa"
[[package]]
name = "aether-proxy"
version = "0.1.4"
version = "0.2.4"
dependencies = [
"anyhow",
"arc-swap",
"base64",
"bytes",
"clap",
@@ -20,32 +21,29 @@ dependencies = [
"flate2",
"futures-util",
"hex",
"hmac",
"http-body-util",
"hyper",
"hyper-util",
"libc",
"ratatui",
"rcgen",
"reqwest",
"rustls",
"rustls-pemfile",
"rustls-pki-types",
"serde",
"serde_json",
"sha2",
"subtle",
"socket2 0.5.10",
"sysinfo",
"tar",
"thiserror 2.0.18",
"tokio",
"tokio-rustls",
"tokio-tungstenite",
"toml",
"tower-service",
"tracing",
"tracing-subscriber",
"url",
"webpki-roots",
"webpki-roots 0.26.11",
]
[[package]]
@@ -119,6 +117,15 @@ version = "1.0.101"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5f0e0fee31ef5ed1ba1316088939cea399010ed7731dba877ed44aeb407a75ea"
[[package]]
name = "arc-swap"
version = "1.8.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f9f3647c145568cec02c42054e07bdf9a5a698e15b466fb2341bfc393cd24aa5"
dependencies = [
"rustversion",
]
[[package]]
name = "atomic"
version = "0.6.1"
@@ -216,6 +223,12 @@ version = "1.25.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c8efb64bd706a16a1bdde310ae86b351e4d21550d98d056f22f8a7f7a2183fec"
[[package]]
name = "byteorder"
version = "1.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1fd0f2584146f6f2ef48085050886acf353beff7305ebd1ae69500e27c67f64b"
[[package]]
name = "bytes"
version = "1.11.1"
@@ -479,6 +492,12 @@ dependencies = [
"syn 2.0.114",
]
[[package]]
name = "data-encoding"
version = "2.10.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d7a1e2f27636f116493b8b860f5546edb47c8d8f8ea73e1d2a20be88e28d1fea"
[[package]]
name = "deltae"
version = "0.3.2"
@@ -524,7 +543,6 @@ checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292"
dependencies = [
"block-buffer",
"crypto-common",
"subtle",
]
[[package]]
@@ -811,15 +829,6 @@ version = "0.4.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7f24254aa9a54b5c858eaee2f5bccdb46aaf0e486a595ed5fd8f86ba55232a70"
[[package]]
name = "hmac"
version = "0.12.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6c49c37c09c17a53d937dfbb742eb3a961d65a994e6bcdcf37e7399d0cc8ab5e"
dependencies = [
"digest",
]
[[package]]
name = "http"
version = "1.4.0"
@@ -859,12 +868,6 @@ version = "1.10.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6dbf3de79e51f3d586ab4cb9d5c3e2c14aa28ed23d180cf89b4df0454a69cc87"
[[package]]
name = "httpdate"
version = "1.0.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "df3b46402a9d5adb4c86a0cf463f42e19994e3ee891101b1841f30a545cb49a9"
[[package]]
name = "hyper"
version = "1.8.1"
@@ -879,7 +882,6 @@ dependencies = [
"http",
"http-body",
"httparse",
"httpdate",
"itoa",
"pin-project-lite",
"pin-utils",
@@ -902,7 +904,7 @@ dependencies = [
"tokio",
"tokio-rustls",
"tower-service",
"webpki-roots",
"webpki-roots 1.0.6",
]
[[package]]
@@ -922,7 +924,7 @@ dependencies = [
"libc",
"percent-encoding",
"pin-project-lite",
"socket2",
"socket2 0.6.2",
"tokio",
"tower-service",
"tracing",
@@ -1416,16 +1418,6 @@ dependencies = [
"windows-link",
]
[[package]]
name = "pem"
version = "3.0.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1d30c53c26bc5b31a98cd02d20f25a7c8567146caf63ed593a9d87b2775291be"
dependencies = [
"base64",
"serde_core",
]
[[package]]
name = "percent-encoding"
version = "2.3.2"
@@ -1591,7 +1583,7 @@ dependencies = [
"quinn-udp",
"rustc-hash",
"rustls",
"socket2",
"socket2 0.6.2",
"thiserror 2.0.18",
"tokio",
"tracing",
@@ -1628,7 +1620,7 @@ dependencies = [
"cfg_aliases",
"libc",
"once_cell",
"socket2",
"socket2 0.6.2",
"tracing",
"windows-sys 0.60.2",
]
@@ -1654,6 +1646,8 @@ version = "0.8.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "34af8d1a0e25924bc5b7c43c079c942339d8f0a8b57c39049bef581b46327404"
dependencies = [
"libc",
"rand_chacha 0.3.1",
"rand_core 0.6.4",
]
@@ -1663,10 +1657,20 @@ version = "0.9.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6db2770f06117d490610c7488547d543617b21bfa07796d7a12f6f1bd53850d1"
dependencies = [
"rand_chacha",
"rand_chacha 0.9.0",
"rand_core 0.9.5",
]
[[package]]
name = "rand_chacha"
version = "0.3.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e6c10a63a0fa32252be49d21e7709d4d4baf8d231c2dbce1eaa8141b9b127d88"
dependencies = [
"ppv-lite86",
"rand_core 0.6.4",
]
[[package]]
name = "rand_chacha"
version = "0.9.0"
@@ -1682,6 +1686,9 @@ name = "rand_core"
version = "0.6.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "ec0be4795e2f6a28069bec0b5ff3e2ac9bafc99e6a9a7dc3547996c5c816922c"
dependencies = [
"getrandom 0.2.17",
]
[[package]]
name = "rand_core"
@@ -1797,19 +1804,6 @@ dependencies = [
"crossbeam-utils",
]
[[package]]
name = "rcgen"
version = "0.13.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "75e669e5202259b5314d1ea5397316ad400819437857b90861765f24c4cf80a2"
dependencies = [
"pem",
"ring",
"rustls-pki-types",
"time",
"yasna",
]
[[package]]
name = "redox_syscall"
version = "0.5.18"
@@ -1896,7 +1890,7 @@ dependencies = [
"wasm-bindgen-futures",
"wasm-streams",
"web-sys",
"webpki-roots",
"webpki-roots 1.0.6",
]
[[package]]
@@ -1970,15 +1964,6 @@ dependencies = [
"zeroize",
]
[[package]]
name = "rustls-pemfile"
version = "2.2.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "dce314e5fee3f39953d46bb63bb8a46d40c2f8fb7cc5a3b6cab2bde9721d6e50"
dependencies = [
"rustls-pki-types",
]
[[package]]
name = "rustls-pki-types"
version = "1.14.0"
@@ -2089,6 +2074,17 @@ dependencies = [
"serde",
]
[[package]]
name = "sha1"
version = "0.10.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e3bf829a2d51ab4a5ddf1352d8470c140cadc8301b2ae1789db023f01cedd6ba"
dependencies = [
"cfg-if",
"cpufeatures",
"digest",
]
[[package]]
name = "sha2"
version = "0.10.9"
@@ -2170,6 +2166,16 @@ version = "1.15.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "67b1b7a3b5fe4f1376887184045fcf45c69e92af734b7aaddc05fb777b6fbd03"
[[package]]
name = "socket2"
version = "0.5.10"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e22376abed350d73dd1cd119b57ffccad95b4e585a7cda43e286245ce23c0678"
dependencies = [
"libc",
"windows-sys 0.52.0",
]
[[package]]
name = "socket2"
version = "0.6.2"
@@ -2462,7 +2468,7 @@ dependencies = [
"parking_lot",
"pin-project-lite",
"signal-hook-registry",
"socket2",
"socket2 0.6.2",
"tokio-macros",
"windows-sys 0.61.2",
]
@@ -2488,6 +2494,22 @@ dependencies = [
"tokio",
]
[[package]]
name = "tokio-tungstenite"
version = "0.24.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "edc5f74e248dc973e0dbb7b74c7e0d6fcc301c694ff50049504004ef4d0cdcd9"
dependencies = [
"futures-util",
"log",
"rustls",
"rustls-pki-types",
"tokio",
"tokio-rustls",
"tungstenite",
"webpki-roots 0.26.11",
]
[[package]]
name = "tokio-util"
version = "0.7.18"
@@ -2667,6 +2689,26 @@ version = "0.2.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e421abadd41a4225275504ea4d6566923418b7f05506fbc9c0fe86ba7396114b"
[[package]]
name = "tungstenite"
version = "0.24.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "18e5b8366ee7a95b16d32197d0b2604b43a0be89dc5fac9f8e96ccafbaedda8a"
dependencies = [
"byteorder",
"bytes",
"data-encoding",
"http",
"httparse",
"log",
"rand 0.8.5",
"rustls",
"rustls-pki-types",
"sha1",
"thiserror 1.0.69",
"utf-8",
]
[[package]]
name = "typenum"
version = "1.19.0"
@@ -2726,6 +2768,12 @@ dependencies = [
"serde",
]
[[package]]
name = "utf-8"
version = "0.7.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "09cc8ee72d2a9becf2f2febe0205bbed8fc6615b7cb429ad062dc7b7ddd036a9"
[[package]]
name = "utf8_iter"
version = "1.0.4"
@@ -2887,6 +2935,15 @@ dependencies = [
"wasm-bindgen",
]
[[package]]
name = "webpki-roots"
version = "0.26.11"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "521bc38abb08001b01866da9f51eb7c5d647a19260e00054a8c7fd5f9e57f7a9"
dependencies = [
"webpki-roots 1.0.6",
]
[[package]]
name = "webpki-roots"
version = "1.0.6"
@@ -3236,15 +3293,6 @@ dependencies = [
"rustix 1.1.3",
]
[[package]]
name = "yasna"
version = "0.5.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e17bb3549cc1321ae1296b9cdc2698e2b6cb1992adfa19a8c72e5b7a738f44cd"
dependencies = [
"time",
]
[[package]]
name = "yoke"
version = "0.8.1"
+12 -14
View File
@@ -1,20 +1,18 @@
[package]
name = "aether-proxy"
version = "0.1.6"
version = "0.2.5"
edition = "2021"
description = "Forward proxy for Aether with HMAC authentication"
description = "Tunnel proxy for Aether"
[dependencies]
tokio = { version = "1", features = ["full"] }
hyper = { version = "1", features = ["http1", "server"] }
hyper-util = { version = "0.1", features = ["tokio", "http1", "http2", "server", "client-legacy"] }
tower-service = "0.3"
http-body-util = "0.1"
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls", "stream", "http2"] }
hyper = { version = "1", features = ["client", "http1", "http2"] }
hyper-util = { version = "0.1", features = ["client", "client-legacy", "http1", "http2", "tokio"] }
http-body-util = "0.1"
tokio-tungstenite = { version = "0.24", features = ["rustls-tls-webpki-roots"] }
tokio-rustls = "0.26"
futures-util = "0.3"
hmac = "0.12"
sha2 = "0.10"
subtle = "2"
base64 = "0.22"
clap = { version = "4", features = ["derive", "env"] }
tracing = "0.1"
@@ -23,15 +21,12 @@ serde = { version = "1", features = ["derive"] }
serde_json = "1"
thiserror = "2"
bytes = "1"
sha2 = "0.10"
hex = "0.4"
anyhow = "1"
arc-swap = "1"
toml = "0.8"
tokio-rustls = "0.26"
webpki-roots = "1"
rustls = { version = "0.23", features = ["ring"] }
rustls-pki-types = "1"
rustls-pemfile = "2"
rcgen = "0.13"
ratatui = "0.30"
crossterm = "0.28"
url = "2"
@@ -39,6 +34,9 @@ sysinfo = "0.32"
libc = "0.2"
flate2 = "1"
tar = "0.4"
socket2 = { version = "0.5", features = ["all"] }
tower-service = "0.3"
webpki-roots = "0.26"
[profile.release]
lto = true
-2
View File
@@ -7,6 +7,4 @@ RUN apt-get update && apt-get install -y --no-install-recommends ca-certificates
COPY build/linux-${TARGETARCH}/aether-proxy /usr/local/bin/aether-proxy
EXPOSE 18080
ENTRYPOINT ["aether-proxy"]
+72 -38
View File
@@ -1,18 +1,16 @@
# aether-proxy
Aether 正向代理节点,部署在海外 VPS 上,为墙内的 Aether 实例中转 API 流量。
Aether Tunnel 代理节点,部署在海外 VPS 上,通过 WebSocket 隧道为 Aether 实例中转 API 流量。
Tunnel 模式下代理节点**无需对外监听端口**,仅需出站连接到 Aether 服务器。
## 安装
### Docker Compose 部署
```bash
# 拉取镜像
docker pull ghcr.io/fawney19/aether-proxy:latest
# 或使用 docker compose
cp .env.example .env
# 编辑 .env 填入 AETHER_PROXY_AETHER_URL, MANAGEMENT_TOKEN, HMAC_KEY
# 编辑 .env 填入 AETHER_PROXY_AETHER_URL 和 AETHER_PROXY_MANAGEMENT_TOKEN
docker compose up -d
```
@@ -21,11 +19,11 @@ docker compose up -d
<!-- DOWNLOAD_TABLE_START -->
| Platform | Download |
|----------|----------|
| Linux x86_64 | [aether-proxy-linux-amd64.tar.gz](https://github.com/fawney19/Aether/releases/download/proxy-v0.1.6/aether-proxy-linux-amd64.tar.gz) |
| Linux ARM64 | [aether-proxy-linux-arm64.tar.gz](https://github.com/fawney19/Aether/releases/download/proxy-v0.1.6/aether-proxy-linux-arm64.tar.gz) |
| macOS x86_64 | [aether-proxy-macos-amd64.tar.gz](https://github.com/fawney19/Aether/releases/download/proxy-v0.1.6/aether-proxy-macos-amd64.tar.gz) |
| macOS ARM64 | [aether-proxy-macos-arm64.tar.gz](https://github.com/fawney19/Aether/releases/download/proxy-v0.1.6/aether-proxy-macos-arm64.tar.gz) |
| Windows x86_64 | [aether-proxy-windows-amd64.zip](https://github.com/fawney19/Aether/releases/download/proxy-v0.1.6/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 -->
## 快速开始
@@ -69,43 +67,79 @@ sudo aether-proxy uninstall
### 参数一览
#### 基础配置
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--aether-url` | `AETHER_PROXY_AETHER_URL` | **必填** | Aether 服务器地址 |
| `--management-token` | `AETHER_PROXY_MANAGEMENT_TOKEN` | **必填** | 管理员 Token(`ae_xxx` 格式) |
| `--hmac-key` | `AETHER_PROXY_HMAC_KEY` | **必填** | HMAC 密钥,需与 Aether 端一致 |
| `--listen-port` | `AETHER_PROXY_LISTEN_PORT` | `18080` | 监听端口 |
| `--public-ip` | `AETHER_PROXY_PUBLIC_IP` | 自动检测 | 公网 IP |
| `--node-name` | `AETHER_PROXY_NODE_NAME` | `proxy-01` | 节点名称标识 |
| `--node-region` | `AETHER_PROXY_NODE_REGION` | 自动检测 | 地区标识 |
| `--heartbeat-interval` | `AETHER_PROXY_HEARTBEAT_INTERVAL` | `30` | 心跳间隔(秒) |
| `--allowed-ports` | `AETHER_PROXY_ALLOWED_PORTS` | `80,443,8080,8443` | 允许代理的目标端口 |
| `--timestamp-tolerance` | `AETHER_PROXY_TIMESTAMP_TOLERANCE` | `300` | HMAC 时间戳容差(秒) |
| `--aether-request-timeout` | `AETHER_PROXY_AETHER_REQUEST_TIMEOUT` | `10` | Aether API 请求总超时(秒) |
| `--aether-connect-timeout` | `AETHER_PROXY_AETHER_CONNECT_TIMEOUT` | `10` | Aether API 建连超时(秒) |
| `--aether-pool-max-idle-per-host` | `AETHER_PROXY_AETHER_POOL_MAX_IDLE_PER_HOST` | `8` | Aether API 每 Host 最大空闲连接数 |
| `--aether-pool-idle-timeout` | `AETHER_PROXY_AETHER_POOL_IDLE_TIMEOUT` | `90` | Aether API 连接池空闲超时(秒) |
| `--aether-tcp-keepalive` | `AETHER_PROXY_AETHER_TCP_KEEPALIVE` | `60` | Aether API TCP keepalive(秒,0 关闭) |
| `--aether-tcp-nodelay` | `AETHER_PROXY_AETHER_TCP_NODELAY` | `true` | Aether API 启用 TCP_NODELAY |
| `--aether-http2` | `AETHER_PROXY_AETHER_HTTP2` | `true` | Aether API 启用 HTTP/2 |
| `--aether-retry-max-attempts` | `AETHER_PROXY_AETHER_RETRY_MAX_ATTEMPTS` | `3` | Aether API 最大重试次数(含首次) |
| `--aether-retry-base-delay-ms` | `AETHER_PROXY_AETHER_RETRY_BASE_DELAY_MS` | `200` | Aether API 重试基础延迟(毫秒) |
| `--aether-retry-max-delay-ms` | `AETHER_PROXY_AETHER_RETRY_MAX_DELAY_MS` | `2000` | Aether API 重试最大延迟(毫秒) |
| `--max-concurrent-connections` | `AETHER_PROXY_MAX_CONCURRENT_CONNECTIONS` | 自动估算 | 最大并发连接数(默认硬件估算) |
| `--connect-timeout` | `AETHER_PROXY_CONNECT_TIMEOUT` | `30` | CONNECT 上游建连超时(秒) |
| `--tls-handshake-timeout` | `AETHER_PROXY_TLS_HANDSHAKE_TIMEOUT` | `10` | TLS 握手超时(秒) |
| `--dns-cache-ttl` | `AETHER_PROXY_DNS_CACHE_TTL` | `60` | DNS 缓存 TTL(秒) |
#### Tunnel 连接
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--tunnel-connections` | `AETHER_PROXY_TUNNEL_CONNECTIONS` | `3` | 到 Aether 的连接池大小 |
| `--tunnel-max-streams` | `AETHER_PROXY_TUNNEL_MAX_STREAMS` | 自动(硬件估算) | 单连接最大并发 stream 数 |
| `--tunnel-connect-timeout-secs` | `AETHER_PROXY_TUNNEL_CONNECT_TIMEOUT_SECS` | `15` | TCP + TLS 握手超时(秒) |
| `--tunnel-tcp-keepalive-secs` | `AETHER_PROXY_TUNNEL_TCP_KEEPALIVE_SECS` | `30` | TCP keepalive 初始延迟(秒) |
| `--tunnel-tcp-nodelay` | `AETHER_PROXY_TUNNEL_TCP_NODELAY` | `true` | 禁用 Nagle 算法 |
| `--tunnel-ping-interval-secs` | `AETHER_PROXY_TUNNEL_PING_INTERVAL_SECS` | `15` | WebSocket Ping 频率(秒) |
| `--tunnel-stale-timeout-secs` | `AETHER_PROXY_TUNNEL_STALE_TIMEOUT_SECS` | `45` | 无数据断连阈值(秒) |
| `--tunnel-reconnect-base-ms` | `AETHER_PROXY_TUNNEL_RECONNECT_BASE_MS` | `500` | 指数退避基础延迟(毫秒) |
| `--tunnel-reconnect-max-ms` | `AETHER_PROXY_TUNNEL_RECONNECT_MAX_MS` | `30000` | 指数退避上限(毫秒) |
#### 上游 HTTP 请求
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--upstream-connect-timeout-secs` | `AETHER_PROXY_UPSTREAM_CONNECT_TIMEOUT_SECS` | `30` | 上游建连超时(秒) |
| `--upstream-pool-max-idle-per-host` | `AETHER_PROXY_UPSTREAM_POOL_MAX_IDLE_PER_HOST` | `64` | 每 Host 最大空闲连接数 |
| `--upstream-pool-idle-timeout-secs` | `AETHER_PROXY_UPSTREAM_POOL_IDLE_TIMEOUT_SECS` | `300` | 连接池空闲超时(秒) |
| `--upstream-tcp-keepalive-secs` | `AETHER_PROXY_UPSTREAM_TCP_KEEPALIVE_SECS` | `60` | TCP keepalive(秒,0 关闭) |
| `--upstream-tcp-nodelay` | `AETHER_PROXY_UPSTREAM_TCP_NODELAY` | `true` | 启用 TCP_NODELAY |
#### Aether API 客户端
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--aether-request-timeout-secs` | `AETHER_PROXY_AETHER_REQUEST_TIMEOUT_SECS` | `10` | 请求总超时(秒) |
| `--aether-connect-timeout-secs` | `AETHER_PROXY_AETHER_CONNECT_TIMEOUT_SECS` | `10` | 建连超时(秒) |
| `--aether-retry-max-attempts` | `AETHER_PROXY_AETHER_RETRY_MAX_ATTEMPTS` | `3` | 最大重试次数 |
#### DNS 与安全
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--dns-cache-ttl-secs` | `AETHER_PROXY_DNS_CACHE_TTL_SECS` | `60` | DNS 缓存 TTL(秒) |
| `--dns-cache-capacity` | `AETHER_PROXY_DNS_CACHE_CAPACITY` | `1024` | DNS 缓存容量(条目数) |
| `--delegate-connect-timeout` | `AETHER_PROXY_DELEGATE_CONNECT_TIMEOUT` | `30` | delegate 上游建连超时(秒) |
| `--delegate-pool-max-idle-per-host` | `AETHER_PROXY_DELEGATE_POOL_MAX_IDLE_PER_HOST` | `64` | delegate 每 Host 最大空闲连接数 |
| `--delegate-pool-idle-timeout` | `AETHER_PROXY_DELEGATE_POOL_IDLE_TIMEOUT` | `300` | delegate 连接池空闲超时(秒) |
| `--delegate-tcp-keepalive` | `AETHER_PROXY_DELEGATE_TCP_KEEPALIVE` | `60` | delegate TCP keepalive(秒,0 关闭) |
| `--delegate-tcp-nodelay` | `AETHER_PROXY_DELEGATE_TCP_NODELAY` | `true` | delegate 启用 TCP_NODELAY |
#### 日志
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--log-level` | `AETHER_PROXY_LOG_LEVEL` | `info` | 日志级别 |
| `--log-json` | `AETHER_PROXY_LOG_JSON` | `false` | JSON 格式日志 |
| `--enable-tls` | `AETHER_PROXY_ENABLE_TLS` | `true` | 启用 TLS |
| `--tls-cert` | `AETHER_PROXY_TLS_CERT` | `aether-proxy-cert.pem` | TLS 证书路径 |
| `--tls-key` | `AETHER_PROXY_TLS_KEY` | `aether-proxy-key.pem` | TLS 私钥路径 |
### 多服务器配置
在 `aether-proxy.toml` 中使用 `[[servers]]` 配置多个 Aether 服务器:
```toml
[[servers]]
aether_url = "https://aether-1.example.com"
management_token = "ae_xxx"
node_name = "jp-proxy-01"
[[servers]]
aether_url = "https://aether-2.example.com"
management_token = "ae_yyy"
node_name = "jp-proxy-02"
```
## 发布新版本
@@ -115,6 +149,6 @@ sudo aether-proxy uninstall
- 更新 README 中的下载链接表格
```bash
git tag proxy-v0.1.0
git push origin proxy-v0.1.0
git tag proxy-v0.2.0
git push origin proxy-v0.2.0
```
-3
View File
@@ -3,12 +3,9 @@ services:
image: ghcr.io/fawney19/aether-proxy:latest
container_name: aether-proxy
restart: unless-stopped
ports:
- "${AETHER_PROXY_LISTEN_PORT:-18080}:18080"
env_file:
- .env
environment:
AETHER_PROXY_LISTEN_PORT: 18080
AETHER_PROXY_LOG_JSON: "true"
logging:
driver: json-file
+243 -96
View File
@@ -1,40 +1,41 @@
//! Application lifecycle: initialization, task orchestration, and shutdown.
//!
//! Extracted from `main.rs` to keep the entry point minimal and consolidate
//! the startup sequence, tracing init, and graceful shutdown logic.
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, Semaphore};
use tracing::{error, info};
use tokio::sync::{watch, Mutex};
use tracing::{error, info, warn};
use crate::config::Config;
use crate::config::{Config, ServerEntry};
use crate::net;
use crate::registration::client::AetherClient;
use crate::runtime::{self, DynamicConfig};
use crate::state::{AppState, ProxyMetrics};
use crate::{hardware, proxy};
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) -> anyhow::Result<()> {
pub async fn run(mut config: Config, servers: Vec<ServerEntry>) -> anyhow::Result<()> {
config.validate()?;
init_tracing(&config);
info!(
version = env!("CARGO_PKG_VERSION"),
port = config.listen_port,
node_name = %config.node_name,
"aether-proxy starting"
server_count = servers.len(),
"aether-proxy starting (tunnel mode)"
);
// Resolve public IP
// 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?,
None => net::detect_public_ip()
.await
.unwrap_or_else(|_| "0.0.0.0".to_string()),
};
info!(public_ip = %public_ip, "using public IP");
// Auto-detect region if not configured
if config.node_region.is_none() {
@@ -43,123 +44,270 @@ pub async fn run(mut config: Config) -> anyhow::Result<()> {
}
}
// Initialize TLS if enabled
let (tls_acceptor, tls_fingerprint) = if config.enable_tls {
let cert_path = std::path::PathBuf::from(&config.tls_cert);
let key_path = std::path::PathBuf::from(&config.tls_key);
proxy::tls::ensure_self_signed_cert(&cert_path, &key_path)?;
let acceptor = proxy::tls::build_tls_acceptor(&cert_path, &key_path)?;
let fingerprint = proxy::tls::cert_sha256_fingerprint(&cert_path)?;
info!(fingerprint = %fingerprint, "TLS enabled");
(Some(acceptor), Some(fingerprint))
} else {
info!("TLS disabled");
(None, None)
};
// Collect hardware info (once at startup)
// Collect hardware info (once at startup, sent during registration)
let hw_info = hardware::collect();
let max_connections_raw = config
.max_concurrent_connections
.unwrap_or(hw_info.estimated_max_concurrency)
.max(1);
let max_connections = usize::try_from(max_connections_raw).unwrap_or(usize::MAX);
// 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!(
max_connections = max_connections_raw,
"connection limit configured"
tunnel_max_streams = auto,
"auto-detected tunnel_max_streams from hardware"
);
}
info!(
max_concurrency = hw_info.estimated_max_concurrency,
"hardware info collected"
);
let connection_semaphore = Arc::new(Semaphore::new(max_connections));
let metrics = Arc::new(ProxyMetrics::new());
let dns_cache = Arc::new(proxy::target_filter::DnsCache::new(
let dns_cache = Arc::new(target_filter::DnsCache::new(
Duration::from_secs(config.dns_cache_ttl_secs),
config.dns_cache_capacity,
));
// Register with Aether
let aether_client = Arc::new(AetherClient::new(&config));
let node_id = aether_client
.register(
// 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,
&public_ip,
config.enable_tls,
tls_fingerprint.as_deref(),
Some(&hw_info),
)
.await?;
&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()));
}
}
}
info!(node_id = %node_id, "node registered");
// Build DynamicConfig before moving config into Arc
let dynamic = Arc::new(RwLock::new(DynamicConfig::from_config(&config)));
// Build delegate HTTP client (for proxy-initiated upstream requests).
let delegate_client = proxy::delegate_client::build_delegate_client(&config);
{
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),
node_id: Arc::new(RwLock::new(node_id)),
dynamic,
aether_client,
hardware_info: Arc::new(hw_info),
public_ip,
tls_fingerprint,
tls_acceptor,
delegate_client,
active_connections: Arc::new(AtomicU64::new(0)),
connection_semaphore,
dns_cache,
metrics,
upstream_client,
tunnel_tls_config,
});
// Shutdown signal channel
let (shutdown_tx, shutdown_rx) = watch::channel(false);
// Start heartbeat task
let heartbeat_handle = {
let state = Arc::clone(&state);
let rx = shutdown_rx.clone();
tokio::spawn(async move {
crate::registration::heartbeat::run(&state, rx).await;
})
};
info!(
active_servers = server_contexts.lock().await.len(),
"running in tunnel mode"
);
// Start proxy server
let server_handle = {
let state = Arc::clone(&state);
// 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();
tokio::spawn(async move {
if let Err(e) = proxy::server::run(&state, rx).await {
error!(error = %e, "proxy server error");
tunnel_handles.push(tokio::spawn(async move {
tunnel::run(&s, &srv, conn_idx, rx).await;
}));
}
}
})
};
// Wait for shutdown signal (SIGTERM or SIGINT)
// 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...");
// Signal all tasks to stop
let _ = shutdown_tx.send(true);
// Graceful unregister (best-effort)
let current_node_id = state.node_id.read().unwrap().clone();
if let Err(e) = state.aether_client.unregister(&current_node_id).await {
error!(error = %e, "unregister failed during shutdown");
// 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 tasks to finish
let _ = tokio::join!(heartbeat_handle, server_handle);
// 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};
@@ -168,7 +316,6 @@ fn init_tracing(config: &Config) {
let (filter_layer, reload_handle) = reload::Layer::new(filter);
// Register log-level hot-reloader
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);
-196
View File
@@ -1,196 +0,0 @@
use base64::Engine;
use hmac::{Hmac, Mac};
use sha2::Sha256;
use subtle::ConstantTimeEq;
use crate::config::Config;
type HmacSha256 = Hmac<Sha256>;
#[derive(Debug)]
pub enum AuthError {
MissingHeader,
InvalidBasicAuth,
InvalidUsername,
InvalidPasswordFormat,
TimestampParseError,
TimestampExpired,
SignatureMismatch,
}
impl std::fmt::Display for AuthError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::MissingHeader => write!(f, "missing Proxy-Authorization header"),
Self::InvalidBasicAuth => write!(f, "invalid Basic auth encoding"),
Self::InvalidUsername => write!(f, "username must be 'hmac'"),
Self::InvalidPasswordFormat => {
write!(f, "password format must be 'timestamp.signature'")
}
Self::TimestampParseError => write!(f, "invalid timestamp"),
Self::TimestampExpired => write!(f, "timestamp outside tolerance window"),
Self::SignatureMismatch => write!(f, "HMAC signature mismatch"),
}
}
}
/// Validate Proxy-Authorization header.
///
/// Expected format: `Basic base64(hmac:{timestamp}.{signature})`
/// where signature = hex(HMAC-SHA256(hmac_key, "{timestamp}"))
///
/// The signature no longer includes `node_id`, eliminating race conditions
/// during re-registration where the Aether server's cached `node_id` could
/// differ from the proxy's freshly assigned `node_id`.
///
/// `timestamp_tolerance` is accepted separately so the caller can supply
/// the value from [`DynamicConfig`](crate::runtime::DynamicConfig) (which
/// may be updated remotely).
pub fn validate_proxy_auth(
proxy_auth_header: Option<&str>,
config: &Config,
timestamp_tolerance: u64,
) -> Result<(), AuthError> {
let header = proxy_auth_header.ok_or(AuthError::MissingHeader)?;
let encoded = header
.strip_prefix("Basic ")
.or_else(|| header.strip_prefix("basic "))
.ok_or(AuthError::InvalidBasicAuth)?;
let decoded_bytes = base64::engine::general_purpose::STANDARD
.decode(encoded.trim())
.map_err(|_| AuthError::InvalidBasicAuth)?;
let decoded = String::from_utf8(decoded_bytes).map_err(|_| AuthError::InvalidBasicAuth)?;
// format: hmac:{timestamp}.{signature}
let (username, password) = decoded.split_once(':').ok_or(AuthError::InvalidBasicAuth)?;
if username != "hmac" {
return Err(AuthError::InvalidUsername);
}
let (timestamp_str, signature_hex) = password
.split_once('.')
.ok_or(AuthError::InvalidPasswordFormat)?;
// Validate timestamp window
let timestamp: u64 = timestamp_str
.parse()
.map_err(|_| AuthError::TimestampParseError)?;
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.expect("system clock before epoch")
.as_secs();
let diff = now.abs_diff(timestamp);
if diff > timestamp_tolerance {
return Err(AuthError::TimestampExpired);
}
// Recompute signature: HMAC-SHA256(key, timestamp)
let mut mac =
HmacSha256::new_from_slice(config.hmac_key.as_bytes()).expect("HMAC accepts any key size");
mac.update(timestamp_str.as_bytes());
let expected = mac.finalize().into_bytes();
let expected_hex = hex::encode(expected);
// Constant-time comparison
let sig_bytes = signature_hex.as_bytes();
let exp_bytes = expected_hex.as_bytes();
if sig_bytes.len() != exp_bytes.len() || sig_bytes.ct_eq(exp_bytes).unwrap_u8() != 1 {
return Err(AuthError::SignatureMismatch);
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
fn make_config() -> Config {
Config {
aether_url: String::new(),
management_token: String::new(),
hmac_key: "test-hmac-key".to_string(),
listen_port: 18080,
public_ip: None,
node_name: "test".to_string(),
node_region: None,
heartbeat_interval: 30,
allowed_ports: vec![80, 443],
timestamp_tolerance: 300,
aether_request_timeout_secs: 10,
aether_connect_timeout_secs: 10,
aether_pool_max_idle_per_host: 8,
aether_pool_idle_timeout_secs: 90,
aether_tcp_keepalive_secs: 60,
aether_tcp_nodelay: true,
aether_http2: true,
aether_retry_max_attempts: 3,
aether_retry_base_delay_ms: 200,
aether_retry_max_delay_ms: 2000,
max_concurrent_connections: None,
connect_timeout_secs: 30,
tls_handshake_timeout_secs: 10,
dns_cache_ttl_secs: 60,
dns_cache_capacity: 1024,
delegate_connect_timeout_secs: 30,
delegate_pool_max_idle_per_host: 64,
delegate_pool_idle_timeout_secs: 300,
delegate_tcp_keepalive_secs: 60,
delegate_tcp_nodelay: true,
log_level: "info".to_string(),
log_json: false,
enable_tls: false,
tls_cert: String::new(),
tls_key: String::new(),
}
}
fn make_valid_auth(config: &Config) -> String {
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs();
let mut mac = HmacSha256::new_from_slice(config.hmac_key.as_bytes()).unwrap();
mac.update(now.to_string().as_bytes());
let sig = hex::encode(mac.finalize().into_bytes());
let cred = format!("hmac:{}.{}", now, sig);
let encoded = base64::engine::general_purpose::STANDARD.encode(cred);
format!("Basic {}", encoded)
}
#[test]
fn test_valid_auth() {
let config = make_config();
let header = make_valid_auth(&config);
assert!(validate_proxy_auth(Some(&header), &config, config.timestamp_tolerance).is_ok());
}
#[test]
fn test_missing_header() {
let config = make_config();
assert!(matches!(
validate_proxy_auth(None, &config, config.timestamp_tolerance),
Err(AuthError::MissingHeader)
));
}
#[test]
fn test_wrong_username() {
let cred = "user:12345.abc";
let encoded = base64::engine::general_purpose::STANDARD.encode(cred);
let header = format!("Basic {}", encoded);
let config = make_config();
assert!(matches!(
validate_proxy_auth(Some(&header), &config, config.timestamp_tolerance),
Err(AuthError::InvalidUsername)
));
}
}
-3
View File
@@ -1,3 +0,0 @@
pub mod hmac;
pub use self::hmac::validate_proxy_auth;
+322 -93
View File
@@ -3,11 +3,41 @@ use std::path::Path;
use clap::Parser;
use serde::{Deserialize, Serialize};
/// Aether forward proxy with HMAC authentication.
/// 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. Registers with Aether, sends heartbeats, and validates
/// incoming proxy requests via HMAC-SHA256 signatures in Basic Auth.
/// 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 {
@@ -19,14 +49,6 @@ pub struct Config {
#[arg(long, env = "AETHER_PROXY_MANAGEMENT_TOKEN")]
pub management_token: String,
/// HMAC-SHA256 key for proxy authentication
#[arg(long, env = "AETHER_PROXY_HMAC_KEY")]
pub hmac_key: String,
/// Port to listen on for proxy connections
#[arg(long, env = "AETHER_PROXY_LISTEN_PORT", default_value_t = 18080)]
pub listen_port: u16,
/// Public IP address of this node (auto-detected if omitted)
#[arg(long, env = "AETHER_PROXY_PUBLIC_IP")]
pub public_ip: Option<String>,
@@ -52,10 +74,6 @@ pub struct Config {
)]
pub allowed_ports: Vec<u16>,
/// Timestamp tolerance window in seconds for HMAC validation
#[arg(long, env = "AETHER_PROXY_TIMESTAMP_TOLERANCE", default_value_t = 300)]
pub timestamp_tolerance: u64,
/// Aether API request timeout in seconds
#[arg(
long,
@@ -128,14 +146,6 @@ pub struct Config {
#[arg(long, env = "AETHER_PROXY_MAX_CONCURRENT_CONNECTIONS")]
pub max_concurrent_connections: Option<u64>,
/// Upstream TCP connect timeout in seconds for CONNECT tunnels
#[arg(long, env = "AETHER_PROXY_CONNECT_TIMEOUT", default_value_t = 30)]
pub connect_timeout_secs: u64,
/// TLS handshake timeout in seconds for incoming TLS connections
#[arg(long, env = "AETHER_PROXY_TLS_HANDSHAKE_TIMEOUT", default_value_t = 10)]
pub tls_handshake_timeout_secs: u64,
/// DNS cache TTL in seconds
#[arg(long, env = "AETHER_PROXY_DNS_CACHE_TTL", default_value_t = 60)]
pub dns_cache_ttl_secs: u64,
@@ -144,45 +154,45 @@ pub struct Config {
#[arg(long, env = "AETHER_PROXY_DNS_CACHE_CAPACITY", default_value_t = 1024)]
pub dns_cache_capacity: usize,
/// Delegate HTTP client connect timeout in seconds
/// Upstream HTTP client connect timeout in seconds
#[arg(
long,
env = "AETHER_PROXY_DELEGATE_CONNECT_TIMEOUT",
env = "AETHER_PROXY_UPSTREAM_CONNECT_TIMEOUT",
default_value_t = 30
)]
pub delegate_connect_timeout_secs: u64,
pub upstream_connect_timeout_secs: u64,
/// Delegate HTTP client max idle connections per host
/// Upstream HTTP client max idle connections per host
#[arg(
long,
env = "AETHER_PROXY_DELEGATE_POOL_MAX_IDLE_PER_HOST",
env = "AETHER_PROXY_UPSTREAM_POOL_MAX_IDLE_PER_HOST",
default_value_t = 64
)]
pub delegate_pool_max_idle_per_host: usize,
pub upstream_pool_max_idle_per_host: usize,
/// Delegate HTTP client idle timeout in seconds
/// Upstream HTTP client idle timeout in seconds
#[arg(
long,
env = "AETHER_PROXY_DELEGATE_POOL_IDLE_TIMEOUT",
env = "AETHER_PROXY_UPSTREAM_POOL_IDLE_TIMEOUT",
default_value_t = 300
)]
pub delegate_pool_idle_timeout_secs: u64,
pub upstream_pool_idle_timeout_secs: u64,
/// Delegate TCP keepalive in seconds (0 disables)
/// Upstream TCP keepalive in seconds (0 disables)
#[arg(
long,
env = "AETHER_PROXY_DELEGATE_TCP_KEEPALIVE",
env = "AETHER_PROXY_UPSTREAM_TCP_KEEPALIVE",
default_value_t = 60
)]
pub delegate_tcp_keepalive_secs: u64,
pub upstream_tcp_keepalive_secs: u64,
/// Delegate TCP_NODELAY
/// Upstream TCP_NODELAY
#[arg(
long,
env = "AETHER_PROXY_DELEGATE_TCP_NODELAY",
env = "AETHER_PROXY_UPSTREAM_TCP_NODELAY",
default_value_t = true
)]
pub delegate_tcp_nodelay: bool,
pub upstream_tcp_nodelay: bool,
/// Log level (trace, debug, info, warn, error)
#[arg(long, env = "AETHER_PROXY_LOG_LEVEL", default_value = "info")]
@@ -192,25 +202,106 @@ pub struct Config {
#[arg(long, env = "AETHER_PROXY_LOG_JSON", default_value_t = false)]
pub log_json: bool,
/// Enable TLS encryption (dual-stack: accepts both HTTP and TLS on same port)
#[arg(long, env = "AETHER_PROXY_ENABLE_TLS", default_value_t = true)]
pub enable_tls: bool,
/// Path to TLS certificate PEM file
/// Tunnel reconnect base delay in milliseconds (used by exponential backoff)
#[arg(
long,
env = "AETHER_PROXY_TLS_CERT",
default_value = "aether-proxy-cert.pem"
env = "AETHER_PROXY_TUNNEL_RECONNECT_BASE_MS",
default_value_t = 500
)]
pub tls_cert: String,
pub tunnel_reconnect_base_ms: u64,
/// Path to TLS private key PEM file
/// Tunnel reconnect max delay in milliseconds (cap for exponential backoff)
#[arg(
long,
env = "AETHER_PROXY_TLS_KEY",
default_value = "aether-proxy-key.pem"
env = "AETHER_PROXY_TUNNEL_RECONNECT_MAX_MS",
default_value_t = 30000
)]
pub tls_key: String,
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>,
}
// ---------------------------------------------------------------------------
@@ -218,7 +309,7 @@ pub struct Config {
// ---------------------------------------------------------------------------
/// Serializable config for TOML file persistence.
/// All fields are optional — only populated values are written.
/// All fields are optional -- only populated values are written.
#[derive(Debug, Default, Serialize, Deserialize)]
pub struct ConfigFile {
#[serde(skip_serializing_if = "Option::is_none")]
@@ -226,10 +317,6 @@ pub struct ConfigFile {
#[serde(skip_serializing_if = "Option::is_none")]
pub management_token: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub hmac_key: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub listen_port: Option<u16>,
#[serde(skip_serializing_if = "Option::is_none")]
pub public_ip: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub node_name: Option<String>,
@@ -240,8 +327,6 @@ pub struct ConfigFile {
#[serde(skip_serializing_if = "Option::is_none")]
pub allowed_ports: Option<Vec<u16>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub timestamp_tolerance: Option<u64>,
#[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>,
@@ -264,33 +349,47 @@ pub struct ConfigFile {
#[serde(skip_serializing_if = "Option::is_none")]
pub max_concurrent_connections: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub connect_timeout_secs: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tls_handshake_timeout_secs: 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 delegate_connect_timeout_secs: Option<u64>,
pub upstream_connect_timeout_secs: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub delegate_pool_max_idle_per_host: Option<usize>,
pub upstream_pool_max_idle_per_host: Option<usize>,
#[serde(skip_serializing_if = "Option::is_none")]
pub delegate_pool_idle_timeout_secs: Option<u64>,
pub upstream_pool_idle_timeout_secs: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub delegate_tcp_keepalive_secs: Option<u64>,
pub upstream_tcp_keepalive_secs: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub delegate_tcp_nodelay: Option<bool>,
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 enable_tls: Option<bool>,
pub tunnel_reconnect_base_ms: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tls_cert: Option<String>,
pub tunnel_reconnect_max_ms: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tls_key: Option<String>,
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 {
@@ -307,6 +406,102 @@ impl ConfigFile {
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
@@ -332,15 +527,30 @@ impl ConfigFile {
}
};
}
set!("AETHER_PROXY_AETHER_URL", self.aether_url);
set!("AETHER_PROXY_MANAGEMENT_TOKEN", self.management_token);
set!("AETHER_PROXY_HMAC_KEY", self.hmac_key);
set!("AETHER_PROXY_LISTEN_PORT", self.listen_port);
// 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", self.node_name);
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_TIMESTAMP_TOLERANCE", self.timestamp_tolerance);
set!(
"AETHER_PROXY_AETHER_REQUEST_TIMEOUT",
self.aether_request_timeout_secs
@@ -379,38 +589,57 @@ impl ConfigFile {
"AETHER_PROXY_MAX_CONCURRENT_CONNECTIONS",
self.max_concurrent_connections
);
set!("AETHER_PROXY_CONNECT_TIMEOUT", self.connect_timeout_secs);
set!(
"AETHER_PROXY_TLS_HANDSHAKE_TIMEOUT",
self.tls_handshake_timeout_secs
);
set!("AETHER_PROXY_DNS_CACHE_TTL", self.dns_cache_ttl_secs);
set!("AETHER_PROXY_DNS_CACHE_CAPACITY", self.dns_cache_capacity);
set!(
"AETHER_PROXY_DELEGATE_CONNECT_TIMEOUT",
self.delegate_connect_timeout_secs
"AETHER_PROXY_UPSTREAM_CONNECT_TIMEOUT",
self.upstream_connect_timeout_secs
);
set!(
"AETHER_PROXY_DELEGATE_POOL_MAX_IDLE_PER_HOST",
self.delegate_pool_max_idle_per_host
"AETHER_PROXY_UPSTREAM_POOL_MAX_IDLE_PER_HOST",
self.upstream_pool_max_idle_per_host
);
set!(
"AETHER_PROXY_DELEGATE_POOL_IDLE_TIMEOUT",
self.delegate_pool_idle_timeout_secs
"AETHER_PROXY_UPSTREAM_POOL_IDLE_TIMEOUT",
self.upstream_pool_idle_timeout_secs
);
set!(
"AETHER_PROXY_DELEGATE_TCP_KEEPALIVE",
self.delegate_tcp_keepalive_secs
"AETHER_PROXY_UPSTREAM_TCP_KEEPALIVE",
self.upstream_tcp_keepalive_secs
);
set!(
"AETHER_PROXY_DELEGATE_TCP_NODELAY",
self.delegate_tcp_nodelay
"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_ENABLE_TLS", self.enable_tls);
set!("AETHER_PROXY_TLS_CERT", self.tls_cert);
set!("AETHER_PROXY_TLS_KEY", self.tls_key);
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 {
+36 -7
View File
@@ -1,13 +1,14 @@
mod app;
mod auth;
mod config;
mod hardware;
mod net;
mod proxy;
mod registration;
mod runtime;
mod setup;
mod state;
mod target_filter;
mod tunnel;
mod upstream_client;
use std::path::PathBuf;
@@ -56,8 +57,13 @@ async fn main() -> anyhow::Result<()> {
// Load config file as env-var defaults (before clap parsing)
let config_file_path =
std::env::var("AETHER_PROXY_CONFIG").unwrap_or_else(|_| DEFAULT_CONFIG.to_string());
if std::path::Path::new(&config_file_path).exists() {
if let Ok(file_cfg) = config::ConfigFile::load(std::path::Path::new(&config_file_path)) {
let config_path = std::path::Path::new(&config_file_path);
if config_path.exists() {
// Migrate legacy 0.1.x config to 0.2.0 format if needed
if let Err(e) = config::ConfigFile::migrate_legacy(config_path) {
eprintln!(" WARNING: config migration failed: {}", e);
}
if let Ok(file_cfg) = config::ConfigFile::load(config_path) {
file_cfg.inject_env();
}
}
@@ -130,10 +136,33 @@ async fn run_proxy(config: Config) -> anyhow::Result<()> {
// Skip this check when we ARE the systemd service (INVOCATION_ID is set by systemd).
if std::env::var_os("INVOCATION_ID").is_none() && setup::service::is_service_active() {
eprintln!("Warning: systemd service is already running.");
eprintln!("Use `aether-proxy stop` to stop it first, or manage via subcommands:");
eprintln!(" aether-proxy status / logs / restart / stop");
eprintln!("Use `./aether-proxy stop` to stop it first, or manage via subcommands:");
eprintln!(" ./aether-proxy status / logs / restart / stop");
std::process::exit(1);
}
app::run(config).await
// Resolve server list: prefer [[servers]] from TOML, fall back to CLI/env single server.
let config_path =
std::env::var("AETHER_PROXY_CONFIG").unwrap_or_else(|_| DEFAULT_CONFIG.to_string());
let servers = if std::path::Path::new(&config_path).exists() {
config::ConfigFile::load(std::path::Path::new(&config_path))
.ok()
.map(|f| f.effective_servers())
.filter(|s| !s.is_empty())
.unwrap_or_else(|| {
vec![config::ServerEntry {
aether_url: config.aether_url.clone(),
management_token: config.management_token.clone(),
node_name: None,
}]
})
} else {
vec![config::ServerEntry {
aether_url: config.aether_url.clone(),
management_token: config.management_token.clone(),
node_name: None,
}]
};
app::run(config, servers).await
}
-163
View File
@@ -1,163 +0,0 @@
use std::collections::HashSet;
use std::sync::Arc;
use std::time::Duration;
use hyper::body::Incoming;
use hyper::{Request, Response};
use tokio::net::TcpStream;
use tokio::time::timeout;
use tracing::{debug, warn};
use crate::auth;
use crate::config::Config;
use crate::proxy::target_filter::{self, DnsCache};
/// Handle HTTP CONNECT tunnel requests.
///
/// Flow: validate auth -> check target filter -> TCP connect -> 200 -> bidirectional copy
pub async fn handle_connect(
req: Request<Incoming>,
config: Arc<Config>,
allowed_ports: &HashSet<u16>,
timestamp_tolerance: u64,
dns_cache: &DnsCache,
) -> Response<http_body_util::Empty<bytes::Bytes>> {
// Extract Proxy-Authorization header
let proxy_auth = req
.headers()
.get("proxy-authorization")
.and_then(|v| v.to_str().ok());
// HMAC authentication
if let Err(e) = auth::validate_proxy_auth(proxy_auth, &config, timestamp_tolerance) {
warn!(error = %e, "CONNECT auth failed");
return proxy_auth_required(&e.to_string());
}
// Parse target host:port from CONNECT URI
let authority = match req.uri().authority() {
Some(auth) => auth.clone(),
None => {
warn!("CONNECT request missing authority");
return bad_request("missing target authority");
}
};
let host = authority.host().to_string();
let port = authority.port_u16().unwrap_or(443);
// Target filter: private IP + port whitelist
let target_addr =
match target_filter::validate_target(&host, port, allowed_ports, dns_cache).await {
Ok(addr) => addr,
Err(e) => {
warn!(host = %host, port, error = %e, "CONNECT target rejected");
return forbidden(&e.to_string());
}
};
debug!(target = %target_addr, "CONNECT tunnel establishing");
// Connect to target
let connect_timeout = Duration::from_secs(config.connect_timeout_secs);
let target_stream = match timeout(connect_timeout, TcpStream::connect(target_addr)).await {
Ok(Ok(s)) => s,
Ok(Err(e)) => {
warn!(target = %target_addr, error = %e, "CONNECT target connection failed");
return bad_gateway(&e.to_string());
}
Err(_) => {
warn!(target = %target_addr, "CONNECT target connection timeout");
return gateway_timeout("connect timeout");
}
};
if let Err(e) = target_stream.set_nodelay(true) {
debug!(target = %target_addr, error = %e, "failed to set TCP_NODELAY");
}
// Respond 200 and upgrade connection to raw TCP tunnel
let target_display = target_addr.to_string();
// Reuse connect_timeout for upgrade: both are connection-phase operations
// and should complete within the same order of magnitude.
let upgrade_timeout = Duration::from_secs(config.connect_timeout_secs);
tokio::task::spawn(async move {
match timeout(upgrade_timeout, hyper::upgrade::on(req)).await {
Ok(Ok(upgraded)) => {
let mut upgraded = hyper_util::rt::TokioIo::new(upgraded);
let mut target = target_stream;
match tokio::io::copy_bidirectional(&mut upgraded, &mut target).await {
Ok((from_client, from_target)) => {
debug!(
target = %target_display,
from_client,
from_target,
"CONNECT tunnel closed"
);
}
Err(e) => {
debug!(target = %target_display, error = %e, "CONNECT tunnel error");
}
}
}
Ok(Err(e)) => {
warn!(target = %target_display, error = %e, "CONNECT upgrade failed");
}
Err(_) => {
warn!(target = %target_display, "CONNECT upgrade timeout");
}
}
});
Response::builder()
.status(200)
.body(http_body_util::Empty::new())
.unwrap()
}
fn proxy_auth_required(msg: &str) -> Response<http_body_util::Empty<bytes::Bytes>> {
Response::builder()
.status(407)
.header("Proxy-Authenticate", "HMAC-SHA256")
.header("Content-Length", "0")
.header("X-Error", msg)
.body(http_body_util::Empty::new())
.unwrap()
}
fn forbidden(msg: &str) -> Response<http_body_util::Empty<bytes::Bytes>> {
Response::builder()
.status(403)
.header("Content-Length", "0")
.header("X-Error", msg)
.body(http_body_util::Empty::new())
.unwrap()
}
fn bad_request(msg: &str) -> Response<http_body_util::Empty<bytes::Bytes>> {
Response::builder()
.status(400)
.header("Content-Length", "0")
.header("X-Error", msg)
.body(http_body_util::Empty::new())
.unwrap()
}
fn bad_gateway(msg: &str) -> Response<http_body_util::Empty<bytes::Bytes>> {
Response::builder()
.status(502)
.header("Content-Length", "0")
.header("X-Error", msg)
.body(http_body_util::Empty::new())
.unwrap()
}
fn gateway_timeout(msg: &str) -> Response<http_body_util::Empty<bytes::Bytes>> {
Response::builder()
.status(504)
.header("Content-Length", "0")
.header("X-Error", msg)
.body(http_body_util::Empty::new())
.unwrap()
}
-389
View File
@@ -1,389 +0,0 @@
use std::collections::HashMap;
use std::collections::HashSet;
use std::error::Error as StdError;
use std::sync::Arc;
use std::time::Instant;
use futures_util::StreamExt;
use http_body_util::{BodyExt, Full, Limited, StreamBody};
use hyper::body::{Frame, Incoming};
use hyper::header::{HeaderName, HeaderValue};
use hyper::{Method, Request, Response, Uri};
use tracing::{debug, warn};
use url::Url;
use super::BoxBody;
use crate::auth;
use crate::config::Config;
use crate::proxy::delegate_client::{ConnectTiming, DelegateClient};
use crate::proxy::target_filter::{self, DnsCache};
/// Handle delegation requests: Aether sends a full request description,
/// and the proxy issues the actual upstream HTTP call using its own TLS stack.
///
/// Endpoint: POST /_aether/delegate
///
/// Wire format: metadata in HTTP headers, upstream body sent directly
/// as HTTP body (optionally gzip-compressed via `Content-Encoding: gzip`).
///
/// Headers:
/// X-Delegate-Method: POST
/// X-Delegate-Url: https://api.anthropic.com/v1/messages
/// X-Delegate-Headers: base64-encoded JSON {"Authorization": "Bearer ...", ...}
/// X-Delegate-Timeout: 30 (accepted but not used — Aether controls timeouts)
/// Content-Encoding: gzip (optional, indicates body is gzip-compressed)
pub async fn handle_delegate(
req: Request<Incoming>,
config: Arc<Config>,
allowed_ports: &HashSet<u16>,
timestamp_tolerance: u64,
dns_cache: &DnsCache,
http_client: &DelegateClient,
) -> Response<BoxBody> {
let total_start = Instant::now();
// ── Auth ──
let auth_header = req
.headers()
.get("authorization")
.and_then(|v| v.to_str().ok());
if let Err(e) = auth::validate_proxy_auth(auth_header, &config, timestamp_tolerance) {
warn!(error = %e, "delegate auth failed");
return error_response(401, "authentication_failed", &e.to_string());
}
let auth_ms = total_start.elapsed().as_millis() as u64;
// ── Parse metadata from headers ──
let meta_start = Instant::now();
let method_str = match req
.headers()
.get("x-delegate-method")
.and_then(|v| v.to_str().ok())
{
Some(m) => m.to_string(),
None => {
warn!("delegate missing X-Delegate-Method");
return error_response(400, "bad_request", "missing X-Delegate-Method header");
}
};
let target_url = match req
.headers()
.get("x-delegate-url")
.and_then(|v| v.to_str().ok())
{
Some(u) => u.to_string(),
None => {
warn!("delegate missing X-Delegate-Url");
return error_response(400, "bad_request", "missing X-Delegate-Url header");
}
};
let upstream_headers: HashMap<String, String> = match req
.headers()
.get("x-delegate-headers")
.and_then(|v| v.to_str().ok())
{
Some(b64) => {
match base64::Engine::decode(&base64::engine::general_purpose::STANDARD, b64) {
Ok(decoded) => match serde_json::from_slice(&decoded) {
Ok(h) => h,
Err(e) => {
warn!(error = %e, "delegate invalid X-Delegate-Headers JSON");
return error_response(
400,
"bad_request",
"invalid X-Delegate-Headers JSON",
);
}
},
Err(e) => {
warn!(error = %e, "delegate invalid X-Delegate-Headers base64");
return error_response(400, "bad_request", "invalid X-Delegate-Headers base64");
}
}
}
None => HashMap::new(),
};
let is_gzip = req
.headers()
.get("content-encoding")
.and_then(|v| v.to_str().ok())
.map(|v| v.eq_ignore_ascii_case("gzip"))
.unwrap_or(false);
let req_content_length: u64 = req
.headers()
.get("content-length")
.and_then(|v| v.to_str().ok())
.and_then(|v| v.parse().ok())
.unwrap_or(0);
let meta_ms = meta_start.elapsed().as_millis() as u64;
// ── Target validation ──
let parsed_url = match Url::parse(&target_url) {
Ok(u) => u,
Err(e) => {
warn!(url = %target_url, error = %e, "delegate invalid target URL");
return error_response(400, "bad_request", &format!("invalid URL: {}", e));
}
};
let host = match parsed_url.host_str() {
Some(h) => h.to_string(),
None => {
warn!(url = %target_url, "delegate target URL missing host");
return error_response(400, "bad_request", "URL missing host");
}
};
let port = parsed_url.port_or_known_default().unwrap_or(443);
let dns_start = Instant::now();
if let Err(e) = target_filter::validate_target(&host, port, allowed_ports, dns_cache).await {
warn!(host = %host, port, error = %e, "delegate target rejected");
return error_response(403, "target_not_allowed", &e.to_string());
}
let dns_ms = dns_start.elapsed().as_millis() as u64;
debug!(method = %method_str, url = %target_url, is_gzip, "delegate request");
// ── Build upstream request ──
let method = match method_str.parse::<Method>() {
Ok(m) => m,
Err(e) => {
warn!(error = %e, method = %method_str, "delegate invalid HTTP method");
return error_response(400, "bad_request", &format!("invalid method: {}", e));
}
};
let uri = match target_url.parse::<Uri>() {
Ok(u) => u,
Err(e) => {
warn!(error = %e, url = %target_url, "delegate invalid target URI");
return error_response(400, "bad_request", &format!("invalid URL: {}", e));
}
};
// ── Stream body passthrough ──
// When body is gzip-compressed, forward it directly to upstream with
// Content-Encoding: gzip header — no collect/decompress needed.
// All major AI API providers (Anthropic, OpenAI, Google) accept gzip request bodies.
let wire_size: u64;
let upstream_body: BoxBody;
if is_gzip {
let body_stream =
http_body_util::BodyStream::new(req.into_body()).filter_map(|result| async {
match result {
Ok(frame) => frame.into_data().ok().map(|data| {
Ok::<_, Box<dyn std::error::Error + Send + Sync>>(Frame::data(data))
}),
Err(e) => Some(Err(Box::new(e) as Box<dyn std::error::Error + Send + Sync>)),
}
});
let stream_body = StreamBody::new(body_stream);
upstream_body = BodyExt::boxed(stream_body);
// wire_size will be reported from Content-Length if available, otherwise 0
wire_size = req_content_length;
} else {
// Non-gzip: read body into memory (legacy path)
const MAX_BODY: usize = 10 * 1024 * 1024;
let body_bytes = match Limited::new(req.into_body(), MAX_BODY).collect().await {
Ok(collected) => collected.to_bytes(),
Err(e) => {
warn!(error = %e, "delegate failed to read request body");
return error_response(413, "payload_too_large", "request body exceeds 10MB limit");
}
};
wire_size = body_bytes.len() as u64;
let body = Full::new(body_bytes)
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> { match e {} })
.boxed();
upstream_body = body;
}
let mut upstream_req = Request::new(upstream_body);
*upstream_req.method_mut() = method;
*upstream_req.uri_mut() = uri;
{
let headers = upstream_req.headers_mut();
// Set headers (skip `host` — hyper sets it from the URI automatically,
// and a duplicate Host header can confuse certain upstreams)
for (name, value) in &upstream_headers {
if name.eq_ignore_ascii_case("host") {
continue;
}
let header_name = match HeaderName::from_bytes(name.as_bytes()) {
Ok(n) => n,
Err(_) => {
warn!(header = %name, "delegate invalid header name");
return error_response(400, "bad_request", "invalid header name");
}
};
let header_value = match HeaderValue::from_str(value) {
Ok(v) => v,
Err(_) => {
warn!(header = %name, "delegate invalid header value");
return error_response(400, "bad_request", "invalid header value");
}
};
headers.insert(header_name, header_value);
}
if is_gzip {
headers.insert(
hyper::header::CONTENT_ENCODING,
HeaderValue::from_static("gzip"),
);
}
}
// ── Send upstream request ──
// NOTE: We intentionally do NOT set a per-request timeout here.
// Connect timeout limits connection establishment; Aether controls
// first-byte / idle timeouts on its own side via asyncio.
let upstream_start = Instant::now();
let upstream_resp = match http_client.request(upstream_req).await {
Ok(resp) => resp,
Err(e) => {
warn!(url = %target_url, error = %e, "delegate upstream request failed");
let safe_detail = sanitize_upstream_error(&root_error_message(&e));
if is_timeout_error(&e) {
return error_response(504, "upstream_timeout", &safe_detail);
}
return error_response(502, "upstream_connection_failed", &safe_detail);
}
};
let ttfb_ms = upstream_start.elapsed().as_millis() as u64;
// ── Build response ──
let status = upstream_resp.status().as_u16();
let resp_headers = upstream_resp.headers().clone();
let (connect_ms, tls_ms) = upstream_resp
.extensions()
.get::<ConnectTiming>()
.map(|t| (t.connect_ms, t.tls_ms))
.unwrap_or((0, 0));
let upstream_processing_ms = ttfb_ms.saturating_sub(connect_ms.saturating_add(tls_ms));
let total_ms = total_start.elapsed().as_millis() as u64;
debug!(
url = %target_url,
status,
dns_ms,
connect_ms,
tls_ms,
ttfb_ms,
upstream_processing_ms,
total_ms,
wire_size,
is_gzip,
"delegate upstream response"
);
let timing = serde_json::json!({
"auth_ms": auth_ms,
"meta_ms": meta_ms,
"wire_size": wire_size,
"passthrough": is_gzip,
"dns_ms": dns_ms,
"connect_ms": connect_ms,
"tls_ms": tls_ms,
"ttfb_ms": ttfb_ms,
"upstream_ms": ttfb_ms,
"upstream_processing_ms": upstream_processing_ms,
"total_ms": total_ms,
});
let stream_body: BoxBody = upstream_resp
.into_body()
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> { Box::new(e) })
.boxed();
let mut builder = Response::builder().status(status);
for (name, value) in resp_headers.iter() {
builder = builder.header(name, value);
}
builder = builder.header("X-Proxy-Timing", timing.to_string());
builder.body(stream_body).unwrap_or_else(|_| {
Response::builder()
.status(500)
.body(super::empty_box_body())
.unwrap()
})
}
fn root_error_message(err: &dyn StdError) -> String {
let mut current = err;
while let Some(source) = current.source() {
current = source;
}
current.to_string()
}
fn is_timeout_error(err: &(dyn StdError + 'static)) -> bool {
if err.is::<tokio::time::error::Elapsed>() {
return true;
}
if let Some(io_err) = err.downcast_ref::<std::io::Error>() {
if io_err.kind() == std::io::ErrorKind::TimedOut {
return true;
}
}
if let Some(source) = err.source() {
// source() returns &(dyn Error + 'static), so this is safe
return is_timeout_error(source);
}
false
}
// ── Sanitisation ─────────────────────────────────────────────────────────────
/// Strip full URLs from error messages to prevent leaking upstream API keys,
/// paths, or query parameters in the delegate error response.
///
/// Replaces `https://api.example.com/v1/chat?key=xxx` with `api.example.com`.
fn sanitize_upstream_error(msg: &str) -> String {
// Simple regex-free approach: find "https://..." or "http://..." spans and
// replace them with just the host portion.
let mut result = msg.to_string();
for scheme in &["https://", "http://"] {
while let Some(start) = result.find(scheme) {
let after_scheme = start + scheme.len();
// Host ends at '/', '?', '#', ' ', or end of string
let host_end = result[after_scheme..]
.find(['/', '?', '#', ' '])
.map(|i| after_scheme + i)
.unwrap_or(result.len());
let host = &result[after_scheme..host_end];
result = format!("{}{}{}", &result[..start], host, &result[host_end..]);
}
}
result
}
// ── Error response helpers ───────────────────────────────────────────────────
fn error_response(status: u16, error: &str, detail: &str) -> Response<BoxBody> {
let body = serde_json::json!({
"error": error,
"detail": detail,
});
let body_bytes = bytes::Bytes::from(body.to_string());
Response::builder()
.status(status)
.header("Content-Type", "application/json")
.header("X-Delegate-Error", "true")
.body(
Full::new(body_bytes)
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> { match e {} })
.boxed(),
)
.unwrap()
}
-282
View File
@@ -1,282 +0,0 @@
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use std::time::{Duration, Instant};
use hyper::rt;
use hyper::Uri;
use hyper_util::client::legacy::connect::{Connected, Connection, HttpConnector};
use hyper_util::client::legacy::Client;
use hyper_util::rt::{TokioExecutor, TokioIo, TokioTimer};
use rustls::ClientConfig;
use rustls_pki_types::ServerName;
use tokio_rustls::TlsConnector;
use tower_service::Service;
use crate::config::Config;
use crate::proxy::BoxBody;
type BoxError = Box<dyn std::error::Error + Send + Sync>;
type DelegateStream = MaybeHttpsStream<TokioIo<tokio::net::TcpStream>>;
type DelegateConn = TimedConn<DelegateStream>;
pub(crate) type DelegateClient = Client<InstrumentedConnector, BoxBody>;
#[derive(Clone, Copy, Debug, Default)]
pub(crate) struct ConnectTiming {
pub connect_ms: u64,
pub tls_ms: u64,
}
pub(crate) fn build_delegate_client(config: &Config) -> DelegateClient {
let mut http = HttpConnector::new();
http.enforce_http(false);
http.set_connect_timeout(Some(Duration::from_secs(
config.delegate_connect_timeout_secs,
)));
http.set_nodelay(config.delegate_tcp_nodelay);
if config.delegate_tcp_keepalive_secs > 0 {
http.set_keepalive(Some(Duration::from_secs(
config.delegate_tcp_keepalive_secs,
)));
} else {
http.set_keepalive(None);
}
let connector = InstrumentedConnector {
http,
tls_config: build_tls_config(),
};
let mut builder = Client::builder(TokioExecutor::new());
builder.pool_max_idle_per_host(config.delegate_pool_max_idle_per_host);
builder.pool_idle_timeout(Duration::from_secs(config.delegate_pool_idle_timeout_secs));
builder.pool_timer(TokioTimer::new());
builder.build::<_, BoxBody>(connector)
}
fn build_tls_config() -> Arc<ClientConfig> {
let root_store =
rustls::RootCertStore::from_iter(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
let mut config = ClientConfig::builder()
.with_root_certificates(root_store)
.with_no_client_auth();
config.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec()];
Arc::new(config)
}
#[derive(Clone)]
pub(crate) struct InstrumentedConnector {
http: HttpConnector,
tls_config: Arc<ClientConfig>,
}
impl Service<Uri> for InstrumentedConnector {
type Response = DelegateConn;
type Error = BoxError;
type Future = Pin<Box<dyn Future<Output = Result<Self::Response, BoxError>> + Send>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.http.poll_ready(cx).map_err(Into::into)
}
fn call(&mut self, dst: Uri) -> Self::Future {
let scheme = dst.scheme_str().map(|s| s.to_ascii_lowercase());
let tls_config = self.tls_config.clone();
let connecting = self.http.call(dst.clone());
let connect_start = Instant::now();
Box::pin(async move {
match scheme.as_deref() {
Some("http") => {
let tcp = connecting.await.map_err(|e| Box::new(e) as BoxError)?;
let connect_ms = connect_start.elapsed().as_millis() as u64;
Ok(TimedConn::new(
MaybeHttpsStream::Http(tcp),
ConnectTiming {
connect_ms,
tls_ms: 0,
},
))
}
Some("https") => {
let server_name = resolve_server_name(&dst)?;
let tcp = connecting.await.map_err(|e| Box::new(e) as BoxError)?;
let connect_ms = connect_start.elapsed().as_millis() as u64;
let tls_start = Instant::now();
let tls_stream = TlsConnector::from(tls_config)
.connect(server_name, TokioIo::new(tcp))
.await
.map_err(std::io::Error::other)?;
let tls_ms = tls_start.elapsed().as_millis() as u64;
Ok(TimedConn::new(
MaybeHttpsStream::Https(TokioIo::new(tls_stream)),
ConnectTiming { connect_ms, tls_ms },
))
}
Some(other) => {
Err(std::io::Error::other(format!("unsupported scheme {other}")).into())
}
None => Err(std::io::Error::other("missing scheme").into()),
}
})
}
}
fn resolve_server_name(uri: &Uri) -> Result<ServerName<'static>, BoxError> {
let host = uri.host().ok_or("missing host")?;
let host = host.trim_start_matches('[').trim_end_matches(']');
Ok(ServerName::try_from(host.to_string())?)
}
pub(crate) struct TimedConn<T> {
inner: T,
timing: ConnectTiming,
}
impl<T> TimedConn<T> {
fn new(inner: T, timing: ConnectTiming) -> Self {
Self { inner, timing }
}
}
impl<T: Connection> Connection for TimedConn<T> {
fn connected(&self) -> Connected {
self.inner.connected().extra(self.timing)
}
}
impl<T: rt::Read + Unpin> rt::Read for TimedConn<T> {
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: rt::ReadBufCursor<'_>,
) -> Poll<Result<(), std::io::Error>> {
Pin::new(&mut self.inner).poll_read(cx, buf)
}
}
impl<T: rt::Write + Unpin> rt::Write for TimedConn<T> {
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<Result<usize, std::io::Error>> {
Pin::new(&mut self.inner).poll_write(cx, buf)
}
fn poll_flush(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Result<(), std::io::Error>> {
Pin::new(&mut self.inner).poll_flush(cx)
}
fn poll_shutdown(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Result<(), std::io::Error>> {
Pin::new(&mut self.inner).poll_shutdown(cx)
}
fn is_write_vectored(&self) -> bool {
self.inner.is_write_vectored()
}
fn poll_write_vectored(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
bufs: &[std::io::IoSlice<'_>],
) -> Poll<Result<usize, std::io::Error>> {
Pin::new(&mut self.inner).poll_write_vectored(cx, bufs)
}
}
#[allow(clippy::large_enum_variant)]
pub(crate) enum MaybeHttpsStream<T> {
Http(T),
Https(TokioIo<tokio_rustls::client::TlsStream<TokioIo<T>>>),
}
impl<T: rt::Read + rt::Write + Connection + Unpin> Connection for MaybeHttpsStream<T> {
fn connected(&self) -> Connected {
match self {
Self::Http(stream) => stream.connected(),
Self::Https(stream) => {
let (tcp, tls) = stream.inner().get_ref();
if tls.alpn_protocol() == Some(b"h2") {
tcp.inner().connected().negotiated_h2()
} else {
tcp.inner().connected()
}
}
}
}
}
impl<T: rt::Read + rt::Write + Unpin> rt::Read for MaybeHttpsStream<T> {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: rt::ReadBufCursor<'_>,
) -> Poll<Result<(), std::io::Error>> {
match Pin::get_mut(self) {
Self::Http(stream) => Pin::new(stream).poll_read(cx, buf),
Self::Https(stream) => Pin::new(stream).poll_read(cx, buf),
}
}
}
impl<T: rt::Write + rt::Read + Unpin> rt::Write for MaybeHttpsStream<T> {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<Result<usize, std::io::Error>> {
match Pin::get_mut(self) {
Self::Http(stream) => Pin::new(stream).poll_write(cx, buf),
Self::Https(stream) => Pin::new(stream).poll_write(cx, buf),
}
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), std::io::Error>> {
match Pin::get_mut(self) {
Self::Http(stream) => Pin::new(stream).poll_flush(cx),
Self::Https(stream) => Pin::new(stream).poll_flush(cx),
}
}
fn poll_shutdown(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Result<(), std::io::Error>> {
match Pin::get_mut(self) {
Self::Http(stream) => Pin::new(stream).poll_shutdown(cx),
Self::Https(stream) => Pin::new(stream).poll_shutdown(cx),
}
}
fn is_write_vectored(&self) -> bool {
match self {
Self::Http(stream) => stream.is_write_vectored(),
Self::Https(stream) => stream.is_write_vectored(),
}
}
fn poll_write_vectored(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
bufs: &[std::io::IoSlice<'_>],
) -> Poll<Result<usize, std::io::Error>> {
match Pin::get_mut(self) {
Self::Http(stream) => Pin::new(stream).poll_write_vectored(cx, bufs),
Self::Https(stream) => Pin::new(stream).poll_write_vectored(cx, bufs),
}
}
}
-19
View File
@@ -1,19 +0,0 @@
pub mod connect;
pub mod delegate;
pub mod delegate_client;
pub mod server;
pub mod target_filter;
pub mod tls;
use http_body_util::BodyExt;
/// Boxed body type used across proxy handlers.
pub type BoxBody =
http_body_util::combinators::BoxBody<bytes::Bytes, Box<dyn std::error::Error + Send + Sync>>;
/// Create an empty [`BoxBody`] (for error responses, 405, etc.).
pub fn empty_box_body() -> BoxBody {
http_body_util::Full::new(bytes::Bytes::new())
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> { match e {} })
.boxed()
}
-213
View File
@@ -1,213 +0,0 @@
use std::net::SocketAddr;
use std::sync::atomic::Ordering;
use std::sync::Arc;
use std::time::{Duration, Instant};
use http_body_util::BodyExt;
use hyper::body::Incoming;
use hyper::rt::{Read, Write};
use hyper::server::conn::http1;
use hyper::service::service_fn;
use hyper::{Method, Request, Response};
use hyper_util::rt::TokioIo;
use tokio::net::TcpListener;
use tokio::sync::watch;
use tokio::time::timeout;
use tracing::{debug, info, warn};
use crate::proxy::{connect, delegate, tls, BoxBody};
use crate::state::AppState;
/// Start the proxy server.
///
/// Listens for incoming TCP connections and dispatches:
/// - CONNECT requests -> tunnel handler
/// - POST /_aether/delegate -> delegate handler
/// - Other requests -> 405 Method Not Allowed
///
/// When TLS is configured, the server operates in dual-stack mode:
/// it peeks at the first byte of each connection to distinguish TLS ClientHello
/// (0x16) from plain HTTP, and handles both on the same port.
pub async fn run(
state: &Arc<AppState>,
mut shutdown_rx: watch::Receiver<bool>,
) -> anyhow::Result<()> {
let addr = SocketAddr::from(([0, 0, 0, 0], state.config.listen_port));
let listener = TcpListener::bind(addr).await?;
if state.tls_acceptor.is_some() {
info!(addr = %addr, "proxy server listening (HTTP+TLS dual-stack)");
} else {
info!(addr = %addr, "proxy server listening (HTTP only)");
}
let handshake_timeout = Duration::from_secs(state.config.tls_handshake_timeout_secs);
loop {
tokio::select! {
result = listener.accept() => {
let (stream, peer_addr) = match result {
Ok(v) => v,
Err(e) => {
warn!(error = %e, "failed to accept connection");
continue;
}
};
debug!(peer = %peer_addr, "new connection");
if let Err(e) = stream.set_nodelay(true) {
debug!(peer = %peer_addr, error = %e, "failed to set TCP_NODELAY");
}
let permit = match state.connection_semaphore.clone().try_acquire_owned() {
Ok(permit) => permit,
Err(_) => {
warn!(peer = %peer_addr, "connection rejected: limit reached");
continue;
}
};
let state = Arc::clone(state);
state.active_connections.fetch_add(1, Ordering::Relaxed);
tokio::task::spawn(async move {
let _permit = permit;
// Dual-stack: peek first byte to decide TLS vs plain HTTP
if let Some(ref acceptor) = state.tls_acceptor {
let is_tls = match timeout(handshake_timeout, tls::is_tls_client_hello(&stream)).await {
Ok(v) => v,
Err(_) => {
debug!(peer = %peer_addr, "TLS detection timeout");
state.active_connections.fetch_sub(1, Ordering::Relaxed);
return;
}
};
if is_tls {
match timeout(handshake_timeout, acceptor.clone().accept(stream)).await {
Ok(Ok(tls_stream)) => {
debug!(peer = %peer_addr, "TLS handshake ok");
serve_connection(
TokioIo::new(tls_stream),
peer_addr,
&state,
)
.await;
}
Ok(Err(e)) => {
debug!(peer = %peer_addr, error = %e, "TLS handshake failed");
}
Err(_) => {
debug!(peer = %peer_addr, "TLS handshake timeout");
}
}
state.active_connections.fetch_sub(1, Ordering::Relaxed);
return;
}
}
// Plain HTTP
serve_connection(
TokioIo::new(stream),
peer_addr,
&state,
)
.await;
state.active_connections.fetch_sub(1, Ordering::Relaxed);
});
}
_ = shutdown_rx.changed() => {
info!("proxy server shutting down");
break;
}
}
}
Ok(())
}
/// Serve a single HTTP/1.1 connection (works over both plain TCP and TLS).
async fn serve_connection<I>(io: I, peer_addr: SocketAddr, state: &Arc<AppState>)
where
I: Read + Write + Unpin + Send + 'static,
{
let config = Arc::clone(&state.config);
let dynamic = Arc::clone(&state.dynamic);
let delegate_client = state.delegate_client.clone();
let dns_cache = Arc::clone(&state.dns_cache);
let metrics = Arc::clone(&state.metrics);
let service = service_fn(move |req: Request<Incoming>| {
let config = Arc::clone(&config);
let dynamic = Arc::clone(&dynamic);
let delegate_client = delegate_client.clone();
let dns_cache = Arc::clone(&dns_cache);
let metrics = Arc::clone(&metrics);
async move {
let start = Instant::now();
// Snapshot current dynamic values (may be updated by remote config)
let (allowed_ports, timestamp_tolerance) = {
let d = dynamic.read().unwrap();
(d.allowed_ports.clone(), d.timestamp_tolerance)
};
if req.method() == Method::CONNECT {
let resp = connect::handle_connect(
req,
config,
&allowed_ports,
timestamp_tolerance,
dns_cache.as_ref(),
)
.await;
let resp = resp.map(|_| -> BoxBody {
http_body_util::Empty::new()
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> { match e {} })
.boxed()
});
metrics.record_request(start.elapsed());
Ok::<_, hyper::Error>(resp)
} else if req.uri().path() == "/_aether/delegate" && req.method() == hyper::Method::POST
{
let resp = delegate::handle_delegate(
req,
config,
&allowed_ports,
timestamp_tolerance,
dns_cache.as_ref(),
&delegate_client,
)
.await;
metrics.record_request(start.elapsed());
Ok(resp)
} else {
// Only CONNECT tunnels and /_aether/delegate are supported;
// plain HTTP forward proxy was removed (all API traffic is HTTPS).
let resp = Response::builder()
.status(405)
.header("Allow", "CONNECT")
.header("Content-Length", "0")
.body(crate::proxy::empty_box_body())
.unwrap();
metrics.record_request(start.elapsed());
Ok(resp)
}
}
});
if let Err(e) = http1::Builder::new()
.preserve_header_case(true)
.title_case_headers(false)
.serve_connection(io, service)
.with_upgrades()
.await
{
if !e.to_string().contains("connection closed") {
debug!(peer = %peer_addr, error = %e, "connection error");
}
}
}
-125
View File
@@ -1,125 +0,0 @@
use std::fs;
use std::io::BufReader;
use std::path::Path;
use std::sync::Arc;
use rcgen::{CertificateParams, KeyPair};
use rustls_pki_types::{CertificateDer, PrivateKeyDer};
use sha2::{Digest, Sha256};
use tokio_rustls::TlsAcceptor;
use tracing::{info, warn};
const SESSION_CACHE_SIZE: usize = 1024;
/// Generate a self-signed certificate if the files do not already exist.
///
/// The certificate includes SANs: `localhost` and `aether-proxy`.
/// The private key file is set to mode 0600 on unix.
pub fn ensure_self_signed_cert(cert_path: &Path, key_path: &Path) -> anyhow::Result<()> {
if cert_path.exists() && key_path.exists() {
info!(
cert = %cert_path.display(),
key = %key_path.display(),
"using existing TLS certificate"
);
return Ok(());
}
info!("generating self-signed TLS certificate");
let mut params = CertificateParams::new(vec!["localhost".into(), "aether-proxy".into()])?;
params.distinguished_name = rcgen::DistinguishedName::new();
params
.distinguished_name
.push(rcgen::DnType::CommonName, "aether-proxy");
let key_pair = KeyPair::generate()?;
let cert = params.self_signed(&key_pair)?;
let cert_pem = cert.pem();
let key_pem = key_pair.serialize_pem();
fs::write(cert_path, &cert_pem)?;
fs::write(key_path, &key_pem)?;
// Set key file permissions to 0600 on unix
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
let perms = fs::Permissions::from_mode(0o600);
fs::set_permissions(key_path, perms)?;
}
info!(
cert = %cert_path.display(),
key = %key_path.display(),
"self-signed TLS certificate generated"
);
Ok(())
}
/// Build a `TlsAcceptor` from PEM certificate and key files.
pub fn build_tls_acceptor(cert_path: &Path, key_path: &Path) -> anyhow::Result<TlsAcceptor> {
let cert_file = fs::File::open(cert_path)?;
let key_file = fs::File::open(key_path)?;
let certs: Vec<CertificateDer<'static>> =
rustls_pemfile::certs(&mut BufReader::new(cert_file)).collect::<Result<Vec<_>, _>>()?;
if certs.is_empty() {
anyhow::bail!("no certificates found in {}", cert_path.display());
}
let key: PrivateKeyDer<'static> =
rustls_pemfile::private_key(&mut BufReader::new(key_file))?
.ok_or_else(|| anyhow::anyhow!("no private key found in {}", key_path.display()))?;
let mut config = rustls::ServerConfig::builder()
.with_no_client_auth()
.with_single_cert(certs, key)?;
config.alpn_protocols = vec![b"http/1.1".to_vec()];
config.session_storage = rustls::server::ServerSessionMemoryCache::new(SESSION_CACHE_SIZE);
match rustls::crypto::ring::Ticketer::new() {
Ok(ticketer) => {
config.ticketer = ticketer;
}
Err(e) => {
warn!(error = %e, "failed to init TLS ticketer; tickets disabled");
}
}
Ok(TlsAcceptor::from(Arc::new(config)))
}
/// Compute the SHA-256 fingerprint of the first certificate in a PEM file.
///
/// Returns the hex-encoded fingerprint (lowercase, no separators).
pub fn cert_sha256_fingerprint(cert_path: &Path) -> anyhow::Result<String> {
let cert_file = fs::File::open(cert_path)?;
let certs: Vec<CertificateDer<'static>> =
rustls_pemfile::certs(&mut BufReader::new(cert_file)).collect::<Result<Vec<_>, _>>()?;
let cert = certs
.first()
.ok_or_else(|| anyhow::anyhow!("no certificates found in {}", cert_path.display()))?;
let digest = Sha256::digest(cert.as_ref());
Ok(hex::encode(digest))
}
/// Peek at the first byte of a TCP stream to determine if it is a TLS ClientHello.
///
/// Returns `true` if the first byte is 0x16 (TLS record type: Handshake).
pub async fn is_tls_client_hello(stream: &tokio::net::TcpStream) -> bool {
let mut buf = [0u8; 1];
match stream.peek(&mut buf).await {
Ok(1) => buf[0] == 0x16,
Ok(_) => false,
Err(e) => {
warn!(error = %e, "failed to peek first byte");
false
}
}
}
+14 -142
View File
@@ -3,30 +3,11 @@ use std::time::{Duration, SystemTime, UNIX_EPOCH};
use reqwest::{Client, StatusCode};
use serde::{Deserialize, Serialize};
use tokio::time::sleep;
use tracing::{debug, error, info, warn};
use tracing::{debug, error, info};
use crate::config::Config;
use crate::hardware::HardwareInfo;
/// Heartbeat-specific error that distinguishes "node not found" (needs
/// re-registration) from transient / other failures.
#[derive(Debug)]
pub enum HeartbeatError {
/// HTTP 404 – the node_id is no longer known to Aether.
NodeNotFound(String),
/// Any other failure (network, 5xx, etc.).
Other(anyhow::Error),
}
impl std::fmt::Display for HeartbeatError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::NodeNotFound(msg) => write!(f, "node not found: {}", msg),
Self::Other(e) => write!(f, "{}", e),
}
}
}
#[derive(Debug, Serialize)]
struct RegisterRequest {
name: String,
@@ -35,14 +16,13 @@ struct RegisterRequest {
#[serde(skip_serializing_if = "Option::is_none")]
region: Option<String>,
heartbeat_interval: u64,
#[serde(skip_serializing_if = "std::ops::Not::not")]
tls_enabled: bool,
#[serde(skip_serializing_if = "Option::is_none")]
tls_cert_fingerprint: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
hardware_info: Option<serde_json::Value>,
#[serde(skip_serializing_if = "Option::is_none")]
estimated_max_concurrency: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
proxy_metadata: Option<serde_json::Value>,
tunnel_mode: bool,
}
#[derive(Debug, Deserialize)]
@@ -50,17 +30,6 @@ pub struct RegisterResponse {
pub node_id: String,
}
#[derive(Debug, Serialize)]
struct HeartbeatRequest {
node_id: String,
#[serde(skip_serializing_if = "Option::is_none")]
active_connections: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
total_requests: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
avg_latency_ms: Option<f64>,
}
/// Remote configuration pushed by the Aether management backend.
#[derive(Debug, Clone, Deserialize)]
pub struct RemoteConfig {
@@ -68,29 +37,6 @@ pub struct RemoteConfig {
pub allowed_ports: Option<Vec<u16>>,
pub log_level: Option<String>,
pub heartbeat_interval: Option<u64>,
pub timestamp_tolerance: Option<u64>,
}
/// Parsed heartbeat response from Aether.
#[derive(Debug, Deserialize)]
struct HeartbeatResponseBody {
#[serde(default)]
node: Option<HeartbeatNodeInfo>,
}
#[derive(Debug, Deserialize)]
struct HeartbeatNodeInfo {
#[serde(default)]
remote_config: Option<RemoteConfig>,
#[serde(default)]
config_version: Option<u64>,
}
/// Heartbeat result returned to the caller.
#[derive(Debug)]
pub struct HeartbeatResult {
pub remote_config: Option<RemoteConfig>,
pub config_version: u64,
}
#[derive(Debug, Serialize)]
@@ -109,7 +55,7 @@ pub struct AetherClient {
}
impl AetherClient {
pub fn new(config: &Config) -> Self {
pub fn new(config: &Config, aether_url: &str, management_token: &str) -> Self {
let mut builder = Client::builder()
.timeout(Duration::from_secs(config.aether_request_timeout_secs))
.connect_timeout(Duration::from_secs(config.aether_connect_timeout_secs))
@@ -136,8 +82,8 @@ impl AetherClient {
Self {
http,
base_url: config.aether_url.trim_end_matches('/').to_string(),
token: config.management_token.clone(),
base_url: aether_url.trim_end_matches('/').to_string(),
token: management_token.to_string(),
retry_max_attempts: config.aether_retry_max_attempts.max(1),
retry_base_delay,
retry_max_delay,
@@ -150,29 +96,29 @@ impl AetherClient {
pub async fn register(
&self,
config: &Config,
node_name: &str,
public_ip: &str,
tls_enabled: bool,
tls_cert_fingerprint: Option<&str>,
hw: Option<&HardwareInfo>,
) -> anyhow::Result<String> {
let url = format!("{}/api/admin/proxy-nodes/register", self.base_url);
let body = RegisterRequest {
name: config.node_name.clone(),
name: node_name.to_string(),
ip: public_ip.to_string(),
port: config.listen_port,
port: 0,
region: config.node_region.clone(),
heartbeat_interval: config.heartbeat_interval,
tls_enabled,
tls_cert_fingerprint: tls_cert_fingerprint.map(|s| s.to_string()),
hardware_info: hw.and_then(|h| serde_json::to_value(h).ok()),
estimated_max_concurrency: hw.map(|h| h.estimated_max_concurrency),
proxy_metadata: Some(serde_json::json!({
"version": env!("CARGO_PKG_VERSION"),
})),
tunnel_mode: true,
};
info!(
url = %url,
name = %body.name,
ip = %body.ip,
port = body.port,
"registering with Aether"
);
@@ -199,80 +145,6 @@ impl AetherClient {
Ok(data.node_id)
}
/// Send heartbeat to Aether.
///
/// On success, returns any remote config included in the response.
/// Returns [`HeartbeatError::NodeNotFound`] on HTTP 404 so the caller
/// can trigger re-registration.
pub async fn heartbeat(
&self,
node_id: &str,
active_connections: Option<i64>,
total_requests: Option<i64>,
avg_latency_ms: Option<f64>,
) -> Result<HeartbeatResult, HeartbeatError> {
let url = format!("{}/api/admin/proxy-nodes/heartbeat", self.base_url);
let body = HeartbeatRequest {
node_id: node_id.to_string(),
active_connections,
total_requests,
avg_latency_ms,
};
debug!(node_id = %node_id, "sending heartbeat");
let resp = self
.send_with_retry(
|| {
self.http
.post(&url)
.header("Authorization", format!("Bearer {}", self.token))
.json(&body)
},
"heartbeat",
)
.await
.map_err(|e| HeartbeatError::Other(e.into()))?;
let status = resp.status();
if !status.is_success() {
let text = resp.text().await.unwrap_or_default();
warn!(status = %status, body = %text, "heartbeat failed");
if status == StatusCode::NOT_FOUND {
return Err(HeartbeatError::NodeNotFound(text));
}
return Err(HeartbeatError::Other(anyhow::anyhow!(
"heartbeat failed (HTTP {}): {}",
status,
text
)));
}
// Parse remote config from response (best-effort)
let result = match resp.json::<HeartbeatResponseBody>().await {
Ok(body) => {
let (remote_config, config_version) = match body.node {
Some(node) => (node.remote_config, node.config_version.unwrap_or(0)),
None => (None, 0),
};
HeartbeatResult {
remote_config,
config_version,
}
}
Err(e) => {
debug!(error = %e, "failed to parse heartbeat response body");
HeartbeatResult {
remote_config: None,
config_version: 0,
}
}
};
debug!(node_id = %node_id, config_version = result.config_version, "heartbeat ok");
Ok(result)
}
/// Unregister this node from Aether (graceful shutdown).
pub async fn unregister(&self, node_id: &str) -> anyhow::Result<()> {
let url = format!("{}/api/admin/proxy-nodes/unregister", self.base_url);
-127
View File
@@ -1,127 +0,0 @@
use std::sync::atomic::Ordering;
use std::sync::Arc;
use tokio::sync::watch;
use tracing::{debug, error, info, warn};
use crate::registration::client::HeartbeatError;
use crate::runtime;
use crate::state::AppState;
/// Run periodic heartbeat task until shutdown signal.
///
/// When Aether responds with 404 (node not found), this task automatically
/// re-registers the node and updates the shared `node_id` so the proxy
/// server and future heartbeats use the new identity.
///
/// When the heartbeat response includes a `remote_config`, it is applied
/// to the [`DynamicConfig`](crate::runtime::DynamicConfig) so the proxy
/// picks up changes without a restart.
pub async fn run(state: &Arc<AppState>, mut shutdown_rx: watch::Receiver<bool>) {
let mut consecutive_failures: u32 = 0;
// Skip the first tick (registration already acts as initial heartbeat)
let initial_interval = state.dynamic.read().unwrap().heartbeat_interval;
tokio::select! {
_ = tokio::time::sleep(std::time::Duration::from_secs(initial_interval)) => {}
_ = shutdown_rx.changed() => {
debug!("heartbeat task stopping (during initial wait)");
return;
}
}
loop {
let current_node_id = state.node_id.read().unwrap().clone();
let active_conns = state.active_connections.load(Ordering::Relaxed) as i64;
// Swap-and-reset: report incremental metrics since last heartbeat
let interval_requests = state.metrics.total_requests.swap(0, Ordering::Relaxed);
let interval_latency_ns = state.metrics.total_latency_ns.swap(0, Ordering::Relaxed);
let interval_requests_i64 = i64::try_from(interval_requests).unwrap_or(i64::MAX);
let avg_latency_ms = if interval_requests > 0 {
Some(interval_latency_ns as f64 / interval_requests as f64 / 1_000_000.0)
} else {
None
};
match state
.aether_client
.heartbeat(
&current_node_id,
Some(active_conns),
Some(interval_requests_i64),
avg_latency_ms,
)
.await
{
Ok(result) => {
if consecutive_failures > 0 {
info!(
previous_failures = consecutive_failures,
"heartbeat recovered"
);
}
consecutive_failures = 0;
// Apply remote config if present and version changed
if let Some(ref remote) = result.remote_config {
runtime::apply_remote_config(&state.dynamic, remote, result.config_version);
}
}
Err(HeartbeatError::NodeNotFound(_)) => {
warn!(
old_node_id = %current_node_id,
"node not found, re-registering"
);
match state
.aether_client
.register(
&state.config,
&state.public_ip,
state.config.enable_tls,
state.tls_fingerprint.as_deref(),
Some(&state.hardware_info),
)
.await
{
Ok(new_id) => {
info!(
old_node_id = %current_node_id,
new_node_id = %new_id,
"re-registered successfully"
);
*state.node_id.write().unwrap() = new_id;
consecutive_failures = 0;
}
Err(e) => {
consecutive_failures += 1;
error!(
error = %e,
consecutive_failures,
"re-registration failed"
);
}
}
}
Err(HeartbeatError::Other(e)) => {
consecutive_failures += 1;
warn!(
error = %e,
consecutive_failures,
"heartbeat failed"
);
}
}
// Read interval from dynamic config (may have been updated remotely)
let interval_secs = state.dynamic.read().unwrap().heartbeat_interval;
tokio::select! {
_ = tokio::time::sleep(std::time::Duration::from_secs(interval_secs)) => {}
_ = shutdown_rx.changed() => {
debug!("heartbeat task stopping");
break;
}
}
}
}
-1
View File
@@ -1,2 +1 @@
pub mod client;
pub mod heartbeat;
+31 -33
View File
@@ -5,18 +5,18 @@
//! management backend through the heartbeat response.
use std::collections::HashSet;
use std::sync::{Arc, OnceLock, RwLock};
use std::sync::{Arc, OnceLock};
use arc_swap::ArcSwap;
use tracing::info;
use crate::config::Config;
/// Configuration that can be changed at runtime without restart.
#[derive(Debug)]
#[derive(Debug, Clone)]
pub struct DynamicConfig {
pub node_name: String,
pub allowed_ports: HashSet<u16>,
pub timestamp_tolerance: u64,
pub allowed_ports: Arc<HashSet<u16>>,
pub log_level: String,
pub heartbeat_interval: u64,
/// Monotonically increasing version from the backend.
@@ -29,8 +29,7 @@ impl DynamicConfig {
pub fn from_config(config: &Config) -> Self {
Self {
node_name: config.node_name.clone(),
allowed_ports: config.allowed_ports.iter().copied().collect(),
timestamp_tolerance: config.timestamp_tolerance,
allowed_ports: Arc::new(config.allowed_ports.iter().copied().collect()),
log_level: config.log_level.clone(),
heartbeat_interval: config.heartbeat_interval,
config_version: 0,
@@ -38,10 +37,10 @@ impl DynamicConfig {
}
}
/// Shared dynamic config handle.
pub type SharedDynamicConfig = Arc<RwLock<DynamicConfig>>;
/// Shared dynamic config handle (lock-free reads via ArcSwap).
pub type SharedDynamicConfig = Arc<ArcSwap<DynamicConfig>>;
// ── Log-level hot-reload ─────────────────────────────────────────────────────
// -- Log-level hot-reload -----
/// Global log-level reloader function, set during tracing init.
type LogReloader = Box<dyn Fn(&str) + Send + Sync>;
@@ -55,53 +54,50 @@ pub fn set_log_reloader(f: LogReloader) {
/// Apply a remote config update to the dynamic config.
///
/// Uses copy-on-write: loads the current snapshot, clones it, applies changes,
/// and stores the new Arc. Reads are always lock-free.
///
/// Returns `true` if the config was actually changed.
pub fn apply_remote_config(
dynamic: &SharedDynamicConfig,
remote: &crate::registration::client::RemoteConfig,
version: u64,
) -> bool {
let mut cfg = dynamic.write().unwrap();
let current = dynamic.load();
if version <= cfg.config_version {
if version <= current.config_version {
return false;
}
let mut new_cfg = (**current).clone();
let mut changed = Vec::new();
if let Some(ref name) = remote.node_name {
if *name != cfg.node_name {
changed.push(format!("node_name → {}", name));
cfg.node_name = name.clone();
if *name != new_cfg.node_name {
changed.push(format!("node_name -> {}", name));
new_cfg.node_name = name.clone();
}
}
if let Some(ref ports) = remote.allowed_ports {
let new_set: HashSet<u16> = ports.iter().copied().collect();
if new_set != cfg.allowed_ports {
changed.push(format!("allowed_ports → {:?}", ports));
cfg.allowed_ports = new_set;
}
}
if let Some(tol) = remote.timestamp_tolerance {
if tol != cfg.timestamp_tolerance {
changed.push(format!("timestamp_tolerance → {}", tol));
cfg.timestamp_tolerance = tol;
if new_set != *new_cfg.allowed_ports {
changed.push(format!("allowed_ports -> {:?}", ports));
new_cfg.allowed_ports = Arc::new(new_set);
}
}
if let Some(interval) = remote.heartbeat_interval {
if interval != cfg.heartbeat_interval {
changed.push(format!("heartbeat_interval → {}s", interval));
cfg.heartbeat_interval = interval;
if interval != new_cfg.heartbeat_interval {
changed.push(format!("heartbeat_interval -> {}s", interval));
new_cfg.heartbeat_interval = interval;
}
}
if let Some(ref level) = remote.log_level {
if *level != cfg.log_level {
changed.push(format!("log_level → {}", level));
cfg.log_level = level.clone();
if *level != new_cfg.log_level {
changed.push(format!("log_level -> {}", level));
new_cfg.log_level = level.clone();
// Hot-reload tracing filter
if let Some(reloader) = LOG_RELOADER.get() {
reloader(level);
@@ -109,15 +105,17 @@ pub fn apply_remote_config(
}
}
cfg.config_version = version;
let has_changes = !changed.is_empty();
if !changed.is_empty() {
if has_changes {
new_cfg.config_version = version;
info!(
version,
changes = %changed.join(", "),
"remote config applied"
);
dynamic.store(Arc::new(new_cfg));
}
!changed.is_empty()
has_changes
}
+66
View File
@@ -0,0 +1,66 @@
//! Safe DNS resolver for reqwest that reuses validated addresses from DnsCache.
//!
//! This resolver ensures reqwest connects only to addresses that have been
//! previously validated by `target_filter::validate_target()`, eliminating
//! the TOCTTOU gap where DNS rebinding could redirect traffic to private IPs.
use std::net::SocketAddr;
use std::sync::Arc;
use reqwest::dns::{Addrs, Name, Resolve, Resolving};
use crate::target_filter::{self, DnsCache};
/// A DNS resolver that serves validated public addresses from the shared DnsCache.
///
/// When reqwest needs to resolve a hostname, this resolver returns addresses
/// from the cache (populated by `validate_target()` during request validation).
/// If the hostname is not in cache (shouldn't happen in normal flow), it
/// performs a fresh resolution with private-IP filtering.
pub struct SafeDnsResolver {
dns_cache: Arc<DnsCache>,
}
impl SafeDnsResolver {
pub fn new(dns_cache: Arc<DnsCache>) -> Self {
Self { dns_cache }
}
}
impl Resolve for SafeDnsResolver {
fn resolve(&self, name: Name) -> Resolving {
let dns_cache = Arc::clone(&self.dns_cache);
Box::pin(async move {
let host = name.as_str();
// Try cache first (should be populated by validate_target).
// reqwest resolves by hostname only (no port), so use host-only lookup.
if let Some(addrs) = dns_cache.get_by_host(host).await {
let socket_addrs: Vec<SocketAddr> = (*addrs).clone();
return Ok(Box::new(socket_addrs.into_iter()) as Addrs);
}
// Fallback: resolve with private-IP filtering (defensive).
// This path should rarely be hit since validate_target() runs first.
// We don't know the real port here (reqwest Resolve only gives hostname),
// so resolve directly without caching to avoid polluting the cache with
// an incorrect port-based key.
let addr_str = format!("{}:0", host);
let resolved: Vec<SocketAddr> = tokio::net::lookup_host(&addr_str)
.await
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> { Box::new(e) })?
.filter(|addr| !target_filter::is_private_ip(&addr.ip()))
.collect();
if resolved.is_empty() {
return Err(Box::new(std::io::Error::other(format!(
"all resolved addresses for {} are private/reserved",
host
)))
as Box<dyn std::error::Error + Send + Sync>);
}
Ok(Box::new(resolved.into_iter()) as Addrs)
})
}
}
+9 -8
View File
@@ -21,7 +21,7 @@ pub fn install_service(config_path: &Path) -> anyhow::Result<()> {
anyhow::bail!("systemd not available");
}
if !is_root() {
anyhow::bail!("root required, use: sudo aether-proxy setup");
anyhow::bail!("root required, use: sudo ./aether-proxy setup");
}
let exe_path = std::env::current_exe()?.canonicalize()?;
@@ -67,6 +67,7 @@ pub fn install_service(config_path: &Path) -> anyhow::Result<()> {
Restart=on-failure\n\
RestartSec=5\n\
LimitNOFILE=65535\n\
UMask=0077\n\
\n\
[Install]\n\
WantedBy=multi-user.target\n",
@@ -93,11 +94,11 @@ pub fn install_service(config_path: &Path) -> anyhow::Result<()> {
eprintln!();
eprintln!(" Commands:");
eprintln!(" aether-proxy status # service status");
eprintln!(" aether-proxy logs # tail logs");
eprintln!(" sudo aether-proxy restart # restart");
eprintln!(" sudo aether-proxy stop # stop");
eprintln!(" sudo aether-proxy uninstall # remove service");
eprintln!(" ./aether-proxy status # service status");
eprintln!(" ./aether-proxy logs # tail logs");
eprintln!(" sudo ./aether-proxy restart # restart");
eprintln!(" sudo ./aether-proxy stop # stop");
eprintln!(" sudo ./aether-proxy uninstall # remove service");
eprintln!();
Ok(())
@@ -165,7 +166,7 @@ pub fn is_service_active() -> bool {
fn ensure_service_installed() -> anyhow::Result<()> {
if !std::path::Path::new(UNIT_PATH).exists() {
anyhow::bail!("service not installed, run `sudo aether-proxy setup` first");
anyhow::bail!("service not installed, run `sudo ./aether-proxy setup` first");
}
Ok(())
}
@@ -173,7 +174,7 @@ fn ensure_service_installed() -> anyhow::Result<()> {
fn ensure_root_and_service() -> anyhow::Result<()> {
ensure_service_installed()?;
if !is_root() {
anyhow::bail!("root required, use: sudo aether-proxy <command>");
anyhow::bail!("root required, use: sudo ./aether-proxy <command>");
}
Ok(())
}
+364 -191
View File
@@ -2,7 +2,8 @@
//!
//! Launched via `aether-proxy setup [path]`. Presents a full-screen form
//! backed by ratatui where the user can navigate fields, edit values, and
//! save to a TOML config file.
//! save to a TOML config file. Supports multi-server configuration via
//! a tabbed interface.
use std::io;
use std::path::PathBuf;
@@ -19,13 +20,13 @@ use ratatui::widgets::{Block, Borders, Paragraph};
use ratatui::Frame;
use ratatui::Terminal;
use crate::config::ConfigFile;
use crate::config::{ConfigFile, ServerEntry};
/// Outcome of the setup wizard, returned to the caller.
pub enum SetupOutcome {
/// Config saved; systemd service installed and started.
ServiceInstalled,
/// Config saved; no service — caller should start the proxy directly.
/// Config saved; no service -- caller should start the proxy directly.
ReadyToRun(PathBuf),
/// User quit without saving.
Cancelled,
@@ -34,13 +35,12 @@ pub enum SetupOutcome {
/// Column width reserved for the field label (chars).
const LABEL_WIDTH: usize = 22;
// ── Field types ──────────────────────────────────────────────────────────────
// -- Field types --------------------------------------------------------------
#[derive(Clone, Copy, PartialEq)]
enum FieldKind {
Text,
Secret,
Number,
Bool,
LogLevel,
}
@@ -53,31 +53,15 @@ struct Field {
required: bool,
help: &'static str,
}
// -- Server tab ---------------------------------------------------------------
// ── App state ────────────────────────────────────────────────────────────────
#[derive(PartialEq)]
enum Mode {
Normal,
Editing,
}
struct App {
/// A single server tab's editable fields.
struct ServerTab {
fields: Vec<Field>,
selected: usize,
mode: Mode,
edit_buffer: String,
edit_cursor: usize, // char index
config_path: PathBuf,
modified: bool,
message: Option<(String, Instant, bool)>, // (text, when, is_error)
scroll_offset: usize,
saved_once: bool,
pending_quit: bool, // true after first q/Esc with unsaved changes
}
impl App {
fn new(config_path: PathBuf) -> Self {
impl ServerTab {
fn new() -> Self {
Self {
fields: vec![
Field {
@@ -86,7 +70,7 @@ impl App {
value: String::new(),
kind: FieldKind::Text,
required: true,
help: "Aether 服务器 URL (如 https://aether.example.com)",
help: "Aether URL (e.g. https://aether.example.com)",
},
Field {
label: "Management Token",
@@ -94,23 +78,7 @@ impl App {
value: String::new(),
kind: FieldKind::Secret,
required: true,
help: "Aether 管理 API Token (ae_xxx)",
},
Field {
label: "HMAC Key",
key: "hmac_key",
value: String::new(),
kind: FieldKind::Secret,
required: true,
help: "HMAC-SHA256 签名密钥,用于代理请求认证",
},
Field {
label: "Listen Port",
key: "listen_port",
value: "18080".into(),
kind: FieldKind::Number,
required: true,
help: "代理服务监听端口",
help: "Aether Management Token (ae_xxx)",
},
Field {
label: "Node Name",
@@ -118,15 +86,60 @@ impl App {
value: "proxy-01".into(),
kind: FieldKind::Text,
required: true,
help: "节点名称,用于在 Aether 后台识别",
help: "Node name for identification in Aether dashboard",
},
],
}
}
fn from_entry(entry: &ServerEntry) -> Self {
let mut tab = Self::new();
tab.fields[0].value = entry.aether_url.clone();
tab.fields[1].value = entry.management_token.clone();
if let Some(ref name) = entry.node_name {
tab.fields[2].value = name.clone();
}
tab
}
}
// -- App state ----------------------------------------------------------------
#[derive(PartialEq)]
enum Mode {
Normal,
Editing,
}
struct App {
server_tabs: Vec<ServerTab>,
active_tab: usize,
global_fields: Vec<Field>,
selected: usize,
mode: Mode,
edit_buffer: String,
edit_cursor: usize,
config_path: PathBuf,
modified: bool,
message: Option<(String, Instant, bool)>,
scroll_offset: usize,
saved_once: bool,
pending_quit: bool,
confirm_delete: bool,
}
impl App {
fn new(config_path: PathBuf) -> Self {
Self {
server_tabs: vec![ServerTab::new()],
active_tab: 0,
global_fields: vec![
Field {
label: "Log Level",
key: "log_level",
value: "info".into(),
kind: FieldKind::LogLevel,
required: true,
help: "日志级别 -- Enter 切换: trace / debug / info / warn / error",
help: "Log level -- Enter to cycle: trace / debug / info / warn / error",
},
Field {
label: "Log JSON",
@@ -134,7 +147,7 @@ impl App {
value: "false".into(),
kind: FieldKind::Bool,
required: true,
help: "是否以 JSON 格式输出日志 -- Enter 切换",
help: "Output logs as JSON -- Enter to toggle",
},
Field {
label: "Install Service",
@@ -147,7 +160,7 @@ impl App {
.into(),
kind: FieldKind::Bool,
required: true,
help: "注册为 systemd 开机启动服务 (需要 root 权限) -- Enter 切换",
help: "Install as systemd service (requires root) -- Enter to toggle",
},
],
selected: 0,
@@ -160,10 +173,47 @@ impl App {
scroll_offset: 0,
saved_once: false,
pending_quit: false,
confirm_delete: false,
}
}
// ── Config ↔ fields ──────────────────────────────────────────────────
// -- Field accessors (unified index across server + global) ---------------
fn server_field_count(&self) -> usize {
self.server_tabs[self.active_tab].fields.len()
}
fn total_field_count(&self) -> usize {
self.server_field_count() + self.global_fields.len()
}
fn selected_field(&self) -> &Field {
let sc = self.server_field_count();
if self.selected < sc {
&self.server_tabs[self.active_tab].fields[self.selected]
} else {
&self.global_fields[self.selected - sc]
}
}
fn selected_field_mut(&mut self) -> &mut Field {
let sc = self.server_field_count();
if self.selected < sc {
&mut self.server_tabs[self.active_tab].fields[self.selected]
} else {
&mut self.global_fields[self.selected - sc]
}
}
fn clamp_selection(&mut self) {
let max = self.total_field_count();
if self.selected >= max {
self.selected = max.saturating_sub(1);
}
self.scroll_offset = 0;
self.confirm_delete = false;
}
// -- Config <-> fields -----------------------------------------------------
fn load_from_file(&mut self) {
if let Ok(cfg) = ConfigFile::load(&self.config_path) {
@@ -172,13 +222,9 @@ impl App {
}
fn apply_config(&mut self, cfg: &ConfigFile) {
for field in &mut self.fields {
// Global fields
for field in &mut self.global_fields {
let val: Option<String> = match field.key {
"aether_url" => cfg.aether_url.clone(),
"management_token" => cfg.management_token.clone(),
"hmac_key" => cfg.hmac_key.clone(),
"listen_port" => cfg.listen_port.map(|v| v.to_string()),
"node_name" => cfg.node_name.clone(),
"log_level" => cfg.log_level.clone(),
"log_json" => cfg.log_json.map(|v| v.to_string()),
_ => None,
@@ -187,59 +233,76 @@ impl App {
field.value = v;
}
}
// Server tabs
let servers = cfg.effective_servers();
if servers.is_empty() {
let mut tab = ServerTab::new();
// Single-server fallback: use top-level node_name
if let Some(ref name) = cfg.node_name {
tab.fields[2].value = name.clone();
}
self.server_tabs = vec![tab];
} else {
self.server_tabs = servers.iter().map(ServerTab::from_entry).collect();
// For single-server mode, node_name might be in top-level only
if self.server_tabs.len() == 1 && self.server_tabs[0].fields[2].value.is_empty() {
if let Some(ref name) = cfg.node_name {
self.server_tabs[0].fields[2].value = name.clone();
}
}
}
self.active_tab = 0;
self.selected = 0;
self.scroll_offset = 0;
}
fn to_config(&self) -> ConfigFile {
let get = |key: &str| -> Option<String> {
self.fields
let get_global = |key: &str| -> Option<String> {
self.global_fields
.iter()
.find(|f| f.key == key)
.map(|f| f.value.clone())
.filter(|v| !v.is_empty())
};
ConfigFile {
aether_url: get("aether_url"),
management_token: get("management_token"),
hmac_key: get("hmac_key"),
listen_port: get("listen_port").and_then(|v| v.parse().ok()),
public_ip: None,
node_name: get("node_name"),
node_region: None,
heartbeat_interval: None,
allowed_ports: None,
timestamp_tolerance: None,
aether_request_timeout_secs: None,
aether_connect_timeout_secs: None,
aether_pool_max_idle_per_host: None,
aether_pool_idle_timeout_secs: None,
aether_tcp_keepalive_secs: None,
aether_tcp_nodelay: None,
aether_http2: None,
aether_retry_max_attempts: None,
aether_retry_base_delay_ms: None,
aether_retry_max_delay_ms: None,
max_concurrent_connections: None,
connect_timeout_secs: None,
tls_handshake_timeout_secs: None,
dns_cache_ttl_secs: None,
dns_cache_capacity: None,
delegate_connect_timeout_secs: None,
delegate_pool_max_idle_per_host: None,
delegate_pool_idle_timeout_secs: None,
delegate_tcp_keepalive_secs: None,
delegate_tcp_nodelay: None,
log_level: get("log_level"),
log_json: get("log_json").and_then(|v| v.parse().ok()),
enable_tls: None,
tls_cert: None,
tls_key: None,
}
let get_tab = |tab: &ServerTab, key: &str| -> Option<String> {
tab.fields
.iter()
.find(|f| f.key == key)
.map(|f| f.value.clone())
.filter(|v| !v.is_empty())
};
let mut cfg = ConfigFile {
log_level: get_global("log_level"),
log_json: get_global("log_json").and_then(|v| v.parse().ok()),
..ConfigFile::default()
};
// Always write [[servers]] format; old top-level fields are read-only compat
cfg.servers = self
.server_tabs
.iter()
.map(|tab| ServerEntry {
aether_url: get_tab(tab, "aether_url").unwrap_or_default(),
management_token: get_tab(tab, "management_token").unwrap_or_default(),
node_name: get_tab(tab, "node_name"),
})
.collect();
cfg
}
fn save(&mut self) -> anyhow::Result<()> {
let cfg = self.to_config();
cfg.save(&self.config_path)?;
// Restrict config file permissions to owner-only (contains management token).
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
let _ =
std::fs::set_permissions(&self.config_path, std::fs::Permissions::from_mode(0o600));
}
self.modified = false;
self.saved_once = true;
self.message = Some((
@@ -249,27 +312,33 @@ impl App {
));
Ok(())
}
// ── Scrolling ────────────────────────────────────────────────────────
// -- Scrolling ---------------------------------------------------------------
fn ensure_visible(&mut self, visible_rows: usize) {
if visible_rows == 0 {
return;
}
if self.selected < self.scroll_offset {
self.scroll_offset = self.selected;
} else if self.selected >= self.scroll_offset + visible_rows {
self.scroll_offset = self.selected - visible_rows + 1;
// Account for separator line between server and global fields
let display_row = if self.selected >= self.server_field_count() {
self.selected + 1
} else {
self.selected
};
if display_row < self.scroll_offset {
self.scroll_offset = display_row;
} else if display_row >= self.scroll_offset + visible_rows {
self.scroll_offset = display_row - visible_rows + 1;
}
}
// ── Key handling ─────────────────────────────────────────────────────
// -- Key handling -------------------------------------------------------------
/// Returns `true` when the app should exit.
fn handle_key(&mut self, key: KeyEvent) -> bool {
// Expire old messages (but keep quit-confirmation messages alive)
if let Some((_, when, _)) = &self.message {
if !self.pending_quit && when.elapsed() > Duration::from_secs(4) {
if !self.pending_quit && !self.confirm_delete && when.elapsed() > Duration::from_secs(4)
{
self.message = None;
}
}
@@ -284,7 +353,7 @@ impl App {
}
fn handle_normal(&mut self, key: KeyEvent) -> bool {
// ── Quit handling (with unsaved-changes confirmation) ─────────
// -- Quit handling (with unsaved-changes confirmation) -----------------
let is_quit_key = matches!(key.code, KeyCode::Char('q') | KeyCode::Esc);
if is_quit_key {
@@ -292,6 +361,7 @@ impl App {
return true;
}
self.pending_quit = true;
self.confirm_delete = false;
self.message = Some((
"unsaved changes! q again to discard, ^S to save".into(),
Instant::now(),
@@ -300,14 +370,21 @@ impl App {
return false;
}
// Any other key cancels the pending quit
// Any other key cancels pending quit / pending delete
if self.pending_quit {
self.pending_quit = false;
self.message = None;
}
if self.confirm_delete && !matches!(key.code, KeyCode::Delete | KeyCode::Char('x')) {
self.confirm_delete = false;
self.message = None;
}
match key.code {
KeyCode::Char('s') if key.modifiers.contains(KeyModifiers::CONTROL) => {
KeyCode::Char('s')
if key.modifiers.contains(KeyModifiers::CONTROL)
|| key.modifiers.contains(KeyModifiers::SUPER) =>
{
if let Err(e) = self.save() {
self.message = Some((format!("error: {}", e), Instant::now(), true));
}
@@ -316,23 +393,20 @@ impl App {
self.selected = self.selected.saturating_sub(1);
}
KeyCode::Down | KeyCode::Char('j') => {
if self.selected + 1 < self.fields.len() {
if self.selected + 1 < self.total_field_count() {
self.selected += 1;
}
}
KeyCode::Home => self.selected = 0,
KeyCode::End => self.selected = self.fields.len() - 1,
KeyCode::End => self.selected = self.total_field_count() - 1,
KeyCode::Enter | KeyCode::Char(' ') => {
let field = &self.fields[self.selected];
match field.kind {
let kind = self.selected_field().kind;
let key_str = self.selected_field().key;
let value = self.selected_field().value.clone();
match kind {
FieldKind::Bool => {
let toggled = if field.value == "true" {
"false"
} else {
"true"
};
// Block enabling service install without root/systemd
if field.key == "install_service"
let toggled = if value == "true" { "false" } else { "true" };
if key_str == "install_service"
&& toggled == "true"
&& !super::service::is_available()
{
@@ -342,27 +416,79 @@ impl App {
true,
));
} else {
self.fields[self.selected].value = toggled.into();
self.selected_field_mut().value = toggled.into();
self.modified = true;
}
}
FieldKind::LogLevel => {
const LEVELS: &[&str] = &["trace", "debug", "info", "warn", "error"];
let idx = LEVELS.iter().position(|l| *l == field.value).unwrap_or(2);
self.fields[self.selected].value = LEVELS[(idx + 1) % LEVELS.len()].into();
let idx = LEVELS.iter().position(|l| *l == value).unwrap_or(2);
self.selected_field_mut().value = LEVELS[(idx + 1) % LEVELS.len()].into();
self.modified = true;
}
_ => {
self.edit_buffer = field.value.clone();
self.edit_buffer = value;
self.edit_cursor = self.edit_buffer.chars().count();
self.mode = Mode::Editing;
}
}
}
// -- Tab navigation --
KeyCode::Tab => {
// Quick save shortcut
if let Err(e) = self.save() {
self.message = Some((format!("error: {}", e), Instant::now(), true));
if self.server_tabs.len() > 1 {
self.active_tab = (self.active_tab + 1) % self.server_tabs.len();
self.clamp_selection();
}
}
KeyCode::BackTab => {
if self.server_tabs.len() > 1 {
self.active_tab = if self.active_tab == 0 {
self.server_tabs.len() - 1
} else {
self.active_tab - 1
};
self.clamp_selection();
}
}
KeyCode::Char(c @ '1'..='9') if !key.modifiers.contains(KeyModifiers::CONTROL) => {
let idx = (c as usize) - ('1' as usize);
if idx < self.server_tabs.len() && idx != self.active_tab {
self.active_tab = idx;
self.clamp_selection();
}
}
// -- Add / remove server --
KeyCode::Char('+') | KeyCode::Char('a') => {
self.server_tabs.push(ServerTab::new());
self.active_tab = self.server_tabs.len() - 1;
self.selected = 0;
self.scroll_offset = 0;
self.modified = true;
self.message = Some((
format!("added server {}", self.server_tabs.len()),
Instant::now(),
false,
));
}
KeyCode::Delete | KeyCode::Char('x') => {
if self.server_tabs.len() <= 1 {
self.message =
Some(("cannot remove the last server".into(), Instant::now(), true));
} else if self.confirm_delete {
let removed = self.active_tab + 1;
self.server_tabs.remove(self.active_tab);
self.active_tab = self.active_tab.min(self.server_tabs.len() - 1);
self.clamp_selection();
self.modified = true;
self.message =
Some((format!("server {} removed", removed), Instant::now(), false));
} else {
self.confirm_delete = true;
self.message = Some((
"press Delete/x again to remove this server".into(),
Instant::now(),
true,
));
}
}
_ => {}
@@ -373,12 +499,11 @@ impl App {
fn handle_edit(&mut self, key: KeyEvent) {
match key.code {
KeyCode::Esc => {
// Cancel -- discard changes to this field
self.mode = Mode::Normal;
}
KeyCode::Enter => {
if self.validate_edit() {
self.fields[self.selected].value = self.edit_buffer.clone();
self.selected_field_mut().value = self.edit_buffer.clone();
self.modified = true;
self.mode = Mode::Normal;
} else {
@@ -419,12 +544,7 @@ impl App {
}
fn validate_edit(&self) -> bool {
let kind = self.fields[self.selected].kind;
let buf = &self.edit_buffer;
match kind {
FieldKind::Number => buf.is_empty() || buf.parse::<u64>().is_ok(),
_ => true,
}
true
}
/// Byte offset of the char at `char_idx`.
@@ -436,13 +556,11 @@ impl App {
.unwrap_or(self.edit_buffer.len())
}
}
// ── Rendering ────────────────────────────────────────────────────────────────
// -- Rendering ----------------------------------------------------------------
fn ui(f: &mut Frame, app: &mut App) {
let area = f.area();
// Outer block
let title = if app.modified {
" Aether Proxy Setup [*] "
} else {
@@ -458,28 +576,82 @@ fn ui(f: &mut Frame, app: &mut App) {
let inner = outer.inner(area);
f.render_widget(outer, area);
// Split: fields | footer
let chunks = Layout::vertical([Constraint::Min(1), Constraint::Length(4)]).split(inner);
// Split: fields | tab bar | footer
let chunks = Layout::vertical([
Constraint::Min(1),
Constraint::Length(1),
Constraint::Length(4),
])
.split(inner);
let fields_area = chunks[0];
let footer_area = chunks[1];
render_fields(f, app, fields_area);
render_footer(f, app, footer_area);
render_fields(f, app, chunks[0]);
render_tab_bar(f, app, chunks[1]);
render_footer(f, app, chunks[2]);
}
fn render_fields(f: &mut Frame, app: &mut App, area: Rect) {
let visible = area.height as usize;
app.ensure_visible(visible);
let server_count = app.server_field_count();
let mut lines: Vec<Line> = Vec::new();
// display_row tracks the actual row index (including separator)
let mut display_row: usize = 0;
for (i, field) in app.fields.iter().enumerate() {
if i < app.scroll_offset || i >= app.scroll_offset + visible {
continue;
// Server fields
for i in 0..server_count {
if display_row >= app.scroll_offset && display_row < app.scroll_offset + visible {
lines.push(build_field_line(app, i, display_row));
}
display_row += 1;
}
let selected = i == app.selected;
// Separator line
if display_row >= app.scroll_offset && display_row < app.scroll_offset + visible {
lines.push(Line::from(Span::styled(
" ----------------------------------------",
Style::default().fg(Color::DarkGray),
)));
}
display_row += 1;
// Global fields
for i in 0..app.global_fields.len() {
let field_idx = server_count + i;
if display_row >= app.scroll_offset && display_row < app.scroll_offset + visible {
lines.push(build_field_line(app, field_idx, display_row));
}
display_row += 1;
}
let paragraph = Paragraph::new(lines);
f.render_widget(paragraph, area);
// Cursor position while editing
if app.mode == Mode::Editing {
let sel_display_row = if app.selected >= server_count {
app.selected + 1
} else {
app.selected
};
let row_in_view = sel_display_row.saturating_sub(app.scroll_offset);
let prefix: u16 = 3 + LABEL_WIDTH as u16 + 2;
let cx = area.x + prefix + app.edit_cursor as u16;
let cy = area.y + row_in_view as u16;
if cx < area.x + area.width && cy < area.y + area.height {
f.set_cursor_position((cx, cy));
}
}
}
fn build_field_line(app: &App, field_idx: usize, _display_row: usize) -> Line<'static> {
let sc = app.server_field_count();
let field = if field_idx < sc {
&app.server_tabs[app.active_tab].fields[field_idx]
} else {
&app.global_fields[field_idx - sc]
};
let selected = field_idx == app.selected;
let indicator = if selected { " > " } else { " " };
let label_style = if selected {
@@ -492,35 +664,18 @@ fn render_fields(f: &mut Frame, app: &mut App, area: Rect) {
let padded_label = format!("{:<width$}", field.label, width = LABEL_WIDTH);
// Value display
let (value_text, value_style) = if app.mode == Mode::Editing && selected {
(app.edit_buffer.clone(), Style::default().fg(Color::Yellow))
} else {
field_display(field)
};
lines.push(Line::from(vec![
Span::styled(indicator, label_style),
Line::from(vec![
Span::styled(indicator.to_string(), label_style),
Span::styled(padded_label, label_style),
Span::raw(" "),
Span::styled(value_text, value_style),
]));
}
let paragraph = Paragraph::new(lines);
f.render_widget(paragraph, area);
// Cursor position while editing
if app.mode == Mode::Editing {
let row_in_view = app.selected - app.scroll_offset;
// prefix: 3 (indicator) + LABEL_WIDTH + 2 (gap) = 27
let prefix: u16 = 3 + LABEL_WIDTH as u16 + 2;
let cx = area.x + prefix + app.edit_cursor as u16;
let cy = area.y + row_in_view as u16;
if cx < area.x + area.width && cy < area.y + area.height {
f.set_cursor_position((cx, cy));
}
}
])
}
/// Returns (display_text, style) for a field in normal mode.
@@ -565,18 +720,54 @@ fn field_display(field: &Field) -> (String, Style) {
_ => (field.value.clone(), Style::default().fg(Color::White)),
}
}
fn render_tab_bar(f: &mut Frame, app: &App, area: Rect) {
let mut spans: Vec<Span> = Vec::new();
spans.push(Span::raw(" "));
for (i, tab) in app.server_tabs.iter().enumerate() {
let num = i + 1;
let name = tab
.fields
.iter()
.find(|f| f.key == "node_name")
.filter(|f| !f.value.is_empty())
.map(|f| f.value.clone())
.unwrap_or_else(|| format!("Server {}", num));
let label = format!(" {} {} ", num, name);
if i == app.active_tab {
spans.push(Span::styled(
label,
Style::default()
.fg(Color::Black)
.bg(Color::Cyan)
.add_modifier(Modifier::BOLD),
));
} else {
spans.push(Span::styled(label, Style::default().fg(Color::DarkGray)));
}
spans.push(Span::raw(" "));
}
spans.push(Span::styled(" + Add ", Style::default().fg(Color::Green)));
f.render_widget(Paragraph::new(Line::from(spans)), area);
}
fn render_footer(f: &mut Frame, app: &App, area: Rect) {
let help = app.fields[app.selected].help;
let help = app.selected_field().help;
let keybindings = if app.mode == Mode::Editing {
"Enter confirm Esc cancel"
} else if app.server_tabs.len() > 1 {
"j/k select Enter edit Tab switch + add x remove ^S save q quit"
} else {
"Up/Down select Enter edit ^S save q quit"
"j/k select Enter edit + add server ^S save q quit"
};
let mut status_spans: Vec<Span> = vec![Span::styled(
keybindings,
format!(" {}", keybindings),
Style::default().fg(Color::DarkGray),
)];
@@ -592,18 +783,7 @@ fn render_footer(f: &mut Frame, app: &App, area: Rect) {
format!(" {}", help),
Style::default().fg(Color::DarkGray),
)),
Line::from(
status_spans
.into_iter()
.map(|mut s| {
// add left padding to first span
if s.content.as_ref() == keybindings {
s.content = format!(" {}", s.content).into();
}
s
})
.collect::<Vec<_>>(),
),
Line::from(status_spans),
];
let footer = Paragraph::new(footer_text).block(
@@ -614,11 +794,9 @@ fn render_footer(f: &mut Frame, app: &App, area: Rect) {
f.render_widget(footer, area);
}
// ── Entry point ──────────────────────────────────────────────────────────────
// -- Entry point --------------------------------------------------------------
pub fn run(config_path: PathBuf) -> anyhow::Result<SetupOutcome> {
// Setup terminal
terminal::enable_raw_mode()?;
let mut stdout = io::stdout();
execute!(stdout, EnterAlternateScreen)?;
@@ -630,14 +808,13 @@ pub fn run(config_path: PathBuf) -> anyhow::Result<SetupOutcome> {
let result = event_loop(&mut terminal, &mut app);
// Restore terminal
terminal::disable_raw_mode()?;
execute!(terminal.backend_mut(), LeaveAlternateScreen)?;
terminal.show_cursor()?;
result?;
// ── Post-TUI: decide outcome ─────────────────────────────────────
// -- Post-TUI: decide outcome ---------------------------------------------
if !app.saved_once {
return Ok(SetupOutcome::Cancelled);
@@ -648,7 +825,7 @@ pub fn run(config_path: PathBuf) -> anyhow::Result<SetupOutcome> {
eprintln!();
let wants_service = app
.fields
.global_fields
.iter()
.find(|f| f.key == "install_service")
.map(|f| f.value == "true")
@@ -662,15 +839,12 @@ pub fn run(config_path: PathBuf) -> anyhow::Result<SetupOutcome> {
eprintln!(" Starting proxy directly instead.\n");
}
}
} else {
// Uninstall service if it was previously installed but toggled off
if super::service::is_installed() {
} else if super::service::is_installed() {
if let Err(e) = super::service::uninstall_service() {
eprintln!(" Service uninstall failed: {}", e);
eprintln!();
}
}
}
Ok(SetupOutcome::ReadyToRun(config_path))
}
@@ -684,7 +858,6 @@ fn event_loop(
if event::poll(Duration::from_millis(200))? {
if let Event::Key(key) = event::read()? {
// Only handle Press events (ignore Release on Windows)
if key.kind == KeyEventKind::Press && app.handle_key(key) {
break;
}
+53 -8
View File
@@ -281,8 +281,17 @@ fn atomic_replace(new_binary: &Path) -> anyhow::Result<PathBuf> {
// ── Public entry point ───────────────────────────────────────────────────────
/// `aether-proxy upgrade [version]` -- self-upgrade from GitHub releases.
pub async fn cmd_upgrade(version: Option<String>) -> anyhow::Result<()> {
#[derive(Clone, Copy)]
enum RestartMode {
BestEffort,
Required,
}
async fn execute_upgrade(
version: Option<&str>,
require_root: bool,
restart_mode: RestartMode,
) -> anyhow::Result<()> {
// Resolve exe path once; reuse throughout the function
let current_exe = std::env::current_exe()?.canonicalize()?;
let exe_dir = current_exe
@@ -290,8 +299,12 @@ pub async fn cmd_upgrade(version: Option<String>) -> anyhow::Result<()> {
.ok_or_else(|| anyhow::anyhow!("cannot determine binary directory"))?;
let temp_path = exe_dir.join(".aether-proxy.upgrade.tmp");
// Check write permission to binary directory
if require_root {
if !super::service::is_root() {
anyhow::bail!("automatic upgrade requires root privileges");
}
} else if !super::service::is_root() {
// Check write permission to binary directory for manual upgrade mode.
let test_path = exe_dir.join(".aether-proxy.write-test");
match std::fs::File::create(&test_path) {
Ok(_) => {
@@ -311,7 +324,7 @@ pub async fn cmd_upgrade(version: Option<String>) -> anyhow::Result<()> {
eprintln!(" Current version: {}", CURRENT_VERSION);
let client = build_github_client()?;
let release = fetch_release(&client, version.as_deref()).await?;
let release = fetch_release(&client, version).await?;
let target_tag = &release.tag_name;
let target_semver = target_tag.strip_prefix("proxy-v").unwrap_or(target_tag);
@@ -341,20 +354,39 @@ pub async fn cmd_upgrade(version: Option<String>) -> anyhow::Result<()> {
}
};
// Restart systemd service if running
match restart_mode {
RestartMode::BestEffort => {
// Restart systemd service if running.
// Use best-effort: binary is already replaced, so a restart failure should
// not abort the whole upgrade -- the user can restart manually.
if super::service::is_service_active() {
if super::service::is_root() {
eprintln!(" Restarting systemd service...");
super::service::run_cmd("systemctl", &["restart", "aether-proxy"])?;
eprintln!(" Service restarted.");
match super::service::run_cmd("systemctl", &["restart", "aether-proxy"]) {
Ok(()) => eprintln!(" Service restarted."),
Err(e) => {
eprintln!(" WARNING: failed to restart service: {}", e);
eprintln!(" Run manually: sudo systemctl restart aether-proxy");
}
}
} else {
eprintln!(" Systemd service is active, but restart requires root.");
eprintln!(" Run: sudo aether-proxy restart");
eprintln!(" Run: sudo systemctl restart aether-proxy");
eprintln!(" Skipping restart.");
}
} else {
eprintln!(" No active systemd service detected, skipping restart.");
}
}
RestartMode::Required => {
if !super::service::is_root() {
anyhow::bail!("automatic upgrade requires root privileges");
}
eprintln!(" Restarting systemd service...");
super::service::run_cmd("systemctl", &["restart", "aether-proxy"])?;
eprintln!(" Service restarted.");
}
}
eprintln!();
eprintln!(" Upgrade complete!");
@@ -364,3 +396,16 @@ pub async fn cmd_upgrade(version: Option<String>) -> anyhow::Result<()> {
);
Ok(())
}
/// `aether-proxy upgrade [version]` -- self-upgrade from GitHub releases.
pub async fn cmd_upgrade(version: Option<String>) -> anyhow::Result<()> {
execute_upgrade(version.as_deref(), false, RestartMode::BestEffort).await
}
/// Perform automatic upgrade to a specific version.
///
/// This path is designed for server-pushed upgrades in systemd/root scenarios:
/// it requires root and requires a successful `systemctl restart aether-proxy`.
pub async fn perform_upgrade(version: &str) -> anyhow::Result<()> {
execute_upgrade(Some(version), true, RestartMode::Required).await
}
+45 -29
View File
@@ -1,48 +1,59 @@
//! Shared application state passed to all subsystems.
//!
//! Consolidates the multiple `Arc<...>` parameters that were previously
//! threaded individually through proxy server, heartbeat, and handlers.
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, RwLock};
use std::time::Duration;
use tokio::sync::Semaphore;
use tokio_rustls::TlsAcceptor;
use crate::config::Config;
use crate::hardware::HardwareInfo;
use crate::proxy::delegate_client::DelegateClient;
use crate::proxy::target_filter::DnsCache;
use crate::registration::client::AetherClient;
use crate::runtime::SharedDynamicConfig;
use crate::target_filter::DnsCache;
use crate::upstream_client::UpstreamClient;
/// Central application state shared across all tasks.
/// Central application state shared across all servers/tunnels.
pub struct AppState {
pub config: Arc<Config>,
pub node_id: Arc<RwLock<String>>,
pub dynamic: SharedDynamicConfig,
pub aether_client: Arc<AetherClient>,
pub hardware_info: Arc<HardwareInfo>,
pub public_ip: String,
pub tls_fingerprint: Option<String>,
pub tls_acceptor: Option<TlsAcceptor>,
/// Shared delegate client for proxy-initiated upstream requests.
pub delegate_client: DelegateClient,
/// Active connection count for metrics reporting.
pub active_connections: Arc<AtomicU64>,
/// Connection concurrency limiter.
pub connection_semaphore: Arc<Semaphore>,
/// DNS cache for upstream target resolution.
/// DNS cache for upstream target resolution (shared).
pub dns_cache: Arc<DnsCache>,
/// Request/latency metrics for heartbeat.
/// Hyper client for tunnel upstream requests with validated DNS and connection timing.
pub upstream_client: UpstreamClient,
/// Shared TLS config for tunnel WebSocket connections (avoids re-parsing root CAs on each reconnect).
pub tunnel_tls_config: Arc<rustls::ClientConfig>,
}
/// Per-server state: one instance per Aether server connection.
pub struct ServerContext {
/// Human-readable label for logging (e.g. "server-0").
pub server_label: String,
/// Aether server URL for this connection.
pub aether_url: String,
/// Management token for this server.
pub management_token: String,
/// Resolved node name at registration time (per-server override or global fallback).
/// After startup, the active node_name is read from `dynamic` (may be updated remotely).
#[allow(dead_code)]
pub node_name: String,
/// Node ID assigned by this Aether server.
pub node_id: Arc<RwLock<String>>,
/// API client for this server.
pub aether_client: Arc<AetherClient>,
/// Dynamic config from this server's heartbeat ACKs.
pub dynamic: SharedDynamicConfig,
/// Per-server active connection count.
pub active_connections: Arc<AtomicU64>,
/// Per-server request/latency metrics.
pub metrics: Arc<ProxyMetrics>,
}
/// Aggregate metrics for reporting to Aether.
pub struct ProxyMetrics {
pub total_requests: AtomicU64,
/// Cumulative connection-establishment latency in nanoseconds
/// (DNS + TCP/TLS + TTFB, excludes response body streaming).
pub total_latency_ns: AtomicU64,
pub failed_requests: AtomicU64,
pub dns_failures: AtomicU64,
pub stream_errors: AtomicU64,
}
impl ProxyMetrics {
@@ -50,12 +61,17 @@ impl ProxyMetrics {
Self {
total_requests: AtomicU64::new(0),
total_latency_ns: AtomicU64::new(0),
failed_requests: AtomicU64::new(0),
dns_failures: AtomicU64::new(0),
stream_errors: AtomicU64::new(0),
}
}
pub fn record_request(&self, elapsed: Duration) {
let nanos = u64::try_from(elapsed.as_nanos()).unwrap_or(u64::MAX);
self.total_requests.fetch_add(1, Ordering::Relaxed);
self.total_latency_ns.fetch_add(nanos, Ordering::Relaxed);
/// Record a completed request with its connection-establishment latency
/// (DNS + TCP/TLS + TTFB, excludes response body streaming).
pub fn record_request(&self, connect_elapsed: Duration) {
let nanos = u64::try_from(connect_elapsed.as_nanos()).unwrap_or(u64::MAX);
self.total_requests.fetch_add(1, Ordering::Release);
self.total_latency_ns.fetch_add(nanos, Ordering::Release);
}
}
@@ -1,11 +1,12 @@
use std::collections::{HashMap, HashSet};
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::RwLock;
/// Check if an IP address belongs to a private/reserved network.
fn is_private_ip(ip: &IpAddr) -> bool {
pub fn is_private_ip(ip: &IpAddr) -> bool {
match ip {
IpAddr::V4(v4) => is_private_ipv4(v4),
IpAddr::V6(v6) => is_private_ipv6(v6),
@@ -38,6 +39,22 @@ fn is_private_ipv4(ip: &Ipv4Addr) -> bool {
if octets[0] == 0 {
return true;
}
// 100.64.0.0/10 (CGNAT / shared address space)
if octets[0] == 100 && (64..=127).contains(&octets[1]) {
return true;
}
// 192.0.0.0/24 (IETF protocol assignments)
if octets[0] == 192 && octets[1] == 0 && octets[2] == 0 {
return true;
}
// 198.18.0.0/15 (benchmark testing)
if octets[0] == 198 && (18..=19).contains(&octets[1]) {
return true;
}
// 240.0.0.0/4 (reserved for future use)
if octets[0] >= 240 {
return true;
}
false
}
@@ -71,6 +88,7 @@ pub enum FilterError {
PrivateIp(IpAddr),
PortNotAllowed(u16),
DnsResolutionFailed(String),
NoPublicAddrs(String),
}
impl std::fmt::Display for FilterError {
@@ -79,17 +97,26 @@ impl std::fmt::Display for FilterError {
Self::PrivateIp(ip) => write!(f, "target IP {} is in private/reserved range", ip),
Self::PortNotAllowed(port) => write!(f, "port {} not in allowed list", port),
Self::DnsResolutionFailed(host) => write!(f, "DNS resolution failed for {}", host),
Self::NoPublicAddrs(host) => {
write!(
f,
"all resolved addresses for {} are private/reserved",
host
)
}
}
}
}
struct DnsCacheEntry {
addr: SocketAddr,
addrs: Arc<Vec<SocketAddr>>,
expires_at: Instant,
inserted_at: Instant,
}
/// Lightweight DNS cache with TTL + capacity bounds.
/// Stores all public resolved addresses per host (used by SafeDnsResolver
/// to ensure reqwest connects to the same validated addresses).
pub struct DnsCache {
ttl: Duration,
capacity: usize,
@@ -105,7 +132,27 @@ impl DnsCache {
}
}
pub async fn get(&self, host: &str, port: u16) -> Option<SocketAddr> {
/// Look up cached public addresses for a host (any port).
///
/// Used by `SafeDnsResolver` which only knows the hostname — returns the
/// first unexpired entry whose key starts with `host:`.
pub async fn get_by_host(&self, host: &str) -> Option<Arc<Vec<SocketAddr>>> {
if self.capacity == 0 || self.ttl.is_zero() {
return None;
}
let prefix = format!("{}:", host.to_ascii_lowercase());
let now = Instant::now();
let entries = self.entries.read().await;
for (key, entry) in entries.iter() {
if key.starts_with(&prefix) && entry.expires_at > now {
return Some(Arc::clone(&entry.addrs));
}
}
None
}
/// Look up cached public addresses for a host + port.
pub async fn get(&self, host: &str, port: u16) -> Option<Arc<Vec<SocketAddr>>> {
if self.capacity == 0 || self.ttl.is_zero() {
return None;
}
@@ -116,7 +163,7 @@ impl DnsCache {
{
let entries = self.entries.read().await;
match entries.get(&key) {
Some(entry) if entry.expires_at > now => return Some(entry.addr),
Some(entry) if entry.expires_at > now => return Some(Arc::clone(&entry.addrs)),
None => return None,
Some(_) => {} // expired, fall through to evict
}
@@ -128,8 +175,9 @@ impl DnsCache {
None
}
pub async fn insert(&self, host: &str, port: u16, addr: SocketAddr) {
if self.capacity == 0 || self.ttl.is_zero() {
/// Insert resolved public addresses into cache.
pub async fn insert(&self, host: &str, port: u16, addrs: Arc<Vec<SocketAddr>>) {
if self.capacity == 0 || self.ttl.is_zero() || addrs.is_empty() {
return;
}
let key = Self::key(host, port);
@@ -150,7 +198,7 @@ impl DnsCache {
entries.insert(
key,
DnsCacheEntry {
addr,
addrs,
expires_at: now + self.ttl,
inserted_at: now,
},
@@ -158,22 +206,62 @@ impl DnsCache {
}
fn key(host: &str, port: u16) -> String {
format!("{}:{}", host, port)
format!("{}:{}", host.to_ascii_lowercase(), port)
}
}
/// Resolve a hostname to public (non-private) socket addresses.
///
/// Results are cached in `dns_cache`. Private/reserved IPs are filtered out.
/// Returns an error if no public addresses remain after filtering.
pub async fn resolve_public_addrs(
host: &str,
port: u16,
dns_cache: &DnsCache,
) -> Result<Vec<SocketAddr>, FilterError> {
// Cache hit
if let Some(addrs) = dns_cache.get(host, port).await {
return Ok((*addrs).clone());
}
// Async DNS resolution
let addr_str = format!("{}:{}", host, port);
let resolved: Vec<SocketAddr> = tokio::net::lookup_host(&addr_str)
.await
.map_err(|_| FilterError::DnsResolutionFailed(host.to_string()))?
.collect();
if resolved.is_empty() {
return Err(FilterError::DnsResolutionFailed(host.to_string()));
}
// Filter out private/reserved addresses
let public: Vec<SocketAddr> = resolved
.into_iter()
.filter(|addr| !is_private_ip(&addr.ip()))
.collect();
if public.is_empty() {
return Err(FilterError::NoPublicAddrs(host.to_string()));
}
// Cache the validated public addresses
let arc_addrs = Arc::new(public);
dns_cache.insert(host, port, Arc::clone(&arc_addrs)).await;
Ok((*arc_addrs).clone())
}
/// Validate that the target host:port is allowed.
///
/// Uses async DNS resolution (via `tokio::net::lookup_host`) to avoid
/// blocking the async runtime on potentially slow DNS lookups.
///
/// Returns the resolved socket address to connect to.
/// Performs port whitelist check, private IP filtering, and DNS resolution
/// with caching. The resolved addresses are stored in the shared DnsCache
/// so that the SafeDnsResolver can reuse them, eliminating the TOCTTOU gap.
pub async fn validate_target(
host: &str,
port: u16,
allowed_ports: &HashSet<u16>,
dns_cache: &DnsCache,
) -> Result<SocketAddr, FilterError> {
) -> Result<Vec<SocketAddr>, FilterError> {
// Port whitelist check
if !allowed_ports.contains(&port) {
return Err(FilterError::PortNotAllowed(port));
@@ -184,35 +272,11 @@ pub async fn validate_target(
if is_private_ip(&ip) {
return Err(FilterError::PrivateIp(ip));
}
return Ok(SocketAddr::new(ip, port));
return Ok(vec![SocketAddr::new(ip, port)]);
}
if let Some(addr) = dns_cache.get(host, port).await {
return Ok(addr);
}
// Async DNS resolution with private IP check (DNS rebinding protection)
let addr_str = format!("{}:{}", host, port);
let addrs: Vec<SocketAddr> = tokio::net::lookup_host(&addr_str)
.await
.map_err(|_| FilterError::DnsResolutionFailed(host.to_string()))?
.collect();
if addrs.is_empty() {
return Err(FilterError::DnsResolutionFailed(host.to_string()));
}
// All resolved addresses must be non-private
for addr in &addrs {
if is_private_ip(&addr.ip()) {
return Err(FilterError::PrivateIp(addr.ip()));
}
}
// Return the first valid address
let selected = addrs[0];
dns_cache.insert(host, port, selected).await;
Ok(selected)
// Resolve and validate DNS (populates cache for SafeDnsResolver)
resolve_public_addrs(host, port, dns_cache).await
}
#[cfg(test)]
@@ -235,6 +299,19 @@ mod tests {
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1))));
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(169, 254, 1, 1))));
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(0, 0, 0, 0))));
// CGNAT
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(100, 64, 0, 1))));
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(
100, 127, 255, 254
))));
assert!(!is_private_ip(&IpAddr::V4(Ipv4Addr::new(
100, 63, 255, 254
))));
// Benchmark testing
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(198, 18, 0, 1))));
// Reserved
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(240, 0, 0, 1))));
// Public
assert!(!is_private_ip(&IpAddr::V4(Ipv4Addr::new(8, 8, 8, 8))));
assert!(!is_private_ip(&IpAddr::V4(Ipv4Addr::new(203, 0, 113, 1))));
}
@@ -272,5 +349,33 @@ mod tests {
let cache = cache();
let result = validate_target("8.8.8.8", 443, &ports(), &cache).await;
assert!(result.is_ok());
let addrs = result.unwrap();
assert_eq!(addrs.len(), 1);
assert_eq!(addrs[0].ip(), IpAddr::V4(Ipv4Addr::new(8, 8, 8, 8)));
}
#[tokio::test]
async fn test_cache_stores_multiple_addrs() {
let cache = cache();
let addrs = vec![
SocketAddr::new(IpAddr::V4(Ipv4Addr::new(1, 1, 1, 1)), 443),
SocketAddr::new(IpAddr::V4(Ipv4Addr::new(1, 0, 0, 1)), 443),
];
cache
.insert("example.com", 443, Arc::new(addrs.clone()))
.await;
let cached = cache.get("example.com", 443).await.unwrap();
assert_eq!(*cached, addrs);
}
#[tokio::test]
async fn test_cache_key_case_insensitive() {
let cache = cache();
let addrs = vec![SocketAddr::new(IpAddr::V4(Ipv4Addr::new(1, 1, 1, 1)), 443)];
cache
.insert("Example.COM", 443, Arc::new(addrs.clone()))
.await;
let cached = cache.get("example.com", 443).await.unwrap();
assert_eq!(*cached, addrs);
}
}
+238
View File
@@ -0,0 +1,238 @@
//! WebSocket tunnel client: connect, authenticate, and run the tunnel.
use std::sync::Arc;
use std::time::Duration;
use tokio::net::TcpStream;
use tokio::sync::watch;
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
use tokio_tungstenite::tungstenite::http;
use tokio_tungstenite::tungstenite::protocol::WebSocketConfig;
use tracing::{debug, info, warn};
use crate::state::{AppState, ServerContext};
use super::{dispatcher, heartbeat, writer};
/// Outcome of a tunnel session.
pub enum TunnelOutcome {
/// Graceful shutdown requested by the local process.
Shutdown,
/// Remote side disconnected or connection lost — should reconnect.
Disconnected,
}
/// Connect to Aether's WebSocket tunnel endpoint and run until disconnected.
///
/// `conn_idx` identifies which connection in the pool this is (0-based).
/// Only connection 0 sends heartbeats to avoid resetting shared metrics.
pub async fn connect_and_run(
state: &Arc<AppState>,
server: &Arc<ServerContext>,
conn_idx: usize,
shutdown: &mut watch::Receiver<bool>,
) -> Result<TunnelOutcome, anyhow::Error> {
let ws_url = build_tunnel_url(server);
info!(url = %ws_url, conn = conn_idx, "connecting tunnel");
// Build WebSocket request with auth headers
let mut request = ws_url.clone().into_client_request()?;
let headers = request.headers_mut();
headers.insert(
"Authorization",
http::HeaderValue::from_str(&format!("Bearer {}", server.management_token))?,
);
let node_id = server.node_id.read().unwrap().clone();
headers.insert("X-Node-Id", http::HeaderValue::from_str(&node_id)?);
// Use dynamic node_name (may be updated by remote config) instead of
// the static server.node_name, so that remote name changes take effect
// on the next reconnect.
let dynamic_node_name = server.dynamic.load().node_name.clone();
headers.insert(
"X-Node-Name",
http::HeaderValue::from_str(&dynamic_node_name)?,
);
// Advertise per-connection max concurrent streams so the backend can
// respect the proxy's capacity limit (backward-compatible: old backends
// ignore this header).
let max_streams = state.config.tunnel_max_streams.unwrap_or(128);
headers.insert("X-Tunnel-Max-Streams", http::HeaderValue::from(max_streams));
// Parse host:port from URL
let uri: http::Uri = ws_url.parse()?;
let host = uri
.host()
.ok_or_else(|| anyhow::anyhow!("missing host in tunnel URL"))?;
let is_tls = uri.scheme_str() == Some("wss");
let port = uri.port_u16().unwrap_or(if is_tls { 443 } else { 80 });
// TCP connect with timeout
let connect_timeout = Duration::from_secs(state.config.tunnel_connect_timeout_secs);
let tcp_stream = tokio::time::timeout(connect_timeout, TcpStream::connect((host, port)))
.await
.map_err(|_| {
anyhow::anyhow!(
"tunnel TCP connect timeout ({}s)",
connect_timeout.as_secs()
)
})??;
// Configure TCP parameters via socket2
configure_tcp_socket(&tcp_stream, state);
// WebSocket upgrade (with TLS if wss://)
let connector = if is_tls {
Some(tokio_tungstenite::Connector::Rustls(Arc::clone(
&state.tunnel_tls_config,
)))
} else {
None
};
// Match Python-side _MAX_FRAME_SIZE (64 MiB) to prevent tungstenite's
// default 16 MiB limit from rejecting large AI API payloads (multi-image
// base64 requests can exceed 16 MiB).
let ws_config = WebSocketConfig {
max_frame_size: Some(64 << 20),
max_message_size: Some(64 << 20),
..Default::default()
};
let handshake_timeout = Duration::from_secs(state.config.tunnel_connect_timeout_secs);
let (ws_stream, _response) = tokio::time::timeout(
handshake_timeout,
tokio_tungstenite::client_async_tls_with_config(
request,
tcp_stream,
Some(ws_config),
connector,
),
)
.await
.map_err(|_| {
anyhow::anyhow!(
"tunnel WebSocket handshake timeout ({}s)",
handshake_timeout.as_secs()
)
})??;
info!(
conn = conn_idx,
tcp_keepalive_secs = state.config.tunnel_tcp_keepalive_secs,
tcp_nodelay = state.config.tunnel_tcp_nodelay,
connect_timeout_secs = state.config.tunnel_connect_timeout_secs,
stale_timeout_secs = state.config.tunnel_stale_timeout_secs,
"tunnel connected"
);
// NOTE: reconnect_attempts reset is handled by the caller (mod.rs)
// based on how long the connection stayed alive.
// Split into read/write halves
let (ws_sink, ws_read) = futures_util::StreamExt::split(ws_stream);
// Spawn writer task (with WebSocket ping keepalive)
let ping_interval = Duration::from_secs(state.config.tunnel_ping_interval_secs);
let (frame_tx, mut writer_handle) = writer::spawn_writer(ws_sink, ping_interval);
// Spawn heartbeat task (only for primary connection to avoid
// resetting shared atomic metrics via swap(0))
let hb_handle = if conn_idx == 0 {
heartbeat::spawn(
Arc::clone(&state.config),
Arc::clone(server),
frame_tx.clone(),
shutdown.clone(),
)
} else {
heartbeat::spawn_noop()
};
// Run dispatcher (blocks until disconnect or shutdown).
// Also watch for writer exit — if the write half dies (e.g. the peer
// closed the connection) but the read half stays open, dispatcher would
// block forever on `ws_stream.next()`. Monitoring `writer_handle`
// ensures we detect this and trigger a reconnect promptly.
let state_clone = Arc::clone(state);
let server_clone = Arc::clone(server);
let outcome = tokio::select! {
result = dispatcher::run(state_clone, server_clone, ws_read, frame_tx.clone(), hb_handle) => {
match result {
Ok(()) => TunnelOutcome::Disconnected,
Err(e) => return Err(e),
}
}
writer_result = &mut writer_handle => {
match writer_result {
Ok(()) => warn!("writer task exited normally, triggering reconnect"),
Err(e) => {
if e.is_panic() {
tracing::error!(error = %e, "writer task panicked, triggering reconnect");
} else {
warn!(error = %e, "writer task cancelled, triggering reconnect");
}
}
}
TunnelOutcome::Disconnected
}
_ = shutdown.changed() => {
debug!("shutdown during tunnel dispatch");
TunnelOutcome::Shutdown
}
};
// Drop our sender; the writer will exit once all stream handler clones
// are also dropped (i.e. after they finish their in-flight work).
drop(frame_tx);
// Wait for the writer task to finish with a generous timeout — the
// dispatcher already waits up to 30s for stream handlers, so 35s here
// covers that plus a small margin.
// Skip if the writer already exited (the select branch that fired).
if !writer_handle.is_finished() {
let _ = tokio::time::timeout(Duration::from_secs(35), writer_handle).await;
}
info!("tunnel disconnected");
Ok(outcome)
}
/// Configure TCP keepalive and NODELAY on an established socket.
fn configure_tcp_socket(stream: &TcpStream, state: &Arc<AppState>) {
let sock_ref = socket2::SockRef::from(stream);
if state.config.tunnel_tcp_keepalive_secs > 0 {
let keepalive = socket2::TcpKeepalive::new()
.with_time(Duration::from_secs(state.config.tunnel_tcp_keepalive_secs))
.with_interval(Duration::from_secs(5));
#[cfg(not(target_os = "windows"))]
let keepalive = keepalive.with_retries(3);
if let Err(e) = sock_ref.set_tcp_keepalive(&keepalive) {
warn!(error = %e, "failed to set TCP keepalive on tunnel socket");
}
}
if state.config.tunnel_tcp_nodelay {
if let Err(e) = sock_ref.set_nodelay(true) {
warn!(error = %e, "failed to set TCP_NODELAY on tunnel socket");
}
}
}
/// Build rustls ClientConfig with system root certificates.
pub fn build_tls_config() -> rustls::ClientConfig {
let root_store =
rustls::RootCertStore::from_iter(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
rustls::ClientConfig::builder()
.with_root_certificates(root_store)
.with_no_client_auth()
}
fn build_tunnel_url(server: &ServerContext) -> String {
let base = server.aether_url.trim_end_matches('/');
let ws_base = if base.starts_with("https://") {
base.replacen("https://", "wss://", 1)
} else if base.starts_with("http://") {
base.replacen("http://", "ws://", 1)
} else {
format!("wss://{}", base)
};
format!("{}/api/internal/proxy-tunnel", ws_base)
}
+249
View File
@@ -0,0 +1,249 @@
//! Frame dispatcher: reads incoming WebSocket frames and routes them.
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use bytes::Bytes;
use futures_util::StreamExt;
use tokio::sync::mpsc;
use tokio::task::JoinHandle;
use tokio_tungstenite::tungstenite::Message;
use tracing::{debug, error, info, warn};
use crate::state::{AppState, ServerContext};
use super::heartbeat::HeartbeatHandle;
use super::protocol::{decompress_if_gzip, Frame, MsgType, RequestMeta};
use super::stream_handler;
use super::writer::FrameSender;
/// Run the dispatcher loop, reading from the WebSocket stream.
pub async fn run<S>(
state: Arc<AppState>,
server: Arc<ServerContext>,
mut ws_stream: S,
frame_tx: FrameSender,
heartbeat: HeartbeatHandle,
) -> Result<(), anyhow::Error>
where
S: StreamExt<Item = Result<Message, tokio_tungstenite::tungstenite::Error>>
+ Unpin
+ Send
+ 'static,
{
// Active streams: stream_id -> body sender
let mut streams: HashMap<u32, mpsc::Sender<Frame>> = HashMap::new();
// Track spawned stream handlers so we can wait for them on shutdown
let mut handler_handles: Vec<JoinHandle<()>> = Vec::new();
let max_streams = state.config.tunnel_max_streams.unwrap_or(128) as usize;
let mut frames_since_cleanup: u32 = 0;
let stale_timeout = Duration::from_secs(state.config.tunnel_stale_timeout_secs);
// Track last time we received any data to detect stale connections
let mut last_data_at = tokio::time::Instant::now();
let read_err = loop {
let msg_result = tokio::select! {
msg = ws_stream.next() => {
match msg {
Some(r) => r,
None => break None,
}
}
_ = tokio::time::sleep_until(last_data_at + stale_timeout) => {
warn!(
stale_secs = stale_timeout.as_secs(),
"tunnel connection stale, no data received"
);
break None;
}
};
let msg = match msg_result {
Ok(m) => m,
Err(e) => {
error!(error = %e, "WebSocket read error");
break Some(e);
}
};
// Any successfully received message proves the connection is alive
last_data_at = tokio::time::Instant::now();
let data = match msg {
Message::Binary(data) => Bytes::from(data),
Message::Ping(_) => continue,
Message::Pong(_) => continue,
Message::Close(_) => {
info!("received WebSocket close");
break None;
}
_ => continue,
};
let frame = match Frame::decode(data) {
Ok(f) => f,
Err(e) => {
warn!(error = %e, "failed to decode frame");
continue;
}
};
match frame.msg_type {
MsgType::RequestHeaders => {
// Decompress if the frame is gzip-compressed, then parse metadata
let payload = match decompress_if_gzip(&frame) {
Ok(p) => p,
Err(e) => {
warn!(stream_id = frame.stream_id, error = %e, "frame decompress failed");
continue;
}
};
let meta: RequestMeta = match serde_json::from_slice(&payload) {
Ok(m) => m,
Err(e) => {
warn!(stream_id = frame.stream_id, error = %e, "invalid request metadata");
// Use try_send to avoid blocking the read loop
if frame_tx
.try_send(Frame::new(
frame.stream_id,
MsgType::StreamError,
0,
Bytes::from(format!("invalid request metadata: {e}")),
))
.is_err()
{
warn!(
stream_id = frame.stream_id,
"writer channel full, StreamError dropped"
);
}
continue;
}
};
if streams.len() >= max_streams {
warn!(
stream_id = frame.stream_id,
"max concurrent streams reached"
);
if frame_tx
.try_send(Frame::new(
frame.stream_id,
MsgType::StreamError,
0,
Bytes::from("max concurrent streams reached"),
))
.is_err()
{
warn!(
stream_id = frame.stream_id,
"writer channel full, StreamError dropped"
);
}
continue;
}
// Create body channel and spawn handler
let (body_tx, body_rx) = mpsc::channel::<Frame>(64);
streams.insert(frame.stream_id, body_tx);
let state_clone = Arc::clone(&state);
let server_clone = Arc::clone(&server);
let tx_clone = frame_tx.clone();
let sid = frame.stream_id;
let handle = tokio::spawn(async move {
stream_handler::handle_stream(
state_clone,
server_clone,
sid,
meta,
body_rx,
tx_clone,
)
.await;
});
handler_handles.push(handle);
debug!(stream_id = frame.stream_id, "new stream started");
}
MsgType::RequestBody => {
if let Some(tx) = streams.get(&frame.stream_id) {
let is_end = frame.is_end_stream();
let sid = frame.stream_id;
let _ = tx.send(frame).await;
if is_end {
streams.remove(&sid);
}
}
}
MsgType::StreamEnd | MsgType::StreamError => {
// Client-side cancellation or end
if let Some(tx) = streams.remove(&frame.stream_id) {
let _ = tx.send(frame).await;
}
}
MsgType::Ping => {
// Use try_send to avoid blocking the read loop when writer is congested
if frame_tx
.try_send(Frame::control(MsgType::Pong, frame.payload))
.is_err()
{
warn!("writer channel full, Pong dropped");
}
}
MsgType::HeartbeatAck => {
heartbeat.on_ack(frame.payload).await;
}
MsgType::GoAway => {
info!("received GOAWAY");
break None;
}
_ => {
debug!(msg_type = ?frame.msg_type, "ignoring unexpected frame type");
}
}
// Periodically clean up finished handles to avoid unbounded growth.
// Trigger every 64 frames OR when the count exceeds max_streams.
frames_since_cleanup += 1;
if frames_since_cleanup >= 64 || handler_handles.len() > max_streams {
handler_handles.retain(|h| !h.is_finished());
frames_since_cleanup = 0;
}
};
// Drop body senders so stream handlers waiting on body_rx will unblock
streams.clear();
// Wait for active stream handlers to finish so their frame_tx clones
// are dropped before the writer closes the sink.
drain_handlers(handler_handles).await;
match read_err {
Some(e) => Err(e.into()),
None => Ok(()),
}
}
/// Wait for all active stream handlers to finish (with a timeout).
async fn drain_handlers(handles: Vec<JoinHandle<()>>) {
if handles.is_empty() {
return;
}
let count = handles.len();
debug!(count, "waiting for active stream handlers to finish");
let _ = tokio::time::timeout(Duration::from_secs(30), async {
for h in handles {
let _ = h.await;
}
})
.await;
}
+340
View File
@@ -0,0 +1,340 @@
//! Tunnel heartbeat: sends metrics over the tunnel, processes ACKs.
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::Duration;
use std::time::SystemTime;
use std::time::UNIX_EPOCH;
use bytes::Bytes;
use tokio::sync::watch;
use tracing::{debug, info, warn};
use crate::config::Config;
use crate::registration::client::RemoteConfig;
use crate::runtime;
use crate::state::ServerContext;
use super::protocol::{Frame, MsgType};
use super::writer::FrameSender;
const CURRENT_VERSION: &str = env!("CARGO_PKG_VERSION");
static UPGRADE_IN_PROGRESS: AtomicBool = AtomicBool::new(false);
static NON_ROOT_UPGRADE_WARNED: AtomicBool = AtomicBool::new(false);
enum AckDecision {
Accept {
heartbeat_id: Option<u64>,
upgrade_to: Option<String>,
},
Ignore,
}
/// Handle for the dispatcher to forward HeartbeatAck frames.
#[derive(Clone)]
pub struct HeartbeatHandle {
ack_tx: tokio::sync::mpsc::Sender<Bytes>,
}
impl HeartbeatHandle {
pub async fn on_ack(&self, payload: Bytes) {
let _ = self.ack_tx.send(payload).await;
}
}
/// Create a no-op heartbeat handle that silently discards ACKs.
/// Used for non-primary tunnel connections (conn_idx > 0) to avoid
/// resetting shared atomic metrics via `swap(0)`.
pub fn spawn_noop() -> HeartbeatHandle {
let (ack_tx, _) = tokio::sync::mpsc::channel::<Bytes>(1);
// receiver is immediately dropped; on_ack() calls will silently fail
HeartbeatHandle { ack_tx }
}
#[derive(Debug, Clone, Copy, Default)]
struct HeartbeatSnapshot {
requests: u64,
latency_ns: u64,
failed: u64,
dns_failures: u64,
stream_errors: u64,
}
/// Spawn the heartbeat task. Returns a handle for forwarding ACKs.
pub fn spawn(
_config: Arc<Config>,
server: Arc<ServerContext>,
frame_tx: FrameSender,
mut shutdown: watch::Receiver<bool>,
) -> HeartbeatHandle {
let (ack_tx, mut ack_rx) = tokio::sync::mpsc::channel::<Bytes>(4);
tokio::spawn(async move {
// Read initial interval from dynamic config (may be updated by remote config).
let initial_interval = Duration::from_secs(server.dynamic.load().heartbeat_interval);
let mut current_interval = initial_interval;
// At most one in-flight heartbeat snapshot is tracked at a time.
// Snapshot is only cleared after receiving an ACK, which avoids losing
// interval counters when ACK/frame delivery is temporarily unstable.
let mut pending: Option<(u64, HeartbeatSnapshot)> = None;
let mut next_heartbeat_id: u64 = 1;
let heartbeat_session_id = format!(
"{}-{}",
std::process::id(),
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_nanos()
);
// Skip first immediate tick by sleeping first.
tokio::time::sleep(current_interval).await;
loop {
tokio::select! {
_ = tokio::time::sleep(current_interval) => {
let (heartbeat_id, snapshot) = if let Some((id, snap)) = pending {
(id, snap)
} else {
let snap = collect_snapshot(&server);
let id = next_heartbeat_id;
next_heartbeat_id = next_heartbeat_id.wrapping_add(1);
if next_heartbeat_id == 0 {
next_heartbeat_id = 1;
}
pending = Some((id, snap));
(id, snap)
};
let payload = build_heartbeat_payload(
&server,
&heartbeat_session_id,
heartbeat_id,
snapshot
);
let frame = Frame::control(MsgType::HeartbeatData, payload);
if frame_tx.send(frame).await.is_err() {
if let Some((_, snap)) = pending.take() {
restore_snapshot(&server, snap);
}
break; // Writer closed
}
debug!("sent heartbeat data");
// Re-read interval from dynamic config (remote config may have
// updated it since the last heartbeat).
let new_interval = Duration::from_secs(
server.dynamic.load().heartbeat_interval
);
if new_interval != current_interval {
debug!(
old_secs = current_interval.as_secs(),
new_secs = new_interval.as_secs(),
"heartbeat interval updated from dynamic config"
);
current_interval = new_interval;
}
}
Some(ack_payload) = ack_rx.recv() => {
match handle_ack(&server, &ack_payload) {
AckDecision::Accept {
heartbeat_id: ack_id,
upgrade_to,
} => {
if let Some((pending_id, _)) = pending {
match ack_id {
Some(id) if id == pending_id => {
pending = None;
}
None => {
// Backward-compatible with servers that don't echo
// heartbeat_id in ACK payload yet.
pending = None;
}
_ => {}
}
}
maybe_trigger_upgrade(upgrade_to);
}
AckDecision::Ignore => {}
}
}
_ = shutdown.changed() => {
debug!("heartbeat task shutting down");
if let Some((_, snap)) = pending.take() {
restore_snapshot(&server, snap);
}
break;
}
}
}
});
HeartbeatHandle { ack_tx }
}
fn collect_snapshot(server: &ServerContext) -> HeartbeatSnapshot {
HeartbeatSnapshot {
requests: server.metrics.total_requests.swap(0, Ordering::AcqRel),
latency_ns: server.metrics.total_latency_ns.swap(0, Ordering::AcqRel),
failed: server.metrics.failed_requests.swap(0, Ordering::AcqRel),
dns_failures: server.metrics.dns_failures.swap(0, Ordering::AcqRel),
stream_errors: server.metrics.stream_errors.swap(0, Ordering::AcqRel),
}
}
fn restore_snapshot(server: &ServerContext, snap: HeartbeatSnapshot) {
if snap.requests > 0 {
server
.metrics
.total_requests
.fetch_add(snap.requests, Ordering::Release);
}
if snap.latency_ns > 0 {
server
.metrics
.total_latency_ns
.fetch_add(snap.latency_ns, Ordering::Release);
}
if snap.failed > 0 {
server
.metrics
.failed_requests
.fetch_add(snap.failed, Ordering::Release);
}
if snap.dns_failures > 0 {
server
.metrics
.dns_failures
.fetch_add(snap.dns_failures, Ordering::Release);
}
if snap.stream_errors > 0 {
server
.metrics
.stream_errors
.fetch_add(snap.stream_errors, Ordering::Release);
}
}
fn build_heartbeat_payload(
server: &ServerContext,
heartbeat_session_id: &str,
heartbeat_id: u64,
snapshot: HeartbeatSnapshot,
) -> Bytes {
let node_id = server.node_id.read().unwrap().clone();
let avg_latency_ms = if snapshot.requests > 0 {
Some(snapshot.latency_ns as f64 / snapshot.requests as f64 / 1_000_000.0)
} else {
None
};
let payload = serde_json::json!({
"node_id": node_id,
"heartbeat_session_id": heartbeat_session_id,
"heartbeat_id": heartbeat_id,
"active_connections": server.active_connections.load(Ordering::Acquire),
"total_requests": snapshot.requests,
"avg_latency_ms": avg_latency_ms,
"failed_requests": snapshot.failed,
"dns_failures": snapshot.dns_failures,
"stream_errors": snapshot.stream_errors,
"proxy_metadata": {
"version": CURRENT_VERSION,
},
});
Bytes::from(serde_json::to_vec(&payload).unwrap_or_default())
}
fn handle_ack(server: &ServerContext, payload: &[u8]) -> AckDecision {
if payload.is_empty() {
return AckDecision::Accept {
heartbeat_id: None,
upgrade_to: None,
};
}
#[derive(serde::Deserialize)]
struct AckPayload {
#[serde(default)]
remote_config: Option<RemoteConfig>,
#[serde(default)]
config_version: u64,
#[serde(default)]
heartbeat_id: Option<u64>,
#[serde(default)]
upgrade_to: Option<String>,
}
match serde_json::from_slice::<AckPayload>(payload) {
Ok(ack) => {
if let Some(ref rc) = ack.remote_config {
runtime::apply_remote_config(&server.dynamic, rc, ack.config_version);
}
AckDecision::Accept {
heartbeat_id: ack.heartbeat_id,
upgrade_to: ack.upgrade_to.and_then(normalize_upgrade_target),
}
}
Err(e) => {
warn!(error = %e, "failed to parse heartbeat ACK");
AckDecision::Ignore
}
}
}
fn normalize_upgrade_target(raw: String) -> Option<String> {
let trimmed = raw.trim();
if trimmed.is_empty() {
return None;
}
let normalized = trimmed.strip_prefix("proxy-v").unwrap_or(trimmed);
if normalized == CURRENT_VERSION {
return None;
}
Some(normalized.to_string())
}
fn maybe_trigger_upgrade(version: Option<String>) {
let Some(target_version) = version else {
return;
};
if !crate::setup::service::is_root() {
if NON_ROOT_UPGRADE_WARNED
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
.is_ok()
{
warn!(
target_version = %target_version,
"remote upgrade skipped: root privileges are required"
);
}
return;
}
if UPGRADE_IN_PROGRESS
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
.is_err()
{
debug!(target_version = %target_version, "upgrade already in progress, ignoring");
return;
}
tokio::spawn(async move {
info!(target_version = %target_version, "received remote upgrade instruction");
match crate::setup::upgrade::perform_upgrade(&target_version).await {
Ok(()) => {
info!(target_version = %target_version, "remote upgrade finished");
}
Err(e) => {
warn!(
target_version = %target_version,
error = %e,
"remote upgrade failed"
);
UPGRADE_IN_PROGRESS.store(false, Ordering::Release);
}
}
});
}
+236
View File
@@ -0,0 +1,236 @@
pub mod client;
pub mod dispatcher;
pub mod heartbeat;
pub mod protocol;
pub mod stream_handler;
pub mod writer;
use std::sync::Arc;
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
use tokio::sync::watch;
use tracing::{error, info};
use crate::state::{AppState, ServerContext};
/// If a tunnel stays connected at least this long, treat the next disconnect
/// as a non-failure and reset reconnect backoff.
const STABLE_SESSION_RESET_AFTER: Duration = Duration::from_secs(30);
/// Startup staggering step per secondary connection, used to avoid
/// simultaneous bursts when a pool of tunnels starts together.
const STARTUP_STAGGER_STEP_MS: u64 = 150;
/// Upper bound for startup staggering.
const MAX_STARTUP_STAGGER_MS: u64 = 1_500;
/// Keep a tiny floor for repeated reconnects; first retry is still immediate.
const MIN_RECONNECT_DELAY_MS: u64 = 50;
/// Even under sustained failures, keep probing frequently so recovery is fast
/// once cross-border network quality improves.
const RECONNECT_PROBE_MAX_DELAY_MS: u64 = 3_000;
/// Run the tunnel mode main loop (connect, dispatch, reconnect).
///
/// `conn_idx` identifies which connection in the pool this is (0-based).
/// Only connection 0 sends heartbeats to avoid resetting shared metrics.
pub async fn run(
state: &Arc<AppState>,
server: &Arc<ServerContext>,
conn_idx: usize,
mut shutdown: watch::Receiver<bool>,
) {
info!(server = %server.server_label, conn = conn_idx, "starting tunnel");
let reconnect_salt = compute_connection_salt(server, conn_idx);
let startup_delay = compute_startup_stagger(conn_idx, reconnect_salt);
if !startup_delay.is_zero() {
info!(
server = %server.server_label,
conn = conn_idx,
delay_ms = startup_delay.as_millis(),
"startup stagger before first connect"
);
tokio::select! {
_ = tokio::time::sleep(startup_delay) => {}
_ = shutdown.changed() => {
info!(server = %server.server_label, conn = conn_idx, "shutdown requested during startup stagger");
return;
}
}
}
let mut consecutive_failures: u32 = 0;
loop {
let started_at = Instant::now();
match client::connect_and_run(state, server, conn_idx, &mut shutdown).await {
Ok(client::TunnelOutcome::Shutdown) => {
info!(server = %server.server_label, conn = conn_idx, "tunnel shut down gracefully");
return;
}
Ok(client::TunnelOutcome::Disconnected) => {
info!(server = %server.server_label, conn = conn_idx, "tunnel disconnected, reconnecting");
}
Err(e) => {
error!(server = %server.server_label, conn = conn_idx, error = %e, "tunnel connection error, reconnecting");
}
}
if *shutdown.borrow() {
info!(server = %server.server_label, conn = conn_idx, "shutdown requested, not reconnecting");
return;
}
// Reset backoff after a stable session to keep recovery snappy when
// failures are only occasional.
let connected_for = started_at.elapsed();
if connected_for >= STABLE_SESSION_RESET_AFTER {
consecutive_failures = 0;
} else {
consecutive_failures = consecutive_failures.saturating_add(1);
}
let reconnect_delay = compute_reconnect_delay(
state.config.tunnel_reconnect_base_ms,
state.config.tunnel_reconnect_max_ms,
consecutive_failures,
reconnect_salt,
);
info!(
server = %server.server_label,
conn = conn_idx,
failures = consecutive_failures,
delay_ms = reconnect_delay.as_millis(),
"waiting before reconnect"
);
tokio::select! {
_ = tokio::time::sleep(reconnect_delay) => {}
_ = shutdown.changed() => {
info!(server = %server.server_label, conn = conn_idx, "shutdown requested during reconnect wait");
return;
}
}
}
}
fn compute_connection_salt(server: &ServerContext, conn_idx: usize) -> u64 {
// FNV-1a style hash over server label + connection index.
let mut h: u64 = 0xcbf29ce484222325;
for &b in server.server_label.as_bytes() {
h ^= b as u64;
h = h.wrapping_mul(0x100000001b3);
}
h ^= conn_idx as u64;
mix_u64(h)
}
fn compute_startup_stagger(conn_idx: usize, salt: u64) -> Duration {
if conn_idx == 0 {
return Duration::ZERO;
}
let base = (conn_idx as u64).saturating_mul(STARTUP_STAGGER_STEP_MS);
let jitter = mix_u64(salt) % 301; // 0..=300ms
Duration::from_millis((base + jitter).min(MAX_STARTUP_STAGGER_MS))
}
fn compute_reconnect_delay(
base_ms: u64,
max_ms: u64,
consecutive_failures: u32,
salt: u64,
) -> Duration {
// First retry should be immediate to maximize recovery speed on transient
// blips (the user's primary expectation in poor networks).
if consecutive_failures <= 1 {
return Duration::ZERO;
}
// Keep a sane minimum for repeated failures.
let base_ms = base_ms.max(MIN_RECONNECT_DELAY_MS);
let max_ms = max_ms.max(base_ms);
let cap_ms = compute_reconnect_cap_ms(base_ms, max_ms, consecutive_failures)
.min(RECONNECT_PROBE_MAX_DELAY_MS.max(base_ms));
// Equal-jitter: randomize in [cap/2, cap], preventing synchronized reconnect
// storms while keeping reconnect latency bounded.
if cap_ms <= 1 {
return Duration::from_millis(cap_ms);
}
let half = cap_ms / 2;
let span = cap_ms - half;
let now_nanos = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.subsec_nanos() as u64)
.unwrap_or(0);
let mixed = mix_u64(now_nanos ^ salt);
let jitter = if span == 0 { 0 } else { mixed % (span + 1) };
Duration::from_millis(half + jitter)
}
fn compute_reconnect_cap_ms(base_ms: u64, max_ms: u64, consecutive_failures: u32) -> u64 {
if consecutive_failures <= 1 {
return base_ms.min(max_ms);
}
let shift = (consecutive_failures - 1).min(31);
let factor = 1u64 << shift;
base_ms.saturating_mul(factor).min(max_ms)
}
fn mix_u64(mut x: u64) -> u64 {
// SplitMix64 finalizer - cheap bit mixing for pseudo-random jitter.
x ^= x >> 30;
x = x.wrapping_mul(0xbf58476d1ce4e5b9);
x ^= x >> 27;
x = x.wrapping_mul(0x94d049bb133111eb);
x ^ (x >> 31)
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use super::{
compute_reconnect_cap_ms, compute_reconnect_delay, compute_startup_stagger,
MAX_STARTUP_STAGGER_MS, RECONNECT_PROBE_MAX_DELAY_MS, STARTUP_STAGGER_STEP_MS,
};
#[test]
fn reconnect_cap_grows_exponentially_and_caps() {
let base = 500;
let max = 30_000;
assert_eq!(compute_reconnect_cap_ms(base, max, 0), 500);
assert_eq!(compute_reconnect_cap_ms(base, max, 1), 500);
assert_eq!(compute_reconnect_cap_ms(base, max, 2), 1_000);
assert_eq!(compute_reconnect_cap_ms(base, max, 3), 2_000);
assert_eq!(compute_reconnect_cap_ms(base, max, 4), 4_000);
assert_eq!(compute_reconnect_cap_ms(base, max, 5), 8_000);
assert_eq!(compute_reconnect_cap_ms(base, max, 6), 16_000);
assert_eq!(compute_reconnect_cap_ms(base, max, 7), 30_000);
assert_eq!(compute_reconnect_cap_ms(base, max, 20), 30_000);
}
#[test]
fn startup_stagger_is_zero_for_primary_and_bounded_for_secondary() {
assert_eq!(compute_startup_stagger(0, 42), Duration::ZERO);
let d1 = compute_startup_stagger(1, 42);
let d2 = compute_startup_stagger(2, 42);
assert!(d1 >= Duration::from_millis(STARTUP_STAGGER_STEP_MS));
assert!(d1 <= Duration::from_millis(MAX_STARTUP_STAGGER_MS));
assert!(d2 >= Duration::from_millis(STARTUP_STAGGER_STEP_MS * 2));
assert!(d2 <= Duration::from_millis(MAX_STARTUP_STAGGER_MS));
}
#[test]
fn reconnect_delay_is_immediate_on_first_failure() {
assert_eq!(compute_reconnect_delay(700, 45_000, 1, 123), Duration::ZERO);
}
#[test]
fn reconnect_delay_stays_within_probe_ceiling_after_many_failures() {
let d = compute_reconnect_delay(500, 45_000, 100, 12345);
assert!(d <= Duration::from_millis(RECONNECT_PROBE_MAX_DELAY_MS));
}
}
+260
View File
@@ -0,0 +1,260 @@
//! Binary frame protocol for WebSocket tunnel multiplexing.
//!
//! Frame layout (10-byte header + variable payload):
//! ```text
//! | stream_id (4B) | msg_type (1B) | flags (1B) | payload_len (4B) | payload (NB) |
//! ```
use bytes::{Buf, BufMut, Bytes, BytesMut};
pub const HEADER_SIZE: usize = 10;
/// Frame flags.
pub mod flags {
pub const END_STREAM: u8 = 0x01;
pub const GZIP_COMPRESSED: u8 = 0x02;
}
/// Message types for the tunnel protocol.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub enum MsgType {
RequestHeaders = 0x01,
RequestBody = 0x02,
ResponseHeaders = 0x03,
ResponseBody = 0x04,
StreamEnd = 0x05,
StreamError = 0x06,
Ping = 0x10,
Pong = 0x11,
GoAway = 0x12,
HeartbeatData = 0x13,
HeartbeatAck = 0x14,
}
impl MsgType {
pub fn from_u8(v: u8) -> Option<Self> {
match v {
0x01 => Some(Self::RequestHeaders),
0x02 => Some(Self::RequestBody),
0x03 => Some(Self::ResponseHeaders),
0x04 => Some(Self::ResponseBody),
0x05 => Some(Self::StreamEnd),
0x06 => Some(Self::StreamError),
0x10 => Some(Self::Ping),
0x11 => Some(Self::Pong),
0x12 => Some(Self::GoAway),
0x13 => Some(Self::HeartbeatData),
0x14 => Some(Self::HeartbeatAck),
_ => None,
}
}
}
/// A single multiplexed frame.
#[derive(Debug, Clone)]
pub struct Frame {
pub stream_id: u32,
pub msg_type: MsgType,
pub flags: u8,
pub payload: Bytes,
}
impl Frame {
pub fn new(stream_id: u32, msg_type: MsgType, flags: u8, payload: impl Into<Bytes>) -> Self {
Self {
stream_id,
msg_type,
flags,
payload: payload.into(),
}
}
/// Control frame (stream_id = 0).
pub fn control(msg_type: MsgType, payload: impl Into<Bytes>) -> Self {
Self::new(0, msg_type, 0, payload)
}
pub fn is_end_stream(&self) -> bool {
self.flags & flags::END_STREAM != 0
}
pub fn is_gzip(&self) -> bool {
self.flags & flags::GZIP_COMPRESSED != 0
}
/// Encode into a binary buffer.
pub fn encode(&self) -> Bytes {
let mut buf = BytesMut::with_capacity(HEADER_SIZE + self.payload.len());
buf.put_u32(self.stream_id);
buf.put_u8(self.msg_type as u8);
buf.put_u8(self.flags);
buf.put_u32(self.payload.len() as u32);
buf.put(self.payload.clone());
buf.freeze()
}
/// Decode from a binary buffer.
pub fn decode(mut data: Bytes) -> Result<Self, ProtocolError> {
if data.len() < HEADER_SIZE {
return Err(ProtocolError::TooShort {
expected: HEADER_SIZE,
actual: data.len(),
});
}
let stream_id = data.get_u32();
let msg_type_raw = data.get_u8();
let frame_flags = data.get_u8();
let payload_len = data.get_u32() as usize;
if data.remaining() < payload_len {
return Err(ProtocolError::Incomplete {
expected: HEADER_SIZE + payload_len,
actual: HEADER_SIZE + data.remaining(),
});
}
let msg_type =
MsgType::from_u8(msg_type_raw).ok_or(ProtocolError::UnknownMsgType(msg_type_raw))?;
let payload = data.split_to(payload_len);
Ok(Self {
stream_id,
msg_type,
flags: frame_flags,
payload,
})
}
}
/// Protocol errors.
#[derive(Debug, thiserror::Error)]
pub enum ProtocolError {
#[error("frame too short: expected {expected} bytes, got {actual}")]
TooShort { expected: usize, actual: usize },
#[error("frame incomplete: expected {expected} bytes, got {actual}")]
Incomplete { expected: usize, actual: usize },
#[error("unknown message type: 0x{0:02x}")]
UnknownMsgType(u8),
}
/// JSON payload for REQUEST_HEADERS frames.
#[derive(Debug, serde::Deserialize)]
pub struct RequestMeta {
pub method: String,
pub url: String,
pub headers: std::collections::HashMap<String, String>,
#[serde(default = "default_timeout", deserialize_with = "deserialize_timeout")]
pub timeout: u64,
}
fn default_timeout() -> u64 {
60
}
fn deserialize_timeout<'de, D>(deserializer: D) -> Result<u64, D::Error>
where
D: serde::Deserializer<'de>,
{
#[derive(serde::Deserialize)]
#[serde(untagged)]
enum TimeoutValue {
Int(u64),
Float(f64),
}
match <TimeoutValue as serde::Deserialize>::deserialize(deserializer)? {
TimeoutValue::Int(v) => Ok(v),
TimeoutValue::Float(v) => {
if !v.is_finite() || v < 0.0 {
return Err(serde::de::Error::custom(
"timeout must be a non-negative finite number",
));
}
if v.fract() != 0.0 {
return Err(serde::de::Error::custom("timeout must be integer seconds"));
}
if v > (u64::MAX as f64) {
return Err(serde::de::Error::custom("timeout is too large"));
}
Ok(v as u64)
}
}
}
/// JSON payload for RESPONSE_HEADERS frames.
#[derive(Debug, serde::Serialize)]
pub struct ResponseMeta {
pub status: u16,
/// Header list preserving duplicates (e.g. multiple Set-Cookie).
pub headers: Vec<(String, String)>,
}
// ---------------------------------------------------------------------------
// Tunnel frame compression helpers
// ---------------------------------------------------------------------------
/// Minimum payload size to attempt gzip compression (bytes).
const COMPRESS_MIN_SIZE: usize = 512;
/// If the frame has the GZIP_COMPRESSED flag, decompress the payload; otherwise
/// return a clone of the raw payload bytes.
pub fn decompress_if_gzip(frame: &Frame) -> Result<Bytes, std::io::Error> {
if frame.is_gzip() {
decompress_gzip(&frame.payload)
} else {
Ok(frame.payload.clone())
}
}
/// Gzip-compress `data` if it is large enough and compression actually shrinks
/// the payload. Returns `(payload, extra_flags)` where `extra_flags` contains
/// `GZIP_COMPRESSED` when compression was applied.
pub fn compress_payload(data: Bytes) -> (Bytes, u8) {
if data.len() >= COMPRESS_MIN_SIZE {
if let Ok(compressed) = compress_gzip(&data) {
if compressed.len() < data.len() {
return (compressed, flags::GZIP_COMPRESSED);
}
}
}
(data, 0)
}
fn decompress_gzip(data: &[u8]) -> Result<Bytes, std::io::Error> {
use flate2::read::GzDecoder;
use std::io::Read;
let mut decoder = GzDecoder::new(data);
let mut buf = Vec::new();
decoder.read_to_end(&mut buf)?;
Ok(Bytes::from(buf))
}
fn compress_gzip(data: &[u8]) -> Result<Bytes, std::io::Error> {
use flate2::write::GzEncoder;
use flate2::Compression;
use std::io::Write;
let mut encoder = GzEncoder::new(Vec::new(), Compression::fast());
encoder.write_all(data)?;
let compressed = encoder.finish()?;
Ok(Bytes::from(compressed))
}
#[cfg(test)]
mod tests {
use super::RequestMeta;
#[test]
fn request_meta_accepts_integer_timeout() {
let raw = br#"{"method":"GET","url":"https://example.com","headers":{},"timeout":15}"#;
let meta: RequestMeta = serde_json::from_slice(raw).expect("parse request meta");
assert_eq!(meta.timeout, 15);
}
#[test]
fn request_meta_accepts_integer_like_float_timeout() {
let raw = br#"{"method":"GET","url":"https://example.com","headers":{},"timeout":15.0}"#;
let meta: RequestMeta = serde_json::from_slice(raw).expect("parse request meta");
assert_eq!(meta.timeout, 15);
}
}
+506
View File
@@ -0,0 +1,506 @@
//! Per-stream request handler.
//!
//! Receives request frames, executes the upstream HTTP request,
//! and sends response frames back through the writer channel.
use std::io;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
use std::sync::Arc;
use std::time::{Duration, Instant};
use bytes::Bytes;
use futures_util::stream;
use futures_util::StreamExt;
use http_body_util::BodyExt;
use hyper::body::Frame as BodyFrame;
use tokio::sync::mpsc;
use tracing::{debug, warn};
use crate::state::{AppState, ServerContext};
use crate::target_filter;
use crate::upstream_client;
use super::protocol::{
compress_payload, decompress_if_gzip, flags, Frame as TunnelFrame, MsgType, RequestMeta,
ResponseMeta,
};
use super::writer::FrameSender;
/// Maximum response body chunk size per frame (32 KB).
const MAX_CHUNK_SIZE: usize = 32 * 1024;
/// Timeout for sending a single frame to the writer channel.
/// If the writer is congested (TCP backpressure), we abandon the stream
/// rather than blocking indefinitely and exhausting the stream pool.
const FRAME_SEND_TIMEOUT: Duration = Duration::from_secs(30);
/// Minimum allowed upstream request timeout (seconds).
const MIN_TIMEOUT_SECS: u64 = 5;
/// Maximum allowed upstream request timeout (seconds).
const MAX_TIMEOUT_SECS: u64 = 300;
/// Headers that must not be forwarded to upstream (hop-by-hop or security-sensitive).
///
/// `host` and `content-length` are managed by the HTTP client (reqwest/hyper):
/// - `host` → translated to `:authority` pseudo-header in HTTP/2; forwarding
/// the original `host` alongside `:authority` triggers PROTOCOL_ERROR on
/// strict H2 implementations (e.g. Google APIs).
/// - `content-length` → recalculated by hyper from the actual body; a stale
/// value from the tunnel (body may have been re-compressed) causes H2
/// PROTOCOL_ERROR when it mismatches the real frame length.
const BLOCKED_HEADERS: &[&str] = &[
"connection",
"content-length",
"host",
"keep-alive",
"proxy-authenticate",
"proxy-authorization",
"proxy-connection",
"te",
"trailer",
"transfer-encoding",
"upgrade",
];
/// Handle a single stream: receive body, execute upstream, send response.
pub async fn handle_stream(
state: Arc<AppState>,
server: Arc<ServerContext>,
stream_id: u32,
meta: RequestMeta,
body_rx: mpsc::Receiver<TunnelFrame>,
frame_tx: FrameSender,
) {
server.active_connections.fetch_add(1, Ordering::Release);
let connect_elapsed =
handle_stream_inner(&state, &server, stream_id, meta, body_rx, &frame_tx).await;
server.active_connections.fetch_sub(1, Ordering::Release);
if let Some(d) = connect_elapsed {
server.metrics.record_request(d);
}
}
/// Send a frame to the writer with a timeout. Returns false if send failed.
async fn send_frame(tx: &FrameSender, frame: TunnelFrame) -> bool {
match tokio::time::timeout(FRAME_SEND_TIMEOUT, tx.send(frame)).await {
Ok(Ok(())) => true,
Ok(Err(_)) => {
// Channel closed (writer exited)
false
}
Err(_) => {
// Timeout — writer is congested
warn!("frame send timeout (writer congested), abandoning stream");
false
}
}
}
/// Returns the connection-establishment duration (DNS + TCP/TLS + TTFB) if the
/// upstream request succeeded, or `None` if the request never reached the
/// response-headers stage.
async fn handle_stream_inner(
state: &AppState,
server: &ServerContext,
stream_id: u32,
meta: RequestMeta,
body_rx: mpsc::Receiver<TunnelFrame>,
frame_tx: &FrameSender,
) -> Option<Duration> {
// Validate target
let target_url = match url::Url::parse(&meta.url) {
Ok(u) => u,
Err(e) => {
send_error(frame_tx, stream_id, &format!("invalid URL: {e}")).await;
return None;
}
};
// Only allow http/https schemes (block file://, data://, etc.)
match target_url.scheme() {
"http" | "https" => {}
other => {
send_error(
frame_tx,
stream_id,
&format!("unsupported URL scheme: {other}"),
)
.await;
return None;
}
}
let host = match target_url.host_str() {
Some(h) => h.to_string(),
None => {
send_error(frame_tx, stream_id, "missing host in URL").await;
return None;
}
};
let port = target_url.port_or_known_default().unwrap_or(443);
// DNS + target validation (populates dns_cache for SafeDnsResolver)
let connect_start = Instant::now();
{
let allowed_ports = Arc::clone(&server.dynamic.load().allowed_ports);
if let Err(e) =
target_filter::validate_target(&host, port, &allowed_ports, &state.dns_cache).await
{
server.metrics.dns_failures.fetch_add(1, Ordering::Release);
send_error(frame_tx, stream_id, &format!("target blocked: {e}")).await;
return None;
}
}
let dns_ms = connect_start.elapsed().as_millis() as u64;
// Execute upstream request
let client = &state.upstream_client;
let timeout = Duration::from_secs(meta.timeout.clamp(MIN_TIMEOUT_SECS, MAX_TIMEOUT_SECS));
let request_body_size = Arc::new(AtomicUsize::new(0));
let request_body = build_streaming_request_body(body_rx, Arc::clone(&request_body_size));
let method: hyper::Method = meta.method.parse().unwrap_or(hyper::Method::GET);
let mut request = match hyper::Request::builder()
.method(method)
.uri(meta.url.as_str())
.body(request_body)
{
Ok(request) => request,
Err(e) => {
send_error(
frame_tx,
stream_id,
&format!("invalid upstream request: {e}"),
)
.await;
return None;
}
};
let headers = request.headers_mut();
for (k, v) in &meta.headers {
let k_lower = k.to_ascii_lowercase();
if BLOCKED_HEADERS.contains(&k_lower.as_str()) {
continue;
}
if let (Ok(name), Ok(value)) = (
hyper::header::HeaderName::from_bytes(k.as_bytes()),
hyper::header::HeaderValue::from_str(v),
) {
headers.insert(name, value);
}
}
let mut captured_connection = upstream_client::capture_connection(&mut request);
let connection_start = Instant::now();
let connection_capture = tokio::spawn(async move {
let connected = captured_connection.wait_for_connection_metadata().await;
connected
.as_ref()
.map(|_| connection_start.elapsed().as_millis() as u64)
});
let upstream_start = Instant::now();
let response = match tokio::time::timeout(timeout, client.request(request)).await {
Ok(Ok(response)) => response,
Ok(Err(e)) => {
connection_capture.abort();
server
.metrics
.failed_requests
.fetch_add(1, Ordering::Release);
let msg = if e.is_connect() {
format!("upstream connect error: {e}")
} else {
format!("upstream error: {e}")
};
send_error(frame_tx, stream_id, &msg).await;
return None;
}
Err(_) => {
connection_capture.abort();
server
.metrics
.failed_requests
.fetch_add(1, Ordering::Release);
send_error(frame_tx, stream_id, "upstream timeout").await;
return None;
}
};
// Capture connection-establishment duration (DNS + TCP/TLS + TTFB)
// before proceeding to stream the response body.
let connect_elapsed = connect_start.elapsed();
// Send RESPONSE_HEADERS
let status = response.status().as_u16();
let ttfb_ms = upstream_start.elapsed().as_millis() as u64;
// Short timeout: on connection reuse hyper may never fire the connect
// callback, so avoid blocking indefinitely.
let connection_acquire_ms =
match tokio::time::timeout(Duration::from_millis(100), connection_capture).await {
Ok(Ok(ms)) => ms,
Ok(Err(_)) => None, // JoinError (task panicked / cancelled)
Err(_) => None, // timeout -- task is detached but lightweight
};
let request_timing =
upstream_client::resolve_request_timing(&response, connection_acquire_ms, ttfb_ms);
let mut resp_headers: Vec<(String, String)> = Vec::with_capacity(response.headers().len() + 1);
for (k, v) in response.headers() {
if let Ok(vs) = v.to_str() {
resp_headers.push((k.as_str().to_string(), vs.to_string()));
}
}
let timing = serde_json::json!({
"dns_ms": dns_ms,
"connection_acquire_ms": request_timing.connection_acquire_ms,
"connection_reused": request_timing.connection_reused,
"connect_ms": request_timing.connect_ms,
"tls_ms": request_timing.tls_ms,
"ttfb_ms": ttfb_ms,
"upstream_ms": ttfb_ms,
"response_wait_ms": request_timing.response_wait_ms,
"upstream_processing_ms": request_timing.response_wait_ms,
"timing_source": "instrumented_connector",
"total_ms": connect_elapsed.as_millis() as u64,
"body_size": request_body_size.load(Ordering::Relaxed),
"mode": "tunnel",
});
resp_headers.push(("x-proxy-timing".to_string(), timing.to_string()));
let resp_meta = ResponseMeta {
status,
headers: resp_headers,
};
let meta_json: Bytes = serde_json::to_vec(&resp_meta).unwrap_or_default().into();
let (meta_payload, meta_flags) = compress_payload(meta_json);
if !send_frame(
frame_tx,
TunnelFrame::new(
stream_id,
MsgType::ResponseHeaders,
meta_flags,
meta_payload,
),
)
.await
{
return Some(connect_elapsed);
}
// Stream response body — relay upstream bytes through the tunnel.
// Apply tunnel-level frame compression for chunks that benefit from it
// (e.g. uncompressed SSE text). Already-compressed data (gzip/br from
// upstream Content-Encoding) won't shrink further and will be sent as-is
// thanks to the size check in compress_payload().
let mut stream = response.into_body().into_data_stream();
while let Some(chunk_result) = stream.next().await {
match chunk_result {
Ok(chunk) => {
if chunk.len() <= MAX_CHUNK_SIZE {
let (payload, extra_flags) = compress_payload(chunk);
if !send_frame(
frame_tx,
TunnelFrame::new(stream_id, MsgType::ResponseBody, extra_flags, payload),
)
.await
{
return Some(connect_elapsed);
}
} else {
// Split oversized chunks, compress each slice
let mut offset = 0;
while offset < chunk.len() {
let end = (offset + MAX_CHUNK_SIZE).min(chunk.len());
let slice = chunk.slice(offset..end);
let (payload, extra_flags) = compress_payload(slice);
if !send_frame(
frame_tx,
TunnelFrame::new(
stream_id,
MsgType::ResponseBody,
extra_flags,
payload,
),
)
.await
{
return Some(connect_elapsed);
}
offset = end;
}
}
}
Err(e) => {
server.metrics.stream_errors.fetch_add(1, Ordering::Release);
warn!(stream_id, error = %e, "upstream body read error");
send_error(frame_tx, stream_id, &format!("body read error: {e}")).await;
return Some(connect_elapsed);
}
}
}
// Send STREAM_END
let _ = send_frame(
frame_tx,
TunnelFrame::new(
stream_id,
MsgType::StreamEnd,
flags::END_STREAM,
Bytes::new(),
),
)
.await;
debug!(stream_id, status, "stream completed");
Some(connect_elapsed)
}
async fn send_error(tx: &FrameSender, stream_id: u32, msg: &str) {
// Error frames use best-effort delivery — don't block if writer is congested
let _ = send_frame(
tx,
TunnelFrame::new(
stream_id,
MsgType::StreamError,
0,
Bytes::from(msg.to_string()),
),
)
.await;
}
fn build_streaming_request_body(
body_rx: mpsc::Receiver<TunnelFrame>,
body_size: Arc<AtomicUsize>,
) -> upstream_client::UpstreamRequestBody {
let body_stream = stream::unfold(
(body_rx, body_size, false),
|(mut body_rx, body_size, finished)| async move {
if finished {
return None;
}
loop {
let frame = match body_rx.recv().await {
Some(frame) => frame,
None => return None,
};
match frame.msg_type {
MsgType::RequestBody => {
let end_stream = frame.is_end_stream();
let payload = match decompress_if_gzip(&frame) {
Ok(payload) => payload,
Err(error) => {
let err =
io::Error::other(format!("gzip decompress failed: {error}"));
return Some((Err(err), (body_rx, body_size, true)));
}
};
if payload.is_empty() {
if end_stream {
return None;
}
continue;
}
body_size.fetch_add(payload.len(), Ordering::Relaxed);
return Some((
Ok(BodyFrame::data(payload)),
(body_rx, body_size, end_stream),
));
}
MsgType::StreamError => {
let message = String::from_utf8(frame.payload.to_vec())
.unwrap_or_else(|_| "client cancelled request body".to_string());
return Some((Err(io::Error::other(message)), (body_rx, body_size, true)));
}
MsgType::StreamEnd => return None,
_ => continue,
}
}
},
);
upstream_client::stream_request_body(body_stream)
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn streaming_request_body_yields_chunks_and_tracks_size() {
let (tx, rx) = mpsc::channel(4);
let body_size = Arc::new(AtomicUsize::new(0));
let mut body = build_streaming_request_body(rx, Arc::clone(&body_size));
tx.send(TunnelFrame::new(
1,
MsgType::RequestBody,
0,
Bytes::from_static(b"abc"),
))
.await
.expect("send first chunk");
tx.send(TunnelFrame::new(
1,
MsgType::RequestBody,
flags::END_STREAM,
Bytes::from_static(b"def"),
))
.await
.expect("send final chunk");
drop(tx);
let first = body
.frame()
.await
.expect("first frame")
.expect("first frame ok")
.into_data()
.expect("first data frame");
let second = body
.frame()
.await
.expect("second frame")
.expect("second frame ok")
.into_data()
.expect("second data frame");
assert_eq!(first, Bytes::from_static(b"abc"));
assert_eq!(second, Bytes::from_static(b"def"));
assert!(body.frame().await.is_none());
assert_eq!(body_size.load(Ordering::Relaxed), 6);
}
#[tokio::test]
async fn streaming_request_body_surfaces_client_cancel_as_error() {
let (tx, rx) = mpsc::channel(4);
let body_size = Arc::new(AtomicUsize::new(0));
let mut body = build_streaming_request_body(rx, Arc::clone(&body_size));
tx.send(TunnelFrame::new(
1,
MsgType::StreamError,
0,
Bytes::from_static(b"client cancelled"),
))
.await
.expect("send cancel frame");
drop(tx);
let err = body
.frame()
.await
.expect("error frame present")
.expect_err("body should surface cancellation error");
assert!(err.to_string().contains("client cancelled"));
assert!(body.frame().await.is_none());
assert_eq!(body_size.load(Ordering::Relaxed), 0);
}
}
+63
View File
@@ -0,0 +1,63 @@
//! Dedicated WebSocket writer task.
//!
//! All frame writes go through an mpsc channel to a single writer task,
//! avoiding contention on the WebSocket sink. The writer also sends
//! periodic WebSocket Ping frames to keep the connection alive through
//! intermediary proxies (Nginx, Cloudflare, etc.).
use std::time::Duration;
use futures_util::SinkExt;
use tokio::sync::mpsc;
use tokio::task::JoinHandle;
use tokio_tungstenite::tungstenite::Message;
use tracing::{debug, error, trace};
use super::protocol::Frame;
/// Sender half — cloned by stream handlers and heartbeat.
pub type FrameSender = mpsc::Sender<Frame>;
/// Spawn the writer task. Returns the sender and a JoinHandle for cleanup.
///
/// `ping_interval` controls WebSocket-level Ping frequency (typically 15s).
/// This keeps the connection alive through intermediary proxies/load-balancers.
pub fn spawn_writer<S>(mut sink: S, ping_interval: Duration) -> (FrameSender, JoinHandle<()>)
where
S: SinkExt<Message, Error = tokio_tungstenite::tungstenite::Error> + Unpin + Send + 'static,
{
let (tx, mut rx) = mpsc::channel::<Frame>(256);
let handle = tokio::spawn(async move {
let mut ping_ticker = tokio::time::interval(ping_interval);
ping_ticker.tick().await; // skip first immediate tick
loop {
tokio::select! {
frame = rx.recv() => {
match frame {
Some(frame) => {
let data = frame.encode();
if let Err(e) = sink.send(Message::Binary(data.into())).await {
error!(error = %e, "failed to write frame to WebSocket");
break;
}
}
None => break, // all senders dropped
}
}
_ = ping_ticker.tick() => {
if let Err(e) = sink.send(Message::Ping(vec![])).await {
error!(error = %e, "failed to send WebSocket ping");
break;
}
trace!("sent WebSocket ping");
}
}
}
debug!("writer task exiting");
let _ = sink.close().await;
});
(tx, handle)
}
+444
View File
@@ -0,0 +1,444 @@
use std::future::Future;
use std::io;
use std::net::IpAddr;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use std::time::Duration;
use bytes::Bytes;
use futures_util::Stream;
use http_body_util::combinators::UnsyncBoxBody;
use http_body_util::{BodyExt, StreamBody};
use hyper::body::Frame;
use hyper::rt;
use hyper::Response;
use hyper::Uri;
pub use hyper_util::client::legacy::connect::capture_connection;
use hyper_util::client::legacy::connect::dns::Name;
use hyper_util::client::legacy::connect::{Connected, Connection, HttpConnector};
use hyper_util::client::legacy::Client;
use hyper_util::rt::{TokioExecutor, TokioIo, TokioTimer};
use rustls::pki_types::ServerName;
use rustls::ClientConfig;
use tokio::net::TcpStream;
use tokio_rustls::TlsConnector;
use tower_service::Service;
use crate::config::Config;
use crate::target_filter::{self, DnsCache};
type BoxError = Box<dyn std::error::Error + Send + Sync>;
type PlainStream = TokioIo<TcpStream>;
type TlsStream = TokioIo<tokio_rustls::client::TlsStream<TcpStream>>;
pub type UpstreamRequestBody = UnsyncBoxBody<Bytes, io::Error>;
pub type UpstreamClient = Client<InstrumentedConnector, UpstreamRequestBody>;
pub fn stream_request_body<S>(stream: S) -> UpstreamRequestBody
where
S: Stream<Item = Result<Frame<Bytes>, io::Error>> + Send + 'static,
{
StreamBody::new(stream).boxed_unsync()
}
#[derive(Clone, Copy, Debug, Default)]
pub struct ConnectTiming {
pub connect_ms: u64,
pub tls_ms: u64,
}
#[derive(Clone, Copy, Debug, Default)]
pub struct RequestTiming {
pub connection_acquire_ms: u64,
pub connect_ms: u64,
pub tls_ms: u64,
pub response_wait_ms: u64,
pub connection_reused: bool,
}
#[derive(Clone)]
pub struct ValidatedResolver {
dns_cache: Arc<DnsCache>,
}
impl ValidatedResolver {
pub fn new(dns_cache: Arc<DnsCache>) -> Self {
Self { dns_cache }
}
}
pub struct ValidatedAddrs {
inner: std::vec::IntoIter<std::net::SocketAddr>,
}
impl Iterator for ValidatedAddrs {
type Item = std::net::SocketAddr;
fn next(&mut self) -> Option<Self::Item> {
self.inner.next()
}
}
impl Service<Name> for ValidatedResolver {
type Response = ValidatedAddrs;
type Error = io::Error;
type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn call(&mut self, name: Name) -> Self::Future {
let dns_cache = Arc::clone(&self.dns_cache);
let host = name.as_str().to_string();
Box::pin(async move {
if let Some(addrs) = dns_cache.get_by_host(&host).await {
return Ok(ValidatedAddrs {
inner: (*addrs).clone().into_iter(),
});
}
let resolved = target_filter::resolve_public_addrs(&host, 0, dns_cache.as_ref())
.await
.map_err(|err| io::Error::other(err.to_string()))?;
Ok(ValidatedAddrs {
inner: resolved.into_iter(),
})
})
}
}
#[derive(Clone)]
pub struct InstrumentedConnector {
http: HttpConnector<ValidatedResolver>,
tls_config: Arc<ClientConfig>,
}
impl Service<Uri> for InstrumentedConnector {
type Response = TimedConn;
type Error = BoxError;
type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.http.poll_ready(cx).map_err(Into::into)
}
fn call(&mut self, dst: Uri) -> Self::Future {
let scheme = dst.scheme_str().map(|value| value.to_ascii_lowercase());
let tls_config = Arc::clone(&self.tls_config);
let connecting = self.http.call(dst.clone());
let connect_start = std::time::Instant::now();
Box::pin(async move {
match scheme.as_deref() {
Some("http") => {
let tcp = connecting.await.map_err(|err| Box::new(err) as BoxError)?;
let connect_ms = connect_start.elapsed().as_millis() as u64;
Ok(TimedConn::new(
MaybeHttpsStream::Http(tcp),
ConnectTiming {
connect_ms,
tls_ms: 0,
},
))
}
Some("https") => {
let server_name = resolve_server_name(&dst)?;
let tcp = connecting.await.map_err(|err| Box::new(err) as BoxError)?;
let connect_ms = connect_start.elapsed().as_millis() as u64;
let tls_start = std::time::Instant::now();
let tls_stream = TlsConnector::from(tls_config)
.connect(server_name, tcp.into_inner())
.await
.map_err(io::Error::other)?;
let tls_ms = tls_start.elapsed().as_millis() as u64;
Ok(TimedConn::new(
MaybeHttpsStream::Https(TokioIo::new(tls_stream)),
ConnectTiming { connect_ms, tls_ms },
))
}
Some(other) => Err(io::Error::other(format!("unsupported scheme {other}")).into()),
None => Err(io::Error::other("missing scheme").into()),
}
})
}
}
pub fn build_upstream_client(config: &Config, dns_cache: Arc<DnsCache>) -> UpstreamClient {
let mut http = HttpConnector::new_with_resolver(ValidatedResolver::new(dns_cache));
http.enforce_http(false);
http.set_connect_timeout(Some(Duration::from_secs(
config.upstream_connect_timeout_secs,
)));
http.set_nodelay(config.upstream_tcp_nodelay);
if config.upstream_tcp_keepalive_secs > 0 {
http.set_keepalive(Some(Duration::from_secs(
config.upstream_tcp_keepalive_secs,
)));
} else {
http.set_keepalive(None);
}
let connector = InstrumentedConnector {
http,
tls_config: build_tls_config(),
};
let mut builder = Client::builder(TokioExecutor::new());
builder.pool_max_idle_per_host(config.upstream_pool_max_idle_per_host);
builder.pool_idle_timeout(Duration::from_secs(config.upstream_pool_idle_timeout_secs));
builder.pool_timer(TokioTimer::new());
builder.build(connector)
}
pub fn resolve_request_timing<B>(
response: &Response<B>,
connection_acquire_ms: Option<u64>,
ttfb_ms: u64,
) -> RequestTiming {
let raw = response
.extensions()
.get::<ConnectTiming>()
.copied()
.unwrap_or_default();
let raw_connection_ms = raw.connect_ms.saturating_add(raw.tls_ms);
let measured_acquire_ms = connection_acquire_ms.unwrap_or(raw_connection_ms.min(ttfb_ms));
let likely_reused = measured_acquire_ms <= 5 && raw_connection_ms > 0;
let connector_matches_request = raw_connection_ms <= measured_acquire_ms.saturating_add(25);
let (connect_ms, tls_ms) = if likely_reused || !connector_matches_request {
(0, 0)
} else {
(raw.connect_ms, raw.tls_ms)
};
RequestTiming {
connection_acquire_ms: measured_acquire_ms,
connect_ms,
tls_ms,
response_wait_ms: ttfb_ms.saturating_sub(measured_acquire_ms),
connection_reused: likely_reused,
}
}
fn build_tls_config() -> Arc<ClientConfig> {
let root_store =
rustls::RootCertStore::from_iter(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
let mut config = ClientConfig::builder()
.with_root_certificates(root_store)
.with_no_client_auth();
config.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec()];
Arc::new(config)
}
fn resolve_server_name(uri: &Uri) -> Result<ServerName<'static>, BoxError> {
let host = uri.host().ok_or_else(|| io::Error::other("missing host"))?;
let host = host.trim_start_matches('[').trim_end_matches(']');
if let Ok(ip) = host.parse::<IpAddr>() {
return Ok(ServerName::from(ip));
}
Ok(ServerName::try_from(host.to_string())?)
}
pub struct TimedConn {
inner: MaybeHttpsStream,
timing: ConnectTiming,
}
impl TimedConn {
fn new(inner: MaybeHttpsStream, timing: ConnectTiming) -> Self {
Self { inner, timing }
}
}
impl Connection for TimedConn {
fn connected(&self) -> Connected {
self.inner.connected().extra(self.timing)
}
}
impl rt::Read for TimedConn {
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: rt::ReadBufCursor<'_>,
) -> Poll<Result<(), io::Error>> {
Pin::new(&mut self.inner).poll_read(cx, buf)
}
}
impl rt::Write for TimedConn {
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<Result<usize, io::Error>> {
Pin::new(&mut self.inner).poll_write(cx, buf)
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
Pin::new(&mut self.inner).poll_flush(cx)
}
fn poll_shutdown(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Result<(), io::Error>> {
Pin::new(&mut self.inner).poll_shutdown(cx)
}
fn is_write_vectored(&self) -> bool {
self.inner.is_write_vectored()
}
fn poll_write_vectored(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
bufs: &[std::io::IoSlice<'_>],
) -> Poll<Result<usize, io::Error>> {
Pin::new(&mut self.inner).poll_write_vectored(cx, bufs)
}
}
pub enum MaybeHttpsStream {
Http(PlainStream),
Https(TlsStream),
}
impl Connection for MaybeHttpsStream {
fn connected(&self) -> Connected {
match self {
Self::Http(stream) => stream.connected(),
Self::Https(stream) => {
let (tcp, tls) = stream.inner().get_ref();
if tls.alpn_protocol() == Some(b"h2") {
tcp.connected().negotiated_h2()
} else {
tcp.connected()
}
}
}
}
}
impl rt::Read for MaybeHttpsStream {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: rt::ReadBufCursor<'_>,
) -> Poll<Result<(), io::Error>> {
match Pin::get_mut(self) {
Self::Http(stream) => Pin::new(stream).poll_read(cx, buf),
Self::Https(stream) => Pin::new(stream).poll_read(cx, buf),
}
}
}
impl rt::Write for MaybeHttpsStream {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<Result<usize, io::Error>> {
match Pin::get_mut(self) {
Self::Http(stream) => Pin::new(stream).poll_write(cx, buf),
Self::Https(stream) => Pin::new(stream).poll_write(cx, buf),
}
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
match Pin::get_mut(self) {
Self::Http(stream) => Pin::new(stream).poll_flush(cx),
Self::Https(stream) => Pin::new(stream).poll_flush(cx),
}
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
match Pin::get_mut(self) {
Self::Http(stream) => Pin::new(stream).poll_shutdown(cx),
Self::Https(stream) => Pin::new(stream).poll_shutdown(cx),
}
}
fn is_write_vectored(&self) -> bool {
match self {
Self::Http(stream) => stream.is_write_vectored(),
Self::Https(stream) => stream.is_write_vectored(),
}
}
fn poll_write_vectored(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
bufs: &[std::io::IoSlice<'_>],
) -> Poll<Result<usize, io::Error>> {
match Pin::get_mut(self) {
Self::Http(stream) => Pin::new(stream).poll_write_vectored(cx, bufs),
Self::Https(stream) => Pin::new(stream).poll_write_vectored(cx, bufs),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use hyper::Response;
#[test]
fn fresh_connection_uses_connector_breakdown() {
let mut response = Response::new(());
response.extensions_mut().insert(ConnectTiming {
connect_ms: 80,
tls_ms: 40,
});
let timing = resolve_request_timing(&response, Some(125), 600);
assert_eq!(timing.connection_acquire_ms, 125);
assert_eq!(timing.connect_ms, 80);
assert_eq!(timing.tls_ms, 40);
assert_eq!(timing.response_wait_ms, 475);
assert!(!timing.connection_reused);
}
#[test]
fn reused_connection_zeroes_stale_connect_timings() {
let mut response = Response::new(());
response.extensions_mut().insert(ConnectTiming {
connect_ms: 70,
tls_ms: 30,
});
let timing = resolve_request_timing(&response, Some(0), 310);
assert_eq!(timing.connection_acquire_ms, 0);
assert_eq!(timing.connect_ms, 0);
assert_eq!(timing.tls_ms, 0);
assert_eq!(timing.response_wait_ms, 310);
assert!(timing.connection_reused);
}
#[test]
fn falls_back_to_connector_timings_when_capture_missing() {
let mut response = Response::new(());
response.extensions_mut().insert(ConnectTiming {
connect_ms: 55,
tls_ms: 25,
});
let timing = resolve_request_timing(&response, None, 400);
assert_eq!(timing.connection_acquire_ms, 80);
assert_eq!(timing.connect_ms, 55);
assert_eq!(timing.tls_ms, 25);
assert_eq!(timing.response_wait_ms, 320);
assert!(!timing.connection_reused);
}
}
+25 -5
View File
@@ -3,13 +3,15 @@ Alembic 环境配置
用于数据库迁移的运行时环境设置
"""
from logging.config import fileConfig
from sqlalchemy import engine_from_config, pool
from alembic import context
import os
import sys
from logging.config import fileConfig
from pathlib import Path
from sqlalchemy import engine_from_config, pool, text
from alembic import context
# 添加项目根目录到 Python 路径
sys.path.insert(0, os.path.dirname(os.path.dirname(__file__)))
@@ -48,6 +50,11 @@ if config.config_file_name is not None:
# 目标元数据(包含所有表定义)
target_metadata = Base.metadata
# PostgreSQL 全局迁移锁,避免多进程并发执行 Alembic 导致竞态(重复加列/索引等)
# 使用会话级 advisory lock(pg_advisory_lock),在迁移完成后手动释放。
# ID 由 crc32("aether-alembic-migration") 拼接生成,仅需全局唯一即可。
MIGRATION_ADVISORY_LOCK_ID = 582694137405821
def run_migrations_offline() -> None:
"""
@@ -83,15 +90,28 @@ def run_migrations_online() -> None:
)
with connectable.connect() as connection:
try:
# 使用会话级 advisory lock(非事务级),避免干扰 Alembic 的事务管理。
# pg_advisory_lock 在会话结束时自动释放,不受 COMMIT/ROLLBACK 影响。
if connection.dialect.name == "postgresql":
connection.execute(
text("SELECT pg_advisory_lock(:lock_id)"),
{"lock_id": MIGRATION_ADVISORY_LOCK_ID},
)
connection.commit()
context.configure(
connection=connection,
target_metadata=target_metadata,
compare_type=True, # 比较列类型变更
compare_server_default=True, # 比较默认值变更
compare_type=True,
compare_server_default=True,
transaction_per_migration=True, # 每个迁移文件独立事务,完成即提交
)
with context.begin_transaction():
context.run_migrations()
except Exception:
raise
# 根据模式选择运行方式
@@ -10,8 +10,7 @@ from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from sqlalchemy import inspect
from sqlalchemy import text
from alembic import op
@@ -24,33 +23,24 @@ depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
conn = op.get_bind()
inspector = inspect(conn)
existing_columns = {col["name"] for col in inspector.get_columns("usage")}
if "provider_request_body" not in existing_columns:
op.add_column("usage", sa.Column("provider_request_body", sa.JSON(), nullable=True))
if "provider_request_body_compressed" not in existing_columns:
op.add_column(
"usage", sa.Column("provider_request_body_compressed", sa.LargeBinary(), nullable=True)
# Use PostgreSQL native IF NOT EXISTS to avoid duplicate-column races
# when migrations are triggered concurrently (e.g. startup + manual run).
conn.execute(text("ALTER TABLE usage ADD COLUMN IF NOT EXISTS provider_request_body JSON"))
conn.execute(
text("ALTER TABLE usage ADD COLUMN IF NOT EXISTS provider_request_body_compressed BYTEA")
)
if "client_response_body" not in existing_columns:
op.add_column("usage", sa.Column("client_response_body", sa.JSON(), nullable=True))
if "client_response_body_compressed" not in existing_columns:
op.add_column(
"usage", sa.Column("client_response_body_compressed", sa.LargeBinary(), nullable=True)
conn.execute(text("ALTER TABLE usage ADD COLUMN IF NOT EXISTS client_response_body JSON"))
conn.execute(
text("ALTER TABLE usage ADD COLUMN IF NOT EXISTS client_response_body_compressed BYTEA")
)
def downgrade() -> None:
conn = op.get_bind()
inspector = inspect(conn)
existing_columns = {col["name"] for col in inspector.get_columns("usage")}
for col in (
"client_response_body_compressed",
"client_response_body",
"provider_request_body_compressed",
"provider_request_body",
):
if col in existing_columns:
op.drop_column("usage", col)
conn.execute(text(f"ALTER TABLE usage DROP COLUMN IF EXISTS {col}"))
@@ -0,0 +1,106 @@
"""Add tunnel mode fields and remove IP forwarding fields
Revision ID: 9a0b1c2d3e4f
Revises: 8f9a0b1c2d3e
Create Date: 2026-02-24 17:00:00.000000
"""
from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from sqlalchemy import inspect
from alembic import op
revision: str = "9a0b1c2d3e4f"
down_revision: str | None = "8f9a0b1c2d3e"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
columns = [c["name"] for c in inspector.get_columns(table_name)]
return column_name in columns
def upgrade() -> None:
# 添加 tunnel 模式字段
if not column_exists("proxy_nodes", "tunnel_mode"):
op.add_column(
"proxy_nodes",
sa.Column(
"tunnel_mode",
sa.Boolean(),
nullable=False,
server_default=sa.text("false"),
comment="是否使用 WebSocket 隧道模式",
),
)
if not column_exists("proxy_nodes", "tunnel_connected"):
op.add_column(
"proxy_nodes",
sa.Column(
"tunnel_connected",
sa.Boolean(),
nullable=False,
server_default=sa.text("false"),
comment="隧道是否已连接",
),
)
if not column_exists("proxy_nodes", "tunnel_connected_at"):
op.add_column(
"proxy_nodes",
sa.Column(
"tunnel_connected_at",
sa.DateTime(timezone=True),
nullable=True,
comment="隧道最近一次建立时间",
),
)
# tunnel 模式节点不需要 port,将其置零
op.execute("UPDATE proxy_nodes SET port = 0 WHERE tunnel_mode = true")
# 移除旧的 IP 转发字段
if column_exists("proxy_nodes", "tls_enabled"):
op.drop_column("proxy_nodes", "tls_enabled")
if column_exists("proxy_nodes", "tls_cert_fingerprint"):
op.drop_column("proxy_nodes", "tls_cert_fingerprint")
def downgrade() -> None:
# 恢复 IP 转发字段
if not column_exists("proxy_nodes", "tls_cert_fingerprint"):
op.add_column(
"proxy_nodes",
sa.Column(
"tls_cert_fingerprint",
sa.String(128),
nullable=True,
comment="TLS 证书 SHA-256 指纹(hex)",
),
)
if not column_exists("proxy_nodes", "tls_enabled"):
op.add_column(
"proxy_nodes",
sa.Column(
"tls_enabled",
sa.Boolean(),
nullable=False,
server_default=sa.text("false"),
comment="是否启用 TLS 加密",
),
)
# 移除 tunnel 模式字段
if column_exists("proxy_nodes", "tunnel_connected_at"):
op.drop_column("proxy_nodes", "tunnel_connected_at")
if column_exists("proxy_nodes", "tunnel_connected"):
op.drop_column("proxy_nodes", "tunnel_connected")
if column_exists("proxy_nodes", "tunnel_mode"):
op.drop_column("proxy_nodes", "tunnel_mode")
@@ -0,0 +1,224 @@
"""Add cache_creation columns, clean up capability settings, add user_model_usage_counts,
enforce global_model_id NOT NULL
1. Add cache_creation_input_tokens_5m and cache_creation_input_tokens_1h to usage table.
2. Clean up cache_1h/context_1m/gemini_files from user-configurable settings
(now auto-detected via REQUEST_PARAM mode).
3. Create user_model_usage_counts table for per-user per-model atomic usage counters.
4. Enforce models.global_model_id NOT NULL (delete orphan models without global model).
Revision ID: b2c3d4e5f6a7
Revises: 9a0b1c2d3e4f
Create Date: 2026-02-28 14:00:00.000000
"""
from __future__ import annotations
import json
import uuid
from collections.abc import Sequence
from datetime import datetime, timezone
import sqlalchemy as sa
from sqlalchemy import inspect
from alembic import op
revision: str = "b2c3d4e5f6a7"
down_revision: str | None = "9a0b1c2d3e4f"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
insp = inspect(bind)
columns = [c["name"] for c in insp.get_columns(table_name)]
return column_name in columns
def table_exists(table_name: str) -> bool:
bind = op.get_bind()
insp = inspect(bind)
return table_name in insp.get_table_names()
def index_exists(table_name: str, index_name: str) -> bool:
bind = op.get_bind()
insp = inspect(bind)
return any(idx["name"] == index_name for idx in insp.get_indexes(table_name))
def upgrade() -> None:
# --- 1. Add cache_creation columns ---
if not column_exists("usage", "cache_creation_input_tokens_5m"):
op.add_column(
"usage",
sa.Column(
"cache_creation_input_tokens_5m",
sa.Integer(),
nullable=False,
server_default=sa.text("0"),
comment="5min TTL cache creation input tokens",
),
)
if not column_exists("usage", "cache_creation_input_tokens_1h"):
op.add_column(
"usage",
sa.Column(
"cache_creation_input_tokens_1h",
sa.Integer(),
nullable=False,
server_default=sa.text("0"),
comment="1h TTL cache creation input tokens",
),
)
# --- 2. Clean up stale capability settings (pure Python, DB-agnostic) ---
stale_keys = {"cache_1h", "context_1m", "gemini_files"}
conn = op.get_bind()
# ApiKey.force_capabilities: dict-like JSON, remove stale keys
rows = conn.execute(
sa.text("SELECT id, force_capabilities FROM api_keys WHERE force_capabilities IS NOT NULL")
).fetchall()
for row in rows:
raw = row[1]
if raw is None:
continue
data = raw if isinstance(raw, dict) else json.loads(raw)
cleaned = {k: v for k, v in data.items() if k not in stale_keys}
new_val = json.dumps(cleaned) if cleaned else None
conn.execute(
sa.text("UPDATE api_keys SET force_capabilities = :val WHERE id = :id"),
{"val": new_val, "id": row[0]},
)
# User.model_capability_settings: nested dict {model_key: {cap: val}}, remove stale keys
rows = conn.execute(
sa.text(
"SELECT id, model_capability_settings FROM users"
" WHERE model_capability_settings IS NOT NULL"
)
).fetchall()
for row in rows:
raw = row[1]
if raw is None:
continue
data = raw if isinstance(raw, dict) else json.loads(raw)
cleaned = {}
for model_key, caps in data.items():
cap_cleaned = {k: v for k, v in caps.items() if k not in stale_keys}
if cap_cleaned:
cleaned[model_key] = cap_cleaned
new_val = json.dumps(cleaned) if cleaned else None
conn.execute(
sa.text("UPDATE users SET model_capability_settings = :val WHERE id = :id"),
{"val": new_val, "id": row[0]},
)
# GlobalModel.supported_capabilities: JSON array, remove stale entries
rows = conn.execute(
sa.text(
"SELECT id, supported_capabilities FROM global_models"
" WHERE supported_capabilities IS NOT NULL"
)
).fetchall()
for row in rows:
raw = row[1]
if raw is None:
continue
data = raw if isinstance(raw, list) else json.loads(raw)
cleaned = [c for c in data if c not in stale_keys]
new_val = json.dumps(cleaned) if cleaned else None
conn.execute(
sa.text("UPDATE global_models SET supported_capabilities = :val WHERE id = :id"),
{"val": new_val, "id": row[0]},
)
# --- 3. Create user_model_usage_counts table ---
if not table_exists("user_model_usage_counts"):
op.create_table(
"user_model_usage_counts",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column(
"user_id",
sa.String(36),
sa.ForeignKey("users.id", ondelete="CASCADE"),
nullable=False,
),
sa.Column("model", sa.String(100), nullable=False),
sa.Column("usage_count", sa.Integer, nullable=False, server_default="0"),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.func.now(),
),
sa.Column(
"updated_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.func.now(),
),
sa.UniqueConstraint("user_id", "model", name="uq_user_model_usage_count"),
)
if not index_exists("user_model_usage_counts", "idx_user_model_usage_user"):
op.create_index("idx_user_model_usage_user", "user_model_usage_counts", ["user_id"])
if not index_exists("user_model_usage_counts", "idx_user_model_usage_model"):
op.create_index("idx_user_model_usage_model", "user_model_usage_counts", ["model"])
# Backfill from existing usage records (truncate first for idempotency)
conn.execute(sa.text("DELETE FROM user_model_usage_counts"))
rows = conn.execute(
sa.text(
"SELECT user_id, model, COUNT(*) AS cnt FROM usage"
" WHERE user_id IS NOT NULL GROUP BY user_id, model"
)
).fetchall()
now = datetime.now(timezone.utc)
for row in rows:
conn.execute(
sa.text(
"INSERT INTO user_model_usage_counts"
" (id, user_id, model, usage_count, created_at, updated_at)"
" VALUES (:id, :user_id, :model, :cnt, :now, :now)"
),
{
"id": str(uuid.uuid4()),
"user_id": row[0],
"model": row[1],
"cnt": row[2],
"now": now,
},
)
# --- 4. Enforce models.global_model_id NOT NULL ---
conn = op.get_bind()
insp = inspect(conn)
model_cols = {c["name"]: c for c in insp.get_columns("models")}
if model_cols.get("global_model_id", {}).get("nullable", True):
op.execute("DELETE FROM models WHERE global_model_id IS NULL")
op.alter_column("models", "global_model_id", existing_type=sa.String(36), nullable=False)
def downgrade() -> None:
# Revert models.global_model_id to nullable
if column_exists("models", "global_model_id"):
op.alter_column("models", "global_model_id", existing_type=sa.String(36), nullable=True)
# Drop user_model_usage_counts
if table_exists("user_model_usage_counts"):
if index_exists("user_model_usage_counts", "idx_user_model_usage_model"):
op.drop_index("idx_user_model_usage_model", table_name="user_model_usage_counts")
if index_exists("user_model_usage_counts", "idx_user_model_usage_user"):
op.drop_index("idx_user_model_usage_user", table_name="user_model_usage_counts")
op.drop_table("user_model_usage_counts")
# Drop cache_creation columns
if column_exists("usage", "cache_creation_input_tokens_1h"):
op.drop_column("usage", "cache_creation_input_tokens_1h")
if column_exists("usage", "cache_creation_input_tokens_5m"):
op.drop_column("usage", "cache_creation_input_tokens_5m")
# capability settings cleanup is not reversible
@@ -0,0 +1,153 @@
"""proxy_node_metrics_and_events
Revision ID: 48afe197cc15
Revises: b2c3d4e5f6a7
Create Date: 2026-02-28 04:33:11.201185+00:00
"""
import sqlalchemy as sa
from sqlalchemy import inspect
from alembic import op
# revision identifiers, used by Alembic.
revision = "48afe197cc15"
down_revision = "b2c3d4e5f6a7"
branch_labels = None
depends_on = None
def _column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
insp = inspect(bind)
columns = [c["name"] for c in insp.get_columns(table_name)]
return column_name in columns
def _table_exists(table_name: str) -> bool:
bind = op.get_bind()
insp = inspect(bind)
return table_name in insp.get_table_names()
def _enum_has_value(enum_name: str, value: str) -> bool:
"""检查 PostgreSQL 枚举类型是否包含指定值"""
bind = op.get_bind()
result = bind.execute(
sa.text(
"SELECT 1 FROM pg_enum e JOIN pg_type t ON e.enumtypid = t.oid"
" WHERE t.typname = :enum_name AND e.enumlabel = :value"
),
{"enum_name": enum_name, "value": value},
)
return result.fetchone() is not None
def upgrade() -> None:
# proxy_nodes: 将已废弃的 unhealthy 状态迁移为 offline,然后从枚举中移除
if _enum_has_value("proxynodestatus", "unhealthy"):
op.execute("UPDATE proxy_nodes SET status = 'offline' WHERE status = 'unhealthy'")
op.execute("ALTER TYPE proxynodestatus RENAME TO proxynodestatus_old")
op.execute("CREATE TYPE proxynodestatus AS ENUM ('online', 'offline')")
# 必须先移除旧枚举类型的 DEFAULT,否则 ALTER TYPE 会因无法转换默认值而报错
op.execute("ALTER TABLE proxy_nodes ALTER COLUMN status DROP DEFAULT")
op.execute(
"ALTER TABLE proxy_nodes ALTER COLUMN status TYPE proxynodestatus"
" USING status::text::proxynodestatus"
)
op.execute("ALTER TABLE proxy_nodes ALTER COLUMN status SET DEFAULT 'online'::proxynodestatus")
op.execute("DROP TYPE proxynodestatus_old")
# proxy_nodes: 新增错误指标字段
if not _column_exists("proxy_nodes", "failed_requests"):
op.add_column(
"proxy_nodes",
sa.Column(
"failed_requests",
sa.BigInteger(),
nullable=False,
server_default="0",
comment="累计失败请求数",
),
)
if not _column_exists("proxy_nodes", "dns_failures"):
op.add_column(
"proxy_nodes",
sa.Column(
"dns_failures",
sa.BigInteger(),
nullable=False,
server_default="0",
comment="累计 DNS 失败数",
),
)
if not _column_exists("proxy_nodes", "stream_errors"):
op.add_column(
"proxy_nodes",
sa.Column(
"stream_errors",
sa.BigInteger(),
nullable=False,
server_default="0",
comment="累计流错误数",
),
)
# proxy_node_events: 连接事件表
if not _table_exists("proxy_node_events"):
op.create_table(
"proxy_node_events",
sa.Column("id", sa.BigInteger(), autoincrement=True, nullable=False),
sa.Column("node_id", sa.String(length=36), nullable=False),
sa.Column(
"event_type",
sa.String(length=20),
nullable=False,
comment="事件类型: connected, disconnected, error",
),
sa.Column(
"detail",
sa.String(length=500),
nullable=True,
comment="事件详情(如断开原因)",
),
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
sa.ForeignKeyConstraint(["node_id"], ["proxy_nodes.id"], ondelete="CASCADE"),
sa.PrimaryKeyConstraint("id"),
)
op.create_index(
"idx_proxy_node_events_node_created",
"proxy_node_events",
["node_id", "created_at"],
)
op.create_index(
op.f("ix_proxy_node_events_node_id"),
"proxy_node_events",
["node_id"],
)
def downgrade() -> None:
# 恢复 proxynodestatus 枚举,加回 unhealthy
if not _enum_has_value("proxynodestatus", "unhealthy"):
op.execute("ALTER TYPE proxynodestatus RENAME TO proxynodestatus_old")
op.execute("CREATE TYPE proxynodestatus AS ENUM ('online', 'unhealthy', 'offline')")
op.execute("ALTER TABLE proxy_nodes ALTER COLUMN status DROP DEFAULT")
op.execute(
"ALTER TABLE proxy_nodes ALTER COLUMN status TYPE proxynodestatus"
" USING status::text::proxynodestatus"
)
op.execute("ALTER TABLE proxy_nodes ALTER COLUMN status SET DEFAULT 'online'::proxynodestatus")
op.execute("DROP TYPE proxynodestatus_old")
if _table_exists("proxy_node_events"):
op.drop_index(op.f("ix_proxy_node_events_node_id"), table_name="proxy_node_events")
op.drop_index("idx_proxy_node_events_node_created", table_name="proxy_node_events")
op.drop_table("proxy_node_events")
if _column_exists("proxy_nodes", "stream_errors"):
op.drop_column("proxy_nodes", "stream_errors")
if _column_exists("proxy_nodes", "dns_failures"):
op.drop_column("proxy_nodes", "dns_failures")
if _column_exists("proxy_nodes", "failed_requests"):
op.drop_column("proxy_nodes", "failed_requests")
@@ -0,0 +1,49 @@
"""add_request_candidates_composite_indexes
Revision ID: 00b9161b8729
Revises: 48afe197cc15
Create Date: 2026-02-28 14:48:00.000000+00:00
"""
from sqlalchemy import inspect
from alembic import op
# revision identifiers, used by Alembic.
revision = "00b9161b8729"
down_revision = "48afe197cc15"
branch_labels = None
depends_on = None
def _index_exists(index_name: str) -> bool:
bind = op.get_bind()
insp = inspect(bind)
indexes = insp.get_indexes("request_candidates")
return any(idx["name"] == index_name for idx in indexes)
def upgrade() -> None:
# (request_id, status) - fallback/retry 查询优化
if not _index_exists("idx_rc_request_id_status"):
op.create_index(
"idx_rc_request_id_status",
"request_candidates",
["request_id", "status"],
)
# (provider_id, status, created_at) - provider 聚合统计优化
if not _index_exists("idx_rc_provider_status_created"):
op.create_index(
"idx_rc_provider_status_created",
"request_candidates",
["provider_id", "status", "created_at"],
)
def downgrade() -> None:
if _index_exists("idx_rc_provider_status_created"):
op.drop_index("idx_rc_provider_status_created", table_name="request_candidates")
if _index_exists("idx_rc_request_id_status"):
op.drop_index("idx_rc_request_id_status", table_name="request_candidates")
@@ -0,0 +1,263 @@
"""vertex_ai_provider_type
Migrate legacy Vertex auth_type/provider_type into the new model:
- provider_type=vertex_ai
- auth_type=service_account (legacy vertex_ai renamed)
- fixed Vertex endpoints: gemini:chat + claude:chat
Revision ID: 2a624af8dd3a
Revises: 00b9161b8729
Create Date: 2026-02-28 15:00:00.000000+00:00
"""
from __future__ import annotations
import uuid
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision = "2a624af8dd3a"
down_revision = "00b9161b8729"
branch_labels = None
depends_on = None
_VERTEX_BASE_URL = "https://aiplatform.googleapis.com"
_VERTEX_ENDPOINTS: tuple[tuple[str, str, str], ...] = (
("gemini:chat", "gemini", "chat"),
("claude:chat", "claude", "chat"),
)
_VERTEX_KEY_FORMATS_SA = '["gemini:chat","claude:chat"]'
_VERTEX_KEY_FORMATS_API_KEY = '["gemini:chat"]'
def _select_vertex_provider_ids(conn: sa.Connection) -> list[str]:
"""Collect providers that should be treated as Vertex after migration."""
rows = conn.execute(sa.text("""
SELECT DISTINCT p.id
FROM providers p
LEFT JOIN provider_api_keys pak ON pak.provider_id = p.id
WHERE lower(COALESCE(p.provider_type, '')) = 'vertex_ai'
OR pak.auth_type = 'vertex_ai'
"""))
return [str(row[0]) for row in rows if row[0]]
def _ensure_fixed_vertex_endpoints(conn: sa.Connection, provider_ids: list[str]) -> None:
"""Ensure every Vertex provider has fixed gemini:chat + claude:chat endpoints."""
for provider_id in provider_ids:
provider_max_retries = (
conn.execute(
sa.text("""
SELECT COALESCE(max_retries, 2)
FROM providers
WHERE id = :provider_id
"""),
{"provider_id": provider_id},
).scalar()
or 2
)
for api_format, api_family, endpoint_kind in _VERTEX_ENDPOINTS:
# Normalize existing fixed endpoint fields.
conn.execute(
sa.text("""
UPDATE provider_endpoints
SET
api_family = :api_family,
endpoint_kind = :endpoint_kind,
base_url = :base_url,
custom_path = NULL,
is_active = TRUE,
updated_at = CURRENT_TIMESTAMP
WHERE provider_id = :provider_id
AND api_format = :api_format
"""),
{
"provider_id": provider_id,
"api_format": api_format,
"api_family": api_family,
"endpoint_kind": endpoint_kind,
"base_url": _VERTEX_BASE_URL,
},
)
exists = conn.execute(
sa.text("""
SELECT 1
FROM provider_endpoints
WHERE provider_id = :provider_id
AND api_format = :api_format
LIMIT 1
"""),
{"provider_id": provider_id, "api_format": api_format},
).first()
if not exists:
conn.execute(
sa.text("""
INSERT INTO provider_endpoints (
id,
provider_id,
api_format,
api_family,
endpoint_kind,
base_url,
custom_path,
header_rules,
body_rules,
max_retries,
is_active,
config,
format_acceptance_config,
proxy,
created_at,
updated_at
)
VALUES (
:id,
:provider_id,
:api_format,
:api_family,
:endpoint_kind,
:base_url,
NULL,
NULL,
NULL,
:max_retries,
TRUE,
NULL,
NULL,
NULL,
CURRENT_TIMESTAMP,
CURRENT_TIMESTAMP
)
"""),
{
"id": str(uuid.uuid4()),
"provider_id": provider_id,
"api_format": api_format,
"api_family": api_family,
"endpoint_kind": endpoint_kind,
"base_url": _VERTEX_BASE_URL,
"max_retries": int(provider_max_retries),
},
)
# Vertex fixed-provider model: disable non-fixed endpoints.
conn.execute(
sa.text("""
UPDATE provider_endpoints
SET
is_active = FALSE,
updated_at = CURRENT_TIMESTAMP
WHERE provider_id = :provider_id
AND api_format NOT IN ('gemini:chat', 'claude:chat')
"""),
{"provider_id": provider_id},
)
def _normalize_vertex_key_formats(conn: sa.Connection, provider_ids: list[str]) -> None:
"""Normalize key.api_formats for Vertex keys by auth type."""
for provider_id in provider_ids:
# Service Account (and legacy vertex_ai) keys: allow Gemini + Claude models.
conn.execute(
sa.text("""
UPDATE provider_api_keys
SET
api_formats = CAST(:api_formats AS json),
updated_at = CURRENT_TIMESTAMP
WHERE provider_id = :provider_id
AND auth_type IN ('service_account', 'vertex_ai')
"""),
{
"provider_id": provider_id,
"api_formats": _VERTEX_KEY_FORMATS_SA,
},
)
# API Key mode on Vertex 仅支持 Gemini(Google publisher)。
conn.execute(
sa.text("""
UPDATE provider_api_keys
SET
api_formats = CAST(:api_formats AS json),
updated_at = CURRENT_TIMESTAMP
WHERE provider_id = :provider_id
AND auth_type = 'api_key'
"""),
{
"provider_id": provider_id,
"api_formats": _VERTEX_KEY_FORMATS_API_KEY,
},
)
def upgrade() -> None:
conn = op.get_bind()
# 1) 收集目标 Provider(兼容重复执行,先识别 legacy/new 两种来源)。
provider_ids = _select_vertex_provider_ids(conn)
# 2) 先重命名 auth_type(legacy vertex_ai -> service_account)。
conn.execute(sa.text("""
UPDATE provider_api_keys
SET auth_type = 'service_account'
WHERE auth_type = 'vertex_ai'
"""))
if not provider_ids:
return
# 3) 归一 provider_type,并启用格式转换(Vertex 同时承载 Gemini/Claude)。
for provider_id in provider_ids:
conn.execute(
sa.text("""
UPDATE providers
SET
provider_type = 'vertex_ai',
enable_format_conversion = TRUE
WHERE id = :provider_id
"""),
{"provider_id": provider_id},
)
# 4) 固定端点落地:gemini:chat + claude:chat。
_ensure_fixed_vertex_endpoints(conn, provider_ids)
# 5) 归一 key 的 api_formats,避免调度命中旧格式。
_normalize_vertex_key_formats(conn, provider_ids)
def downgrade() -> None:
conn = op.get_bind()
provider_rows = conn.execute(sa.text("""
SELECT id
FROM providers
WHERE lower(COALESCE(provider_type, '')) = 'vertex_ai'
"""))
provider_ids = [str(row[0]) for row in provider_rows if row[0]]
if provider_ids:
for provider_id in provider_ids:
conn.execute(
sa.text("""
UPDATE provider_api_keys
SET auth_type = 'vertex_ai'
WHERE provider_id = :provider_id
AND auth_type = 'service_account'
"""),
{"provider_id": provider_id},
)
conn.execute(
sa.text("""
UPDATE providers
SET provider_type = 'custom'
WHERE id = :provider_id
"""),
{"provider_id": provider_id},
)
@@ -0,0 +1,199 @@
"""backfill_codex_compact_endpoint
Backfill Codex reverse-proxy endpoints:
- ensure `openai:cli` endpoint is pinned to force_stream
- ensure `openai:compact` endpoint exists
Revision ID: f0c3a7b9d1e2
Revises: 2a624af8dd3a
Create Date: 2026-03-01 17:00:00.000000+00:00
"""
from __future__ import annotations
import json
import uuid
from typing import Any
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision = "f0c3a7b9d1e2"
down_revision = "2a624af8dd3a"
branch_labels = None
depends_on = None
_CODEX_BASE_URL = "https://chatgpt.com/backend-api/codex"
_COMPACT_FORMAT = "openai:compact"
_CLI_FORMAT = "openai:cli"
_FORCE_STREAM = "force_stream"
def _find_codex_provider_ids(conn: sa.Connection) -> list[str]:
"""Find Codex providers (by provider_type or legacy base_url pattern)."""
rows = conn.execute(sa.text("""
SELECT DISTINCT p.id
FROM providers p
LEFT JOIN provider_endpoints pe ON pe.provider_id = p.id
WHERE lower(COALESCE(p.provider_type, '')) = 'codex'
OR (
lower(COALESCE(pe.api_format, '')) = 'openai:cli'
AND lower(COALESCE(pe.base_url, '')) LIKE '%/backend-api/codex%'
)
"""))
return [str(r[0]) for r in rows if r[0]]
def _get_cli_endpoint(conn: sa.Connection, provider_id: str) -> dict[str, Any] | None:
"""Load existing openai:cli endpoint for the provider."""
row = (
conn.execute(
sa.text("""
SELECT base_url, header_rules, body_rules, max_retries, proxy, config
FROM provider_endpoints
WHERE provider_id = :pid AND api_format = :fmt
LIMIT 1
"""),
{"pid": provider_id, "fmt": _CLI_FORMAT},
)
.mappings()
.first()
)
return dict(row) if row else None
def _pin_cli_force_stream(conn: sa.Connection, provider_id: str, cli: dict[str, Any]) -> None:
"""Set upstream_stream_policy=force_stream on existing cli endpoint."""
cfg = dict(cli.get("config") or {}) if isinstance(cli.get("config"), dict) else {}
cfg.pop("upstreamStreamPolicy", None)
cfg.pop("upstream_stream", None)
cfg["upstream_stream_policy"] = _FORCE_STREAM
conn.execute(
sa.text("""
UPDATE provider_endpoints
SET api_family = 'openai',
endpoint_kind = 'cli',
config = CAST(:config AS json),
updated_at = CURRENT_TIMESTAMP
WHERE provider_id = :pid AND api_format = :fmt
"""),
{
"pid": provider_id,
"fmt": _CLI_FORMAT,
"config": json.dumps(cfg, ensure_ascii=False),
},
)
def _ensure_compact_endpoint(conn: sa.Connection, provider_id: str, cli: dict[str, Any]) -> None:
"""Create openai:compact endpoint if missing (clone from cli)."""
exists = conn.execute(
sa.text(
"SELECT 1 FROM provider_endpoints WHERE provider_id = :pid AND api_format = :fmt LIMIT 1"
),
{"pid": provider_id, "fmt": _COMPACT_FORMAT},
).first()
if exists:
# Already exists, just ensure api_family/endpoint_kind are set.
conn.execute(
sa.text("""
UPDATE provider_endpoints
SET api_family = 'openai', endpoint_kind = 'compact',
updated_at = CURRENT_TIMESTAMP
WHERE provider_id = :pid AND api_format = :fmt
"""),
{"pid": provider_id, "fmt": _COMPACT_FORMAT},
)
return
# Clone from cli endpoint, strip stream policy.
cfg = dict(cli.get("config") or {}) if isinstance(cli.get("config"), dict) else {}
for k in ("upstream_stream_policy", "upstreamStreamPolicy", "upstream_stream"):
cfg.pop(k, None)
def _json(val: Any) -> str | None:
return json.dumps(val, ensure_ascii=False) if val is not None else None
conn.execute(
sa.text("""
INSERT INTO provider_endpoints (
id, provider_id, api_format, api_family, endpoint_kind,
base_url, custom_path, header_rules, body_rules,
max_retries, is_active, config, format_acceptance_config,
proxy, created_at, updated_at
) VALUES (
:id, :pid, :fmt, 'openai', 'compact',
:base_url, NULL, CAST(:header_rules AS json), CAST(:body_rules AS json),
:max_retries, TRUE, CAST(:config AS json), NULL,
CAST(:proxy AS jsonb), CURRENT_TIMESTAMP, CURRENT_TIMESTAMP
)
"""),
{
"id": str(uuid.uuid4()),
"pid": provider_id,
"fmt": _COMPACT_FORMAT,
"base_url": cli.get("base_url") or _CODEX_BASE_URL,
"header_rules": _json(cli.get("header_rules")),
"body_rules": _json(cli.get("body_rules")),
"max_retries": cli.get("max_retries") or 2,
"config": _json(cfg or None),
"proxy": _json(cli.get("proxy")),
},
)
def _add_compact_to_key_formats(conn: sa.Connection, provider_id: str) -> None:
"""Ensure provider keys include openai:compact in api_formats."""
rows = (
conn.execute(
sa.text("SELECT id, api_formats FROM provider_api_keys WHERE provider_id = :pid"),
{"pid": provider_id},
)
.mappings()
.all()
)
for row in rows:
raw = row["api_formats"]
formats: list[str] = []
if isinstance(raw, list):
for item in raw:
v = str(item or "").strip().lower()
if v and v not in formats:
formats.append(v)
if _COMPACT_FORMAT in formats:
continue
# Insert compact right after cli, or at end.
if _CLI_FORMAT in formats:
idx = formats.index(_CLI_FORMAT) + 1
formats.insert(idx, _COMPACT_FORMAT)
else:
formats.append(_COMPACT_FORMAT)
conn.execute(
sa.text("""
UPDATE provider_api_keys
SET api_formats = CAST(:fmts AS json), updated_at = CURRENT_TIMESTAMP
WHERE id = :id
"""),
{"id": row["id"], "fmts": json.dumps(formats, ensure_ascii=False)},
)
def upgrade() -> None:
conn = op.get_bind()
for provider_id in _find_codex_provider_ids(conn):
cli = _get_cli_endpoint(conn, provider_id)
if not cli:
continue # No cli endpoint to clone from; skip.
_pin_cli_force_stream(conn, provider_id, cli)
_ensure_compact_endpoint(conn, provider_id, cli)
_add_compact_to_key_formats(conn, provider_id)
def downgrade() -> None:
# Data backfill: no-op to avoid deleting user-managed data.
return
@@ -0,0 +1,29 @@
"""add_proxy_metadata_to_proxy_nodes
Revision ID: 1d2e3f4a5b6c
Revises: f0c3a7b9d1e2
Create Date: 2026-03-02 13:00:00.000000+00:00
"""
from __future__ import annotations
from collections.abc import Sequence
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "1d2e3f4a5b6c"
down_revision: str | None = "f0c3a7b9d1e2"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
op.execute("ALTER TABLE public.proxy_nodes ADD COLUMN IF NOT EXISTS proxy_metadata json")
op.execute(
"COMMENT ON COLUMN public.proxy_nodes.proxy_metadata IS 'aether-proxy 上报元数据(版本等)'"
)
def downgrade() -> None:
op.execute("ALTER TABLE public.proxy_nodes DROP COLUMN IF EXISTS proxy_metadata")
@@ -0,0 +1,72 @@
"""backfill_codex_default_body_rules
Backfill default body_rules for codex providers with openai:cli endpoints
that currently have body_rules IS NULL.
Rules:
- drop max_output_tokens
- drop temperature
- drop top_p
- set store = false
- set instructions = "You are GPT-5." (when instructions not exists)
Revision ID: dd0278c0a28c
Revises: 1d2e3f4a5b6c
Create Date: 2026-03-02 15:00:00.000000+00:00
"""
from __future__ import annotations
import json
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision = "dd0278c0a28c"
down_revision = "1d2e3f4a5b6c"
branch_labels = None
depends_on = None
_TARGET_FORMATS = ("openai:cli",)
_DEFAULT_BODY_RULES = [
{"action": "drop", "path": "max_output_tokens"},
{"action": "drop", "path": "temperature"},
{"action": "drop", "path": "top_p"},
{"action": "set", "path": "store", "value": False},
{
"action": "set",
"path": "instructions",
"value": "You are GPT-5.",
"condition": {"path": "instructions", "op": "not_exists"},
},
]
def upgrade() -> None:
conn = op.get_bind()
# 幂等性: 仅回填 codex 提供商中 body_rules 为空(SQL NULL 或 JSON null)的记录
rules_json = json.dumps(_DEFAULT_BODY_RULES, ensure_ascii=False)
result = conn.execute(
sa.text("""
UPDATE provider_endpoints pe
SET body_rules = CAST(:rules AS json),
updated_at = CURRENT_TIMESTAMP
FROM providers p
WHERE pe.provider_id = p.id
AND p.provider_type = :ptype
AND pe.api_format = :fmt
AND (pe.body_rules IS NULL OR pe.body_rules::text = 'null')
"""),
{"rules": rules_json, "ptype": "codex", "fmt": _TARGET_FORMATS[0]},
)
if result.rowcount:
print(f" backfilled body_rules for {result.rowcount} endpoint(s)")
def downgrade() -> None:
# Data backfill: no-op to avoid removing user-customized rules.
return
@@ -0,0 +1,45 @@
"""add_idx_usage_provider_key
Add composite index on usage(provider_id, provider_api_key_id) to support
the pool management page's per-key usage stats aggregation query.
Revision ID: 0ba031f328de
Revises: dd0278c0a28c
Create Date: 2026-03-03 10:00:00.000000+00:00
"""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
revision = "0ba031f328de"
down_revision = "dd0278c0a28c"
branch_labels = None
depends_on = None
INDEX_NAME = "idx_usage_provider_key"
TABLE = "usage"
COLUMNS = ["provider_id", "provider_api_key_id"]
def upgrade() -> None:
bind = op.get_bind()
result = bind.execute(
sa.text("SELECT 1 FROM pg_indexes WHERE indexname = :name"),
{"name": INDEX_NAME},
).fetchone()
if result:
return
op.create_index(INDEX_NAME, TABLE, COLUMNS)
def downgrade() -> None:
bind = op.get_bind()
result = bind.execute(
sa.text("SELECT 1 FROM pg_indexes WHERE indexname = :name"),
{"name": INDEX_NAME},
).fetchone()
if not result:
return
op.drop_index(INDEX_NAME, table_name=TABLE)
@@ -0,0 +1,45 @@
"""add_idx_usage_status_user_created
Add composite index on usage(status, user_id, created_at) to speed up
interval timeline and active usage analytics queries.
Revision ID: 5f1d2e3c4b5a
Revises: 0ba031f328de
Create Date: 2026-03-03 17:30:00.000000+00:00
"""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
revision = "5f1d2e3c4b5a"
down_revision = "0ba031f328de"
branch_labels = None
depends_on = None
INDEX_NAME = "idx_usage_status_user_created"
TABLE = "usage"
COLUMNS = ["status", "user_id", "created_at"]
def upgrade() -> None:
bind = op.get_bind()
result = bind.execute(
sa.text("SELECT 1 FROM pg_indexes WHERE indexname = :name"),
{"name": INDEX_NAME},
).fetchone()
if result:
return
op.create_index(INDEX_NAME, TABLE, COLUMNS)
def downgrade() -> None:
bind = op.get_bind()
result = bind.execute(
sa.text("SELECT 1 FROM pg_indexes WHERE indexname = :name"),
{"name": INDEX_NAME},
).fetchone()
if not result:
return
op.drop_index(INDEX_NAME, table_name=TABLE)
@@ -0,0 +1,41 @@
"""add fingerprint column to provider_api_keys
Revision ID: 6a9b8c7d5e4f
Revises: 5f1d2e3c4b5a
Create Date: 2026-03-04 23:50:00.000000
"""
from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from sqlalchemy import inspect
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "6a9b8c7d5e4f"
down_revision: str | None = "5f1d2e3c4b5a"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
columns = [c["name"] for c in inspector.get_columns(table_name)]
return column_name in columns
def upgrade() -> None:
if not column_exists("provider_api_keys", "fingerprint"):
op.add_column(
"provider_api_keys",
sa.Column("fingerprint", sa.JSON(), nullable=True),
)
def downgrade() -> None:
if column_exists("provider_api_keys", "fingerprint"):
op.drop_column("provider_api_keys", "fingerprint")
@@ -0,0 +1,65 @@
"""remove standalone api key locking
Revision ID: 7c91d2e4f8a1
Revises: 6f7a8b9c0d1e
Create Date: 2026-03-05 17:00:00.000000+00:00
"""
from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from sqlalchemy import inspect
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "7c91d2e4f8a1"
down_revision: str | None = "6f7a8b9c0d1e"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
_CONSTRAINT_NAME = "ck_api_keys_standalone_not_locked"
def _column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
insp = inspect(bind)
insp.clear_cache()
return column_name in [c["name"] for c in insp.get_columns(table_name)]
def _check_constraint_exists(table_name: str, constraint_name: str) -> bool:
bind = op.get_bind()
insp = inspect(bind)
insp.clear_cache()
return any(c.get("name") == constraint_name for c in insp.get_check_constraints(table_name))
def upgrade() -> None:
if not (
_column_exists("api_keys", "is_standalone")
and _column_exists("api_keys", "is_locked")
and _column_exists("api_keys", "is_active")
):
return
op.execute(sa.text("""
UPDATE api_keys
SET is_active = FALSE,
is_locked = FALSE
WHERE is_standalone IS TRUE AND is_locked IS TRUE
"""))
if not _check_constraint_exists("api_keys", _CONSTRAINT_NAME):
op.create_check_constraint(
_CONSTRAINT_NAME,
"api_keys",
"(NOT is_standalone) OR (NOT is_locked)",
)
def downgrade() -> None:
if _check_constraint_exists("api_keys", _CONSTRAINT_NAME):
op.drop_constraint(_CONSTRAINT_NAME, "api_keys", type_="check")
@@ -0,0 +1,220 @@
"""tighten wallet transaction snapshots and remove wallet version
Revision ID: 8e71f2a4c9b0
Revises: 7c91d2e4f8a1
Create Date: 2026-03-07 13:00:00.000000+00:00
"""
from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from sqlalchemy import inspect
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "8e71f2a4c9b0"
down_revision: str | None = "7c91d2e4f8a1"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
_WALLET_TX_BEFORE_CHECK = "ck_wallet_tx_balance_before_consistent"
_WALLET_TX_AFTER_CHECK = "ck_wallet_tx_balance_after_consistent"
_WALLET_LIMIT_MODE_INDEX = "idx_wallets_limit_mode"
def _table_exists(table_name: str) -> bool:
bind = op.get_bind()
insp = inspect(bind)
insp.clear_cache()
return table_name in insp.get_table_names()
def _column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
insp = inspect(bind)
insp.clear_cache()
return column_name in [c["name"] for c in insp.get_columns(table_name)]
def _index_exists(table_name: str, index_name: str) -> bool:
bind = op.get_bind()
insp = inspect(bind)
insp.clear_cache()
return any(index.get("name") == index_name for index in insp.get_indexes(table_name))
def _check_constraint_exists(table_name: str, constraint_name: str) -> bool:
bind = op.get_bind()
insp = inspect(bind)
insp.clear_cache()
return any(c.get("name") == constraint_name for c in insp.get_check_constraints(table_name))
def _tighten_wallet_transaction_snapshots() -> None:
if not _table_exists("wallet_transactions"):
return
required_columns = {
"balance_before",
"balance_after",
"recharge_balance_before",
"recharge_balance_after",
"gift_balance_before",
"gift_balance_after",
}
existing_columns = {
column["name"] for column in inspect(op.get_bind()).get_columns("wallet_transactions")
}
if not required_columns.issubset(existing_columns):
return
op.execute(
sa.text(
"""
UPDATE wallet_transactions
SET recharge_balance_before = balance_before
WHERE recharge_balance_before IS NULL
"""
)
)
op.execute(
sa.text(
"""
UPDATE wallet_transactions
SET recharge_balance_after = balance_after
WHERE recharge_balance_after IS NULL
"""
)
)
op.execute(
sa.text(
"""
UPDATE wallet_transactions
SET gift_balance_before = 0
WHERE gift_balance_before IS NULL
"""
)
)
op.execute(
sa.text(
"""
UPDATE wallet_transactions
SET gift_balance_after = 0
WHERE gift_balance_after IS NULL
"""
)
)
op.execute(
sa.text(
"""
UPDATE wallet_transactions
SET balance_before = recharge_balance_before + gift_balance_before,
balance_after = recharge_balance_after + gift_balance_after
"""
)
)
if not _check_constraint_exists("wallet_transactions", _WALLET_TX_BEFORE_CHECK):
op.create_check_constraint(
_WALLET_TX_BEFORE_CHECK,
"wallet_transactions",
"balance_before = recharge_balance_before + gift_balance_before",
)
if not _check_constraint_exists("wallet_transactions", _WALLET_TX_AFTER_CHECK):
op.create_check_constraint(
_WALLET_TX_AFTER_CHECK,
"wallet_transactions",
"balance_after = recharge_balance_after + gift_balance_after",
)
op.alter_column(
"wallet_transactions",
"recharge_balance_before",
existing_type=sa.Numeric(20, 8),
nullable=False,
)
op.alter_column(
"wallet_transactions",
"recharge_balance_after",
existing_type=sa.Numeric(20, 8),
nullable=False,
)
op.alter_column(
"wallet_transactions",
"gift_balance_before",
existing_type=sa.Numeric(20, 8),
nullable=False,
)
op.alter_column(
"wallet_transactions",
"gift_balance_after",
existing_type=sa.Numeric(20, 8),
nullable=False,
)
def _drop_wallet_cleanup_artifacts() -> None:
if not _table_exists("wallets"):
return
if _index_exists("wallets", _WALLET_LIMIT_MODE_INDEX):
op.drop_index(_WALLET_LIMIT_MODE_INDEX, table_name="wallets")
if _column_exists("wallets", "version"):
op.drop_column("wallets", "version")
def upgrade() -> None:
_tighten_wallet_transaction_snapshots()
_drop_wallet_cleanup_artifacts()
def downgrade() -> None:
if _table_exists("wallets"):
if not _column_exists("wallets", "version"):
op.add_column(
"wallets",
sa.Column("version", sa.Integer(), nullable=False, server_default="0"),
)
if not _index_exists("wallets", _WALLET_LIMIT_MODE_INDEX):
op.create_index(_WALLET_LIMIT_MODE_INDEX, "wallets", ["limit_mode"])
if not _table_exists("wallet_transactions"):
return
if _column_exists("wallet_transactions", "recharge_balance_before"):
op.alter_column(
"wallet_transactions",
"recharge_balance_before",
existing_type=sa.Numeric(20, 8),
nullable=True,
)
if _column_exists("wallet_transactions", "recharge_balance_after"):
op.alter_column(
"wallet_transactions",
"recharge_balance_after",
existing_type=sa.Numeric(20, 8),
nullable=True,
)
if _column_exists("wallet_transactions", "gift_balance_before"):
op.alter_column(
"wallet_transactions",
"gift_balance_before",
existing_type=sa.Numeric(20, 8),
nullable=True,
)
if _column_exists("wallet_transactions", "gift_balance_after"):
op.alter_column(
"wallet_transactions",
"gift_balance_after",
existing_type=sa.Numeric(20, 8),
nullable=True,
)
if _check_constraint_exists("wallet_transactions", _WALLET_TX_AFTER_CHECK):
op.drop_constraint(_WALLET_TX_AFTER_CHECK, "wallet_transactions", type_="check")
if _check_constraint_exists("wallet_transactions", _WALLET_TX_BEFORE_CHECK):
op.drop_constraint(_WALLET_TX_BEFORE_CHECK, "wallet_transactions", type_="check")
@@ -0,0 +1,80 @@
"""add missing foreign key indexes for cascade delete performance
Revision ID: 2d932114930d
Revises: 8e71f2a4c9b0
Create Date: 2026-03-07 16:28:48.633531+00:00
"""
from alembic import op
from sqlalchemy import inspect
# revision identifiers, used by Alembic.
revision = '2d932114930d'
down_revision = '8e71f2a4c9b0'
branch_labels = None
depends_on = None
def _index_exists(table_name: str, index_name: str) -> bool:
bind = op.get_bind()
insp = inspect(bind)
return any(idx["name"] == index_name for idx in insp.get_indexes(table_name))
def _create_index_if_not_exists(index_name: str, table_name: str, columns: list[str]) -> None:
if not _index_exists(table_name, index_name):
op.create_index(op.f(index_name), table_name, columns, unique=False)
def _drop_index_if_exists(index_name: str, table_name: str) -> None:
if _index_exists(table_name, index_name):
op.drop_index(op.f(index_name), table_name=table_name)
# (index_name, table_name, columns)
_INDEXES = [
# api_keys.user_id (CASCADE -> users.id)
('ix_api_keys_user_id', 'api_keys', ['user_id']),
# usage: wallet_id, provider_endpoint_id, provider_api_key_id (SET NULL)
('ix_usage_wallet_id', 'usage', ['wallet_id']),
('ix_usage_provider_endpoint_id', 'usage', ['provider_endpoint_id']),
('ix_usage_provider_api_key_id', 'usage', ['provider_api_key_id']),
# wallet_transactions.operator_id (SET NULL -> users.id)
('ix_wallet_transactions_operator_id', 'wallet_transactions', ['operator_id']),
# payment_callbacks.payment_order_id (SET NULL -> payment_orders.id)
('ix_payment_callbacks_payment_order_id', 'payment_callbacks', ['payment_order_id']),
# refund_requests: payment_order_id, requested_by, approved_by, processed_by (SET NULL)
('ix_refund_requests_payment_order_id', 'refund_requests', ['payment_order_id']),
('ix_refund_requests_requested_by', 'refund_requests', ['requested_by']),
('ix_refund_requests_approved_by', 'refund_requests', ['approved_by']),
('ix_refund_requests_processed_by', 'refund_requests', ['processed_by']),
# proxy_nodes.registered_by (SET NULL -> users.id)
('ix_proxy_nodes_registered_by', 'proxy_nodes', ['registered_by']),
# video_tasks: api_key_id, provider_id, endpoint_id, key_id, remixed_from_task_id
('ix_video_tasks_api_key_id', 'video_tasks', ['api_key_id']),
('ix_video_tasks_provider_id', 'video_tasks', ['provider_id']),
('ix_video_tasks_endpoint_id', 'video_tasks', ['endpoint_id']),
('ix_video_tasks_key_id', 'video_tasks', ['key_id']),
('ix_video_tasks_remixed_from_task_id', 'video_tasks', ['remixed_from_task_id']),
# user_preferences.default_provider_id (-> providers.id)
('ix_user_preferences_default_provider_id', 'user_preferences', ['default_provider_id']),
# announcements.author_id (SET NULL -> users.id)
('ix_announcements_author_id', 'announcements', ['author_id']),
# announcement_reads.announcement_id (-> announcements.id)
('ix_announcement_reads_announcement_id', 'announcement_reads', ['announcement_id']),
# request_candidates: user_id, api_key_id, endpoint_id, key_id (CASCADE)
('ix_request_candidates_user_id', 'request_candidates', ['user_id']),
('ix_request_candidates_api_key_id', 'request_candidates', ['api_key_id']),
('ix_request_candidates_endpoint_id', 'request_candidates', ['endpoint_id']),
('ix_request_candidates_key_id', 'request_candidates', ['key_id']),
]
def upgrade() -> None:
for index_name, table_name, columns in _INDEXES:
_create_index_if_not_exists(index_name, table_name, columns)
def downgrade() -> None:
for index_name, table_name, _columns in reversed(_INDEXES):
_drop_index_if_exists(index_name, table_name)
@@ -0,0 +1,224 @@
"""usage stats retention: SET NULL on delete and add name snapshots
Revision ID: 45b118150a78
Revises: 2d932114930d
Create Date: 2026-03-08 03:48:49.622091+00:00
"""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision = "45b118150a78"
down_revision = "2d932114930d"
branch_labels = None
depends_on = None
_TABLES = ["usage", "stats_user_daily", "stats_daily_api_key"]
# ---------------------------------------------------------------------------
# Inline helpers
# ---------------------------------------------------------------------------
class _SchemaCache:
def __init__(self) -> None:
self._columns: dict[str, dict[str, str]] = {}
self._fk_rules: dict[tuple[str, str], str] = {}
self._fk_loaded_tables: set[str] = set()
def load_columns(self, tables: list[str]) -> None:
need = [t for t in tables if t not in self._columns]
if not need:
return
bind = op.get_bind()
rows = bind.execute(
sa.text(
"SELECT table_name, column_name, data_type "
"FROM information_schema.columns "
"WHERE table_name = ANY(:tables) "
" AND table_schema = current_schema()"
),
{"tables": need},
).fetchall()
for t in need:
self._columns.setdefault(t, {})
for table, col, dtype in rows:
self._columns[table][col] = dtype
def load_fk_rules(self, tables: list[str]) -> None:
need = [t for t in tables if t not in self._fk_loaded_tables]
if not need:
return
bind = op.get_bind()
rows = bind.execute(
sa.text(
"SELECT tc.table_name, tc.constraint_name, rc.delete_rule "
"FROM information_schema.referential_constraints rc "
"JOIN information_schema.table_constraints tc "
" ON rc.constraint_name = tc.constraint_name "
" AND rc.constraint_schema = tc.constraint_schema "
"WHERE tc.table_name = ANY(:tables) "
" AND tc.table_schema = current_schema()"
),
{"tables": need},
).fetchall()
for table, name, rule in rows:
self._fk_rules[(table, name)] = rule
self._fk_loaded_tables.update(need)
def column_exists(self, table: str, column: str) -> bool:
return column in self._columns.get(table, {})
def fk_ondelete(self, table: str, constraint: str) -> str | None:
return self._fk_rules.get((table, constraint))
def _fk_exists(constraint_name: str, table_name: str) -> bool:
bind = op.get_bind()
result = bind.execute(
sa.text(
"SELECT 1 FROM pg_constraint c "
"JOIN pg_class r ON c.conrelid = r.oid "
"JOIN pg_namespace n ON r.relnamespace = n.oid "
"WHERE c.conname = :name AND r.relname = :table "
" AND n.nspname = current_schema() AND c.contype = 'f'"
),
{"name": constraint_name, "table": table_name},
)
return result.scalar() is not None
def _replace_fk_if_needed(
cache: _SchemaCache,
constraint_name: str,
table_name: str,
ref_table: str,
local_cols: list[str],
remote_cols: list[str],
desired_ondelete: str,
) -> None:
current = cache.fk_ondelete(table_name, constraint_name)
if current and current.upper() == desired_ondelete.upper():
return
if current or _fk_exists(constraint_name, table_name):
op.drop_constraint(constraint_name, table_name, type_="foreignkey")
op.create_foreign_key(
constraint_name,
table_name,
ref_table,
local_cols,
remote_cols,
ondelete=desired_ondelete,
)
# ---------------------------------------------------------------------------
def upgrade() -> None:
c = _SchemaCache()
c.load_columns(_TABLES)
c.load_fk_rules(["stats_user_daily", "stats_daily_api_key"])
# --- Usage: add name snapshot columns ---
if not c.column_exists("usage", "username"):
op.add_column(
"usage", sa.Column("username", sa.String(100), nullable=True, comment="用户名快照")
)
if not c.column_exists("usage", "api_key_name"):
op.add_column(
"usage",
sa.Column("api_key_name", sa.String(200), nullable=True, comment="API Key 名称快照"),
)
# --- StatsUserDaily: CASCADE -> SET NULL, add username snapshot ---
_replace_fk_if_needed(
c,
"stats_user_daily_user_id_fkey",
"stats_user_daily",
"users",
["user_id"],
["id"],
"SET NULL",
)
op.alter_column("stats_user_daily", "user_id", existing_type=sa.String(36), nullable=True)
if not c.column_exists("stats_user_daily", "username"):
op.add_column(
"stats_user_daily",
sa.Column(
"username",
sa.String(100),
nullable=True,
comment="用户名快照(删除用户后仍可追溯)",
),
)
# --- StatsDailyApiKey: CASCADE -> SET NULL, add api_key_name snapshot ---
_replace_fk_if_needed(
c,
"stats_daily_api_key_api_key_id_fkey",
"stats_daily_api_key",
"api_keys",
["api_key_id"],
["id"],
"SET NULL",
)
op.alter_column("stats_daily_api_key", "api_key_id", existing_type=sa.String(36), nullable=True)
if not c.column_exists("stats_daily_api_key", "api_key_name"):
op.add_column(
"stats_daily_api_key",
sa.Column(
"api_key_name",
sa.String(200),
nullable=True,
comment="API Key 名称快照(删除 Key 后仍可追溯)",
),
)
def downgrade() -> None:
c = _SchemaCache()
c.load_columns(["stats_daily_api_key", "stats_user_daily", "usage"])
c.load_fk_rules(["stats_daily_api_key", "stats_user_daily"])
# --- Remove snapshot columns ---
if c.column_exists("stats_daily_api_key", "api_key_name"):
op.drop_column("stats_daily_api_key", "api_key_name")
if c.column_exists("stats_user_daily", "username"):
op.drop_column("stats_user_daily", "username")
if c.column_exists("usage", "api_key_name"):
op.drop_column("usage", "api_key_name")
if c.column_exists("usage", "username"):
op.drop_column("usage", "username")
# --- StatsDailyApiKey: SET NULL -> CASCADE ---
_replace_fk_if_needed(
c,
"stats_daily_api_key_api_key_id_fkey",
"stats_daily_api_key",
"api_keys",
["api_key_id"],
["id"],
"CASCADE",
)
op.alter_column(
"stats_daily_api_key", "api_key_id", existing_type=sa.String(36), nullable=False
)
# --- StatsUserDaily: SET NULL -> CASCADE ---
_replace_fk_if_needed(
c,
"stats_user_daily_user_id_fkey",
"stats_user_daily",
"users",
["user_id"],
["id"],
"CASCADE",
)
op.alter_column("stats_user_daily", "user_id", existing_type=sa.String(36), nullable=False)
@@ -0,0 +1,257 @@
"""request_candidates/video_tasks retention: SET NULL and add snapshots
Revision ID: 13a4c8f6d9e0
Revises: 45b118150a78
Create Date: 2026-03-08 12:15:00.000000+00:00
"""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision = "13a4c8f6d9e0"
down_revision = "45b118150a78"
branch_labels = None
depends_on = None
_TABLES = ["request_candidates", "video_tasks"]
# ---------------------------------------------------------------------------
# Inline helpers
# ---------------------------------------------------------------------------
class _SchemaCache:
def __init__(self) -> None:
self._columns: dict[str, dict[str, str]] = {}
self._fk_rules: dict[tuple[str, str], str] = {}
self._fk_loaded_tables: set[str] = set()
def load_columns(self, tables: list[str]) -> None:
need = [t for t in tables if t not in self._columns]
if not need:
return
bind = op.get_bind()
rows = bind.execute(
sa.text(
"SELECT table_name, column_name, data_type "
"FROM information_schema.columns "
"WHERE table_name = ANY(:tables) "
" AND table_schema = current_schema()"
),
{"tables": need},
).fetchall()
for t in need:
self._columns.setdefault(t, {})
for table, col, dtype in rows:
self._columns[table][col] = dtype
def load_fk_rules(self, tables: list[str]) -> None:
need = [t for t in tables if t not in self._fk_loaded_tables]
if not need:
return
bind = op.get_bind()
rows = bind.execute(
sa.text(
"SELECT tc.table_name, tc.constraint_name, rc.delete_rule "
"FROM information_schema.referential_constraints rc "
"JOIN information_schema.table_constraints tc "
" ON rc.constraint_name = tc.constraint_name "
" AND rc.constraint_schema = tc.constraint_schema "
"WHERE tc.table_name = ANY(:tables) "
" AND tc.table_schema = current_schema()"
),
{"tables": need},
).fetchall()
for table, name, rule in rows:
self._fk_rules[(table, name)] = rule
self._fk_loaded_tables.update(need)
def column_exists(self, table: str, column: str) -> bool:
return column in self._columns.get(table, {})
def fk_ondelete(self, table: str, constraint: str) -> str | None:
return self._fk_rules.get((table, constraint))
def _fk_exists(constraint_name: str, table_name: str) -> bool:
bind = op.get_bind()
result = bind.execute(
sa.text(
"SELECT 1 FROM pg_constraint c "
"JOIN pg_class r ON c.conrelid = r.oid "
"JOIN pg_namespace n ON r.relnamespace = n.oid "
"WHERE c.conname = :name AND r.relname = :table "
" AND n.nspname = current_schema() AND c.contype = 'f'"
),
{"name": constraint_name, "table": table_name},
)
return result.scalar() is not None
def _replace_fk_if_needed(
cache: _SchemaCache,
constraint_name: str,
table_name: str,
ref_table: str,
local_cols: list[str],
remote_cols: list[str],
desired_ondelete: str,
) -> None:
current = cache.fk_ondelete(table_name, constraint_name)
if current and current.upper() == desired_ondelete.upper():
return
if current or _fk_exists(constraint_name, table_name):
op.drop_constraint(constraint_name, table_name, type_="foreignkey")
op.create_foreign_key(
constraint_name,
table_name,
ref_table,
local_cols,
remote_cols,
ondelete=desired_ondelete,
)
# ---------------------------------------------------------------------------
def upgrade() -> None:
c = _SchemaCache()
c.load_columns(_TABLES)
c.load_fk_rules(_TABLES)
# --- request_candidates: add snapshot columns ---
if not c.column_exists("request_candidates", "username"):
op.add_column(
"request_candidates",
sa.Column("username", sa.String(length=100), nullable=True, comment="用户名快照"),
)
if not c.column_exists("request_candidates", "api_key_name"):
op.add_column(
"request_candidates",
sa.Column(
"api_key_name",
sa.String(length=200),
nullable=True,
comment="API Key 名称快照",
),
)
# --- request_candidates: CASCADE -> SET NULL ---
_replace_fk_if_needed(
c,
"request_candidates_user_id_fkey",
"request_candidates",
"users",
["user_id"],
["id"],
"SET NULL",
)
_replace_fk_if_needed(
c,
"request_candidates_api_key_id_fkey",
"request_candidates",
"api_keys",
["api_key_id"],
["id"],
"SET NULL",
)
# --- video_tasks: add snapshot columns ---
if not c.column_exists("video_tasks", "username"):
op.add_column(
"video_tasks",
sa.Column("username", sa.String(length=100), nullable=True, comment="用户名快照"),
)
if not c.column_exists("video_tasks", "api_key_name"):
op.add_column(
"video_tasks",
sa.Column(
"api_key_name",
sa.String(length=200),
nullable=True,
comment="API Key 名称快照",
),
)
# --- video_tasks: CASCADE -> SET NULL, user_id nullable ---
op.alter_column("video_tasks", "user_id", existing_type=sa.String(length=36), nullable=True)
_replace_fk_if_needed(
c,
"video_tasks_user_id_fkey",
"video_tasks",
"users",
["user_id"],
["id"],
"SET NULL",
)
_replace_fk_if_needed(
c,
"video_tasks_api_key_id_fkey",
"video_tasks",
"api_keys",
["api_key_id"],
["id"],
"SET NULL",
)
def downgrade() -> None:
c = _SchemaCache()
c.load_columns(_TABLES)
c.load_fk_rules(_TABLES)
# --- video_tasks: SET NULL -> default (no action), restore NOT NULL ---
_replace_fk_if_needed(
c,
"video_tasks_api_key_id_fkey",
"video_tasks",
"api_keys",
["api_key_id"],
["id"],
"NO ACTION",
)
_replace_fk_if_needed(
c,
"video_tasks_user_id_fkey",
"video_tasks",
"users",
["user_id"],
["id"],
"NO ACTION",
)
op.alter_column("video_tasks", "user_id", existing_type=sa.String(length=36), nullable=False)
if c.column_exists("video_tasks", "api_key_name"):
op.drop_column("video_tasks", "api_key_name")
if c.column_exists("video_tasks", "username"):
op.drop_column("video_tasks", "username")
# --- request_candidates: SET NULL -> CASCADE ---
_replace_fk_if_needed(
c,
"request_candidates_api_key_id_fkey",
"request_candidates",
"api_keys",
["api_key_id"],
["id"],
"CASCADE",
)
_replace_fk_if_needed(
c,
"request_candidates_user_id_fkey",
"request_candidates",
"users",
["user_id"],
["id"],
"CASCADE",
)
if c.column_exists("request_candidates", "api_key_name"):
op.drop_column("request_candidates", "api_key_name")
if c.column_exists("request_candidates", "username"):
op.drop_column("request_candidates", "username")
@@ -0,0 +1,230 @@
"""cost fields: Float -> Numeric(20,8) + provider_api_keys composite index
Revision ID: 2053ab8ed764
Revises: 13a4c8f6d9e0
Create Date: 2026-03-08 15:30:00.000000+00:00
"""
from __future__ import annotations
import re
from collections import defaultdict
from collections.abc import Callable
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision = "2053ab8ed764"
down_revision = "13a4c8f6d9e0"
branch_labels = None
depends_on = None
# (table_name, column_name, nullable, server_default)
_COST_COLUMNS: list[tuple[str, str, bool, str | None]] = [
# api_keys
("api_keys", "total_cost_usd", True, "0.0"),
# usage
("usage", "input_cost_usd", True, "0.0"),
("usage", "output_cost_usd", True, "0.0"),
("usage", "cache_cost_usd", True, "0.0"),
("usage", "cache_creation_cost_usd", True, "0.0"),
("usage", "cache_read_cost_usd", True, "0.0"),
("usage", "request_cost_usd", True, "0.0"),
("usage", "total_cost_usd", True, "0.0"),
("usage", "actual_input_cost_usd", True, "0.0"),
("usage", "actual_output_cost_usd", True, "0.0"),
("usage", "actual_cache_creation_cost_usd", True, "0.0"),
("usage", "actual_cache_read_cost_usd", True, "0.0"),
("usage", "actual_request_cost_usd", True, "0.0"),
("usage", "actual_total_cost_usd", True, "0.0"),
("usage", "rate_multiplier", True, "1.0"),
("usage", "input_price_per_1m", True, None),
("usage", "output_price_per_1m", True, None),
("usage", "cache_creation_price_per_1m", True, None),
("usage", "cache_read_price_per_1m", True, None),
("usage", "price_per_request", True, None),
# providers
("providers", "monthly_quota_usd", True, None),
("providers", "monthly_used_usd", True, "0.0"),
# global_models
("global_models", "default_price_per_request", True, None),
# models
("models", "price_per_request", True, None),
# stats_hourly
("stats_hourly", "total_cost", False, "0.0"),
("stats_hourly", "actual_total_cost", False, "0.0"),
# stats_hourly_user
("stats_hourly_user", "total_cost", False, "0.0"),
# stats_hourly_model
("stats_hourly_model", "total_cost", False, "0.0"),
# stats_hourly_provider
("stats_hourly_provider", "total_cost", False, "0.0"),
# stats_daily
("stats_daily", "total_cost", False, "0.0"),
("stats_daily", "actual_total_cost", False, "0.0"),
("stats_daily", "input_cost", False, "0.0"),
("stats_daily", "output_cost", False, "0.0"),
("stats_daily", "cache_creation_cost", False, "0.0"),
("stats_daily", "cache_read_cost", False, "0.0"),
# stats_daily_model
("stats_daily_model", "total_cost", False, "0.0"),
# stats_daily_provider
("stats_daily_provider", "total_cost", False, "0.0"),
# stats_daily_api_key
("stats_daily_api_key", "total_cost", False, "0.0"),
# stats_summary
("stats_summary", "all_time_cost", False, "0.0"),
("stats_summary", "all_time_actual_cost", False, "0.0"),
# stats_user_daily
("stats_user_daily", "total_cost", False, "0.0"),
]
_ALL_TABLES = list({t for t, *_ in _COST_COLUMNS})
# ---------------------------------------------------------------------------
# Inline helpers
# ---------------------------------------------------------------------------
class _SchemaCache:
def __init__(self) -> None:
self._columns: dict[str, dict[str, str]] = {}
def load_columns(self, tables: list[str]) -> None:
need = [t for t in tables if t not in self._columns]
if not need:
return
bind = op.get_bind()
rows = bind.execute(
sa.text(
"SELECT table_name, column_name, data_type "
"FROM information_schema.columns "
"WHERE table_name = ANY(:tables) "
" AND table_schema = current_schema()"
),
{"tables": need},
).fetchall()
for t in need:
self._columns.setdefault(t, {})
for table, col, dtype in rows:
self._columns[table][col] = dtype
def column_exists(self, table: str, column: str) -> bool:
return column in self._columns.get(table, {})
def column_type(self, table: str, column: str) -> str | None:
return self._columns.get(table, {}).get(column)
def is_numeric(self, table: str, column: str) -> bool:
return self.column_type(table, column) == "numeric"
def _index_exists(index_name: str) -> bool:
bind = op.get_bind()
result = bind.execute(
sa.text(
"SELECT 1 FROM pg_indexes "
"WHERE indexname = :name AND schemaname = current_schema()::text"
),
{"name": index_name},
)
return result.scalar() is not None
def _numeric_max(type_spec: str) -> float | None:
m = re.match(r"NUMERIC\((\d+),(\d+)\)", type_spec, re.IGNORECASE)
if not m:
return None
precision, scale = int(m.group(1)), int(m.group(2))
return 10 ** (precision - scale) - 10 ** (-scale)
def _batch_alter_type(
cache: _SchemaCache,
columns: list[tuple[str, str, bool, str | None]],
cast_suffix: str,
type_fn: Callable[[str], str],
) -> None:
by_table: dict[str, list[tuple[str, str]]] = defaultdict(list)
for table, col, _nullable, _default in columns:
if not cache.column_exists(table, col):
continue
by_table[table].append((col, type_fn(col)))
bind = op.get_bind()
for table, col_types in by_table.items():
for col, target in col_types:
cap = _numeric_max(target)
if cap is not None:
bind.execute(
sa.text(
f"UPDATE {table} SET {col} = :cap "
f"WHERE {col} IS NOT NULL AND abs({col}) > :cap"
),
{"cap": cap},
)
parts = [
f"ALTER COLUMN {col} TYPE {target} USING {col}::{cast_suffix}"
for col, target in col_types
]
if parts:
bind.execute(sa.text(f"ALTER TABLE {table} " + ", ".join(parts)))
# ---------------------------------------------------------------------------
def _type_spec(col: str) -> str:
"""Return the SQL type literal for a given column name."""
return "NUMERIC(10,6)" if col == "rate_multiplier" else "NUMERIC(20,8)"
def upgrade() -> None:
c = _SchemaCache()
c.load_columns(_ALL_TABLES)
# -- 1. cost fields: Float -> Numeric (batched per table)
cols_to_convert = [
(t, col, n, d)
for t, col, n, d in _COST_COLUMNS
if c.column_exists(t, col) and not c.is_numeric(t, col)
]
_batch_alter_type(c, cols_to_convert, cast_suffix="numeric", type_fn=_type_spec)
# -- 2. provider_api_keys composite index
if not _index_exists("idx_provider_api_keys_provider_active"):
op.create_index(
"idx_provider_api_keys_provider_active",
"provider_api_keys",
["provider_id", "is_active"],
)
def downgrade() -> None:
# -- 2. drop composite index
if _index_exists("idx_provider_api_keys_provider_active"):
op.drop_index(
"idx_provider_api_keys_provider_active",
table_name="provider_api_keys",
)
# -- 1. Numeric -> Float (batched per table)
c = _SchemaCache()
c.load_columns(_ALL_TABLES)
cols_to_revert = [
(t, col, n, d)
for t, col, n, d in _COST_COLUMNS
if c.column_exists(t, col) and c.is_numeric(t, col)
]
_batch_alter_type(
c,
cols_to_revert,
cast_suffix="double precision",
type_fn=lambda _col: "DOUBLE PRECISION",
)
@@ -0,0 +1,127 @@
"""video_tasks.key_id: add ondelete SET NULL
Revision ID: d7649c1f8e21
Revises: 2053ab8ed764
Create Date: 2026-03-09 01:00:00.000000+00:00
"""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision = "d7649c1f8e21"
down_revision = "2053ab8ed764"
branch_labels = None
depends_on = None
_TABLE = "video_tasks"
_FK_NAME = "video_tasks_key_id_fkey"
# ---------------------------------------------------------------------------
# Inline helpers
# ---------------------------------------------------------------------------
class _SchemaCache:
def __init__(self) -> None:
self._fk_rules: dict[tuple[str, str], str] = {}
self._fk_loaded_tables: set[str] = set()
def load_fk_rules(self, tables: list[str]) -> None:
need = [t for t in tables if t not in self._fk_loaded_tables]
if not need:
return
bind = op.get_bind()
rows = bind.execute(
sa.text(
"SELECT tc.table_name, tc.constraint_name, rc.delete_rule "
"FROM information_schema.referential_constraints rc "
"JOIN information_schema.table_constraints tc "
" ON rc.constraint_name = tc.constraint_name "
" AND rc.constraint_schema = tc.constraint_schema "
"WHERE tc.table_name = ANY(:tables) "
" AND tc.table_schema = current_schema()"
),
{"tables": need},
).fetchall()
for table, name, rule in rows:
self._fk_rules[(table, name)] = rule
self._fk_loaded_tables.update(need)
def fk_ondelete(self, table: str, constraint: str) -> str | None:
return self._fk_rules.get((table, constraint))
def _fk_exists(constraint_name: str, table_name: str) -> bool:
bind = op.get_bind()
result = bind.execute(
sa.text(
"SELECT 1 FROM pg_constraint c "
"JOIN pg_class r ON c.conrelid = r.oid "
"JOIN pg_namespace n ON r.relnamespace = n.oid "
"WHERE c.conname = :name AND r.relname = :table "
" AND n.nspname = current_schema() AND c.contype = 'f'"
),
{"name": constraint_name, "table": table_name},
)
return result.scalar() is not None
def _replace_fk_if_needed(
cache: _SchemaCache,
constraint_name: str,
table_name: str,
ref_table: str,
local_cols: list[str],
remote_cols: list[str],
desired_ondelete: str,
) -> None:
current = cache.fk_ondelete(table_name, constraint_name)
if current and current.upper() == desired_ondelete.upper():
return
if current or _fk_exists(constraint_name, table_name):
op.drop_constraint(constraint_name, table_name, type_="foreignkey")
op.create_foreign_key(
constraint_name,
table_name,
ref_table,
local_cols,
remote_cols,
ondelete=desired_ondelete,
)
# ---------------------------------------------------------------------------
def upgrade() -> None:
c = _SchemaCache()
c.load_fk_rules([_TABLE])
_replace_fk_if_needed(
c,
_FK_NAME,
_TABLE,
"provider_api_keys",
["key_id"],
["id"],
"SET NULL",
)
def downgrade() -> None:
c = _SchemaCache()
c.load_fk_rules([_TABLE])
_replace_fk_if_needed(
c,
_FK_NAME,
_TABLE,
"provider_api_keys",
["key_id"],
["id"],
"NO ACTION",
)
@@ -0,0 +1,42 @@
"""Strip request_results_window from health_by_format JSON.
This data is now maintained in process memory only, no longer persisted to DB.
Revision ID: a3f1b7c9d2e4
Revises: d7649c1f8e21
Create Date: 2026-03-10 12:00:00.000000+00:00
"""
from alembic import op
# revision identifiers, used by Alembic.
revision = "a3f1b7c9d2e4"
down_revision = "d7649c1f8e21"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.execute("""
UPDATE provider_api_keys
SET health_by_format = (
SELECT jsonb_object_agg(
fmt_key,
fmt_value - 'request_results_window'
)
FROM jsonb_each(health_by_format) AS x(fmt_key, fmt_value)
)
WHERE health_by_format IS NOT NULL
AND health_by_format != '{}'::jsonb
AND EXISTS (
SELECT 1
FROM jsonb_each(health_by_format) AS x(fmt_key, fmt_value)
WHERE fmt_value ? 'request_results_window'
)
""")
def downgrade() -> None:
# No-op: window data is rebuilt from scratch on process start
pass
@@ -0,0 +1,159 @@
"""tighten usage billing state machine
Revision ID: 9e4f1a2b3c4d
Revises: a3f1b7c9d2e4
Create Date: 2026-03-11 19:00:00.000000+00:00
This migration does two things:
1. Change new `usage.billing_status` default from `settled` to `pending`.
2. Repair only the clearly-safe inconsistent historical rows for production:
- failed/cancelled zero-cost rows that were marked settled are converted to void
- terminal rows missing finalized_at are backfilled from created_at
Ambiguous positive-cost settled rows are intentionally left untouched for manual audit.
All data updates are batched (10000 rows per iteration) to avoid long-held locks
and excessive WAL generation on large usage tables.
"""
from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from sqlalchemy import inspect
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "9e4f1a2b3c4d"
down_revision: str | None = "a3f1b7c9d2e4"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
BATCH_SIZE = 10000
def _table_exists(table_name: str) -> bool:
bind = op.get_bind()
insp = inspect(bind)
insp.clear_cache()
return table_name in insp.get_table_names()
def _column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
insp = inspect(bind)
insp.clear_cache()
return column_name in [col["name"] for col in insp.get_columns(table_name)]
def upgrade() -> None:
if not _table_exists("usage"):
return
if _column_exists("usage", "billing_status"):
op.alter_column(
"usage",
"billing_status",
existing_type=sa.String(length=20),
server_default="pending",
existing_nullable=False,
)
required_columns = {
"billing_status",
"status",
"total_cost_usd",
"request_cost_usd",
"actual_total_cost_usd",
"actual_request_cost_usd",
"wallet_balance_after",
"finalized_at",
"created_at",
}
if not required_columns.issubset(
{col for col in required_columns if _column_exists("usage", col)}
):
return
conn = op.get_bind()
# Step 1: billing_status IS NULL -> 'pending' (batched)
while True:
result = conn.execute(
sa.text("""
WITH batch AS (
SELECT id FROM usage
WHERE billing_status IS NULL
LIMIT :batch_size
FOR UPDATE SKIP LOCKED
)
UPDATE usage
SET billing_status = 'pending'
FROM batch WHERE usage.id = batch.id
"""),
{"batch_size": BATCH_SIZE},
)
if result.rowcount < BATCH_SIZE:
break
# Step 2: failed/cancelled zero-cost settled -> void (batched)
while True:
result = conn.execute(
sa.text("""
WITH batch AS (
SELECT id FROM usage
WHERE billing_status = 'settled'
AND status IN ('failed', 'cancelled')
AND COALESCE(total_cost_usd, 0) = 0
AND wallet_balance_after IS NULL
LIMIT :batch_size
FOR UPDATE SKIP LOCKED
)
UPDATE usage
SET billing_status = 'void',
finalized_at = COALESCE(usage.finalized_at, usage.created_at),
total_cost_usd = 0,
request_cost_usd = 0,
actual_total_cost_usd = 0,
actual_request_cost_usd = 0
FROM batch WHERE usage.id = batch.id
"""),
{"batch_size": BATCH_SIZE},
)
if result.rowcount < BATCH_SIZE:
break
# Step 3: backfill finalized_at for terminal rows (batched)
while True:
result = conn.execute(
sa.text("""
WITH batch AS (
SELECT id FROM usage
WHERE billing_status IN ('settled', 'void')
AND finalized_at IS NULL
LIMIT :batch_size
FOR UPDATE SKIP LOCKED
)
UPDATE usage
SET finalized_at = COALESCE(usage.finalized_at, usage.created_at)
FROM batch WHERE usage.id = batch.id
"""),
{"batch_size": BATCH_SIZE},
)
if result.rowcount < BATCH_SIZE:
break
def downgrade() -> None:
if not _table_exists("usage") or not _column_exists("usage", "billing_status"):
return
op.alter_column(
"usage",
"billing_status",
existing_type=sa.String(length=20),
server_default="settled",
existing_nullable=False,
)
@@ -0,0 +1,113 @@
"""add wallet daily usage ledgers
Revision ID: d4e5f6a7b8c9
Revises: 9e4f1a2b3c4d
Create Date: 2026-03-11 21:00:00.000000+00:00
"""
from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from sqlalchemy import inspect
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "d4e5f6a7b8c9"
down_revision: str | None = "9e4f1a2b3c4d"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def _table_exists(table_name: str) -> bool:
bind = op.get_bind()
insp = inspect(bind)
insp.clear_cache()
return table_name in insp.get_table_names()
def _column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
insp = inspect(bind)
insp.clear_cache()
return column_name in [col["name"] for col in insp.get_columns(table_name)]
def _index_exists(table_name: str, index_name: str) -> bool:
bind = op.get_bind()
insp = inspect(bind)
insp.clear_cache()
return any(idx["name"] == index_name for idx in insp.get_indexes(table_name))
def upgrade() -> None:
if not _table_exists("wallet_daily_usage_ledgers"):
op.create_table(
"wallet_daily_usage_ledgers",
sa.Column("id", sa.String(length=36), nullable=False),
sa.Column("wallet_id", sa.String(length=36), nullable=False),
sa.Column("billing_date", sa.Date(), nullable=False),
sa.Column("billing_timezone", sa.String(length=64), nullable=False),
sa.Column("total_cost_usd", sa.Numeric(20, 8), nullable=False, server_default="0"),
sa.Column("total_requests", sa.Integer(), nullable=False, server_default="0"),
sa.Column("input_tokens", sa.BigInteger(), nullable=False, server_default="0"),
sa.Column("output_tokens", sa.BigInteger(), nullable=False, server_default="0"),
sa.Column("cache_creation_tokens", sa.BigInteger(), nullable=False, server_default="0"),
sa.Column("cache_read_tokens", sa.BigInteger(), nullable=False, server_default="0"),
sa.Column("first_finalized_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("last_finalized_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("aggregated_at", sa.DateTime(timezone=True), nullable=False),
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False),
sa.ForeignKeyConstraint(["wallet_id"], ["wallets.id"], ondelete="CASCADE"),
sa.PrimaryKeyConstraint("id"),
sa.UniqueConstraint(
"wallet_id",
"billing_date",
"billing_timezone",
name="uq_wallet_daily_usage_ledgers_wallet_date_tz",
),
)
if not _index_exists("wallet_daily_usage_ledgers", "idx_wallet_daily_usage_wallet_date"):
op.create_index(
"idx_wallet_daily_usage_wallet_date",
"wallet_daily_usage_ledgers",
["wallet_id", "billing_date"],
)
if not _index_exists("wallet_daily_usage_ledgers", "idx_wallet_daily_usage_date"):
op.create_index(
"idx_wallet_daily_usage_date",
"wallet_daily_usage_ledgers",
["billing_date"],
)
if (
_table_exists("usage")
and all(
_column_exists("usage", col) for col in ["billing_status", "finalized_at", "wallet_id"]
)
and not _index_exists("usage", "idx_usage_billing_finalized_wallet")
):
op.create_index(
"idx_usage_billing_finalized_wallet",
"usage",
["billing_status", "finalized_at", "wallet_id"],
)
def downgrade() -> None:
if _table_exists("usage") and _index_exists("usage", "idx_usage_billing_finalized_wallet"):
op.drop_index("idx_usage_billing_finalized_wallet", table_name="usage")
if _table_exists("wallet_daily_usage_ledgers"):
if _index_exists("wallet_daily_usage_ledgers", "idx_wallet_daily_usage_date"):
op.drop_index("idx_wallet_daily_usage_date", table_name="wallet_daily_usage_ledgers")
if _index_exists("wallet_daily_usage_ledgers", "idx_wallet_daily_usage_wallet_date"):
op.drop_index(
"idx_wallet_daily_usage_wallet_date",
table_name="wallet_daily_usage_ledgers",
)
op.drop_table("wallet_daily_usage_ledgers")
@@ -0,0 +1,59 @@
"""add provider_api_keys usage total columns
Revision ID: 9b7c6d5e4f3a
Revises: d4e5f6a7b8c9
Create Date: 2026-03-11 22:00:00.000000+00:00
"""
from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from sqlalchemy import inspect
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "9b7c6d5e4f3a"
down_revision: str | None = "d4e5f6a7b8c9"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
columns = [c["name"] for c in inspector.get_columns(table_name)]
return column_name in columns
def upgrade() -> None:
if not column_exists("provider_api_keys", "total_tokens"):
op.add_column(
"provider_api_keys",
sa.Column("total_tokens", sa.BigInteger(), nullable=False, server_default="0"),
)
if column_exists("provider_api_keys", "total_tokens"):
op.alter_column("provider_api_keys", "total_tokens", server_default=None)
if not column_exists("provider_api_keys", "total_cost_usd"):
op.add_column(
"provider_api_keys",
sa.Column(
"total_cost_usd",
sa.Numeric(20, 8),
nullable=False,
server_default="0.0",
),
)
if column_exists("provider_api_keys", "total_cost_usd"):
op.alter_column("provider_api_keys", "total_cost_usd", server_default=None)
def downgrade() -> None:
if column_exists("provider_api_keys", "total_cost_usd"):
op.drop_column("provider_api_keys", "total_cost_usd")
if column_exists("provider_api_keys", "total_tokens"):
op.drop_column("provider_api_keys", "total_tokens")
@@ -0,0 +1,116 @@
"""cleanup stale provider references after provider deletion
Revision ID: c1d2e3f4a5b6
Revises: 9b7c6d5e4f3a
Create Date: 2026-03-11 23:00:00.000000+00:00
"""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision = "c1d2e3f4a5b6"
down_revision = "9b7c6d5e4f3a"
branch_labels = None
depends_on = None
_users = sa.table(
"users",
sa.column("id", sa.String(36)),
sa.column("allowed_providers", sa.JSON()),
)
_api_keys = sa.table(
"api_keys",
sa.column("id", sa.String(36)),
sa.column("allowed_providers", sa.JSON()),
)
_user_preferences = sa.table(
"user_preferences",
sa.column("id", sa.String(36)),
sa.column("default_provider_id", sa.String(36)),
)
_video_tasks = sa.table(
"video_tasks",
sa.column("id", sa.String(36)),
sa.column("provider_id", sa.String(36)),
sa.column("endpoint_id", sa.String(36)),
)
_providers = sa.table("providers", sa.column("id", sa.String(36)))
_provider_endpoints = sa.table("provider_endpoints", sa.column("id", sa.String(36)))
def _load_valid_ids(conn: sa.Connection, table: sa.Table) -> set[str]:
return {str(row[0]) for row in conn.execute(sa.select(table.c.id)).fetchall() if row[0]}
def _cleanup_allowed_providers(
conn: sa.Connection,
table: sa.Table,
valid_provider_ids: set[str],
) -> None:
rows = conn.execute(
sa.select(table.c.id, table.c.allowed_providers).where(
table.c.allowed_providers.isnot(None)
)
).fetchall()
for row_id, allowed_providers in rows:
if not isinstance(allowed_providers, list):
continue
filtered = [
provider_id for provider_id in allowed_providers if provider_id in valid_provider_ids
]
if filtered == allowed_providers:
continue
conn.execute(table.update().where(table.c.id == row_id).values(allowed_providers=filtered))
def _nullify_missing_fk(
conn: sa.Connection,
table: sa.Table,
id_column: sa.ColumnElement[str],
fk_column: sa.ColumnElement[str],
valid_ids: set[str],
) -> None:
rows = conn.execute(sa.select(id_column, fk_column).where(fk_column.isnot(None))).fetchall()
invalid_row_ids = [row_id for row_id, fk_value in rows if fk_value not in valid_ids]
if not invalid_row_ids:
return
conn.execute(table.update().where(id_column.in_(invalid_row_ids)).values({fk_column.key: None}))
def upgrade() -> None:
conn = op.get_bind()
valid_provider_ids = _load_valid_ids(conn, _providers)
valid_endpoint_ids = _load_valid_ids(conn, _provider_endpoints)
_cleanup_allowed_providers(conn, _users, valid_provider_ids)
_cleanup_allowed_providers(conn, _api_keys, valid_provider_ids)
_nullify_missing_fk(
conn,
_user_preferences,
_user_preferences.c.id,
_user_preferences.c.default_provider_id,
valid_provider_ids,
)
_nullify_missing_fk(
conn,
_video_tasks,
_video_tasks.c.id,
_video_tasks.c.provider_id,
valid_provider_ids,
)
_nullify_missing_fk(
conn,
_video_tasks,
_video_tasks.c.id,
_video_tasks.c.endpoint_id,
valid_endpoint_ids,
)
def downgrade() -> None:
pass
@@ -0,0 +1,64 @@
"""decouple request_candidates.key_id foreign key from provider_api_keys lifecycle
Revision ID: b7c8d9e0f1a2
Revises: c1d2e3f4a5b6
Create Date: 2026-03-12 19:15:00.000000+00:00
"""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision = "b7c8d9e0f1a2"
down_revision = "c1d2e3f4a5b6"
branch_labels = None
depends_on = None
def _fk_exists(constraint_name: str, table_name: str) -> bool:
bind = op.get_bind()
result = bind.execute(
sa.text(
"SELECT 1 FROM pg_constraint c "
"JOIN pg_class r ON c.conrelid = r.oid "
"JOIN pg_namespace n ON r.relnamespace = n.oid "
"WHERE c.conname = :name AND r.relname = :table "
" AND n.nspname = current_schema() AND c.contype = 'f'"
),
{"name": constraint_name, "table": table_name},
)
return result.scalar() is not None
def upgrade() -> None:
if _fk_exists("request_candidates_key_id_fkey", "request_candidates"):
op.drop_constraint(
"request_candidates_key_id_fkey", "request_candidates", type_="foreignkey"
)
def downgrade() -> None:
bind = op.get_bind()
bind.execute(
sa.text(
"UPDATE request_candidates rc "
"SET key_id = NULL "
"WHERE key_id IS NOT NULL "
" AND NOT EXISTS ("
" SELECT 1 FROM provider_api_keys pak WHERE pak.id = rc.key_id"
" )"
)
)
if not _fk_exists("request_candidates_key_id_fkey", "request_candidates"):
op.create_foreign_key(
"request_candidates_key_id_fkey",
"request_candidates",
"provider_api_keys",
["key_id"],
["id"],
ondelete="CASCADE",
)
@@ -0,0 +1,54 @@
"""add user rate_limit and backfill normal api key limits
Revision ID: b7e8f9a0c1d2
Revises: b7c8d9e0f1a2
Create Date: 2026-03-13 12:00:00.000000+00:00
"""
from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from sqlalchemy import inspect
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "b7e8f9a0c1d2"
down_revision: str | None = "b7c8d9e0f1a2"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
columns = [c["name"] for c in inspector.get_columns(table_name)]
return column_name in columns
def upgrade() -> None:
if not column_exists("users", "rate_limit"):
op.add_column("users", sa.Column("rate_limit", sa.Integer(), nullable=True))
# 普通 Key 新语义不再允许 NULL;存量 NULL 统一回填为 0(不限制)。
op.execute(sa.text("""
UPDATE api_keys
SET rate_limit = 0
WHERE is_standalone = FALSE
AND rate_limit IS NULL
"""))
def downgrade() -> None:
# 恢复普通 Key 的 rate_limit 为 NULL(与 upgrade 中回填 0 对应)
op.execute(sa.text("""
UPDATE api_keys
SET rate_limit = NULL
WHERE is_standalone = FALSE
AND rate_limit = 0
"""))
if column_exists("users", "rate_limit"):
op.drop_column("users", "rate_limit")
@@ -0,0 +1,87 @@
"""add user sessions table for device-level auth
Revision ID: f6e7d8c9b0a1
Revises: b7e8f9a0c1d2
Create Date: 2026-03-15 12:00:00.000000+00:00
"""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision = "f6e7d8c9b0a1"
down_revision = "b7e8f9a0c1d2"
branch_labels = None
depends_on = None
def upgrade() -> None:
bind = op.get_bind()
inspector = sa.inspect(bind)
if "user_sessions" in inspector.get_table_names():
return
op.create_table(
"user_sessions",
sa.Column("id", sa.String(length=36), nullable=False),
sa.Column("user_id", sa.String(length=36), nullable=False),
sa.Column("client_device_id", sa.String(length=128), nullable=False),
sa.Column("device_label", sa.String(length=120), nullable=True),
sa.Column("device_type", sa.String(length=20), nullable=False, server_default="unknown"),
sa.Column("browser_name", sa.String(length=50), nullable=True),
sa.Column("browser_version", sa.String(length=50), nullable=True),
sa.Column("os_name", sa.String(length=50), nullable=True),
sa.Column("os_version", sa.String(length=50), nullable=True),
sa.Column("device_model", sa.String(length=100), nullable=True),
sa.Column("ip_address", sa.String(length=45), nullable=True),
sa.Column("user_agent", sa.String(length=1000), nullable=True),
sa.Column("client_hints", sa.JSON(), nullable=True),
sa.Column("refresh_token_hash", sa.String(length=64), nullable=False),
sa.Column("prev_refresh_token_hash", sa.String(length=64), nullable=True),
sa.Column("rotated_at", sa.DateTime(timezone=True), nullable=True),
sa.Column(
"last_seen_at", sa.DateTime(timezone=True), nullable=False, server_default=sa.func.now()
),
sa.Column("expires_at", sa.DateTime(timezone=True), nullable=False),
sa.Column("revoked_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("revoke_reason", sa.String(length=100), nullable=True),
sa.Column(
"created_at", sa.DateTime(timezone=True), nullable=False, server_default=sa.func.now()
),
sa.Column(
"updated_at", sa.DateTime(timezone=True), nullable=False, server_default=sa.func.now()
),
sa.ForeignKeyConstraint(["user_id"], ["users.id"], ondelete="CASCADE"),
sa.PrimaryKeyConstraint("id"),
)
op.create_index("ix_user_sessions_user_id", "user_sessions", ["user_id"], unique=False)
op.create_index(
"ix_user_sessions_client_device_id",
"user_sessions",
["client_device_id"],
unique=False,
)
op.create_index(
"idx_user_sessions_user_active",
"user_sessions",
["user_id", "revoked_at", "expires_at"],
unique=False,
)
op.create_index(
"idx_user_sessions_user_device",
"user_sessions",
["user_id", "client_device_id"],
unique=False,
)
def downgrade() -> None:
op.drop_index("idx_user_sessions_user_device", table_name="user_sessions")
op.drop_index("idx_user_sessions_user_active", table_name="user_sessions")
op.drop_index("ix_user_sessions_client_device_id", table_name="user_sessions")
op.drop_index("ix_user_sessions_user_id", table_name="user_sessions")
op.drop_table("user_sessions")
@@ -0,0 +1,41 @@
"""add status_snapshot column to provider_api_keys
Revision ID: c9d8e7f6a5b4
Revises: f6e7d8c9b0a1
Create Date: 2026-03-20 12:00:00.000000
"""
from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from sqlalchemy import inspect
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "c9d8e7f6a5b4"
down_revision: str | None = "f6e7d8c9b0a1"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
columns = [c["name"] for c in inspector.get_columns(table_name)]
return column_name in columns
def upgrade() -> None:
if not column_exists("provider_api_keys", "status_snapshot"):
op.add_column(
"provider_api_keys",
sa.Column("status_snapshot", sa.JSON(), nullable=True),
)
def downgrade() -> None:
if column_exists("provider_api_keys", "status_snapshot"):
op.drop_column("provider_api_keys", "status_snapshot")
@@ -0,0 +1,317 @@
"""usage token semantics v2
Revision ID: c3d4e5f6a7b8
Revises: c9d8e7f6a5b4
Create Date: 2026-03-24 14:00:00.000000+00:00
"""
from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from sqlalchemy import inspect
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "c3d4e5f6a7b8"
down_revision: str | None = "c9d8e7f6a5b4"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
BACKFILL_BATCH_SIZE = 2000
# 最大批次数,防止因数据异常导致死循环(2000 * 500000 = 10亿行上限)
_MAX_BATCHES = 500000
# 使用子查询中间层展开 input_output_total_tokens 的计算,
# 确保 total_tokens 引用的是本次 SET 后的新值而非旧值。
_UPGRADE_BACKFILL_SQL = sa.text(
"""
UPDATE usage
SET
input_output_total_tokens = src.new_iot,
input_context_tokens = src.new_ict,
total_tokens = src.new_total,
cache_creation_cost_usd_5m = src.new_cc5m,
cache_creation_cost_usd_1h = src.new_cc1h,
actual_cache_creation_cost_usd_5m = src.new_acc5m,
actual_cache_creation_cost_usd_1h = src.new_acc1h,
actual_cache_cost_usd = src.new_accu,
cache_creation_price_per_1m_5m = src.new_cp5m,
cache_creation_price_per_1m_1h = src.new_cp1h,
cache_cost_usd = src.new_ccu
FROM (
SELECT
id,
COALESCE(input_output_total_tokens,
COALESCE(input_tokens, 0) + COALESCE(output_tokens, 0))
AS new_iot,
COALESCE(input_tokens, 0) + COALESCE(cache_read_input_tokens, 0)
AS new_ict,
/* total_tokens 引用本行计算出的 new_iot,避免依赖 SET 顺序 */
COALESCE(input_output_total_tokens,
COALESCE(input_tokens, 0) + COALESCE(output_tokens, 0))
+ COALESCE(cache_creation_input_tokens, 0)
+ COALESCE(cache_read_input_tokens, 0)
AS new_total,
CASE
WHEN COALESCE(cache_creation_input_tokens_5m, 0) > 0
AND COALESCE(cache_creation_input_tokens_1h, 0) = 0
THEN COALESCE(cache_creation_cost_usd, 0)
WHEN COALESCE(cache_creation_input_tokens_5m, 0) > 0
AND COALESCE(cache_creation_input_tokens, 0) > 0
THEN COALESCE(cache_creation_cost_usd, 0)
* (COALESCE(cache_creation_input_tokens_5m, 0) * 1.0
/ GREATEST(COALESCE(cache_creation_input_tokens, 0), 1))
ELSE 0
END AS new_cc5m,
CASE
WHEN COALESCE(cache_creation_input_tokens_1h, 0) > 0
AND COALESCE(cache_creation_input_tokens_5m, 0) = 0
THEN COALESCE(cache_creation_cost_usd, 0)
WHEN COALESCE(cache_creation_input_tokens_1h, 0) > 0
AND COALESCE(cache_creation_input_tokens, 0) > 0
THEN COALESCE(cache_creation_cost_usd, 0)
* (COALESCE(cache_creation_input_tokens_1h, 0) * 1.0
/ GREATEST(COALESCE(cache_creation_input_tokens, 0), 1))
ELSE 0
END AS new_cc1h,
CASE
WHEN COALESCE(cache_creation_input_tokens_5m, 0) > 0
AND COALESCE(cache_creation_input_tokens_1h, 0) = 0
THEN COALESCE(actual_cache_creation_cost_usd, 0)
WHEN COALESCE(cache_creation_input_tokens_5m, 0) > 0
AND COALESCE(cache_creation_input_tokens, 0) > 0
THEN COALESCE(actual_cache_creation_cost_usd, 0)
* (COALESCE(cache_creation_input_tokens_5m, 0) * 1.0
/ GREATEST(COALESCE(cache_creation_input_tokens, 0), 1))
ELSE 0
END AS new_acc5m,
CASE
WHEN COALESCE(cache_creation_input_tokens_1h, 0) > 0
AND COALESCE(cache_creation_input_tokens_5m, 0) = 0
THEN COALESCE(actual_cache_creation_cost_usd, 0)
WHEN COALESCE(cache_creation_input_tokens_1h, 0) > 0
AND COALESCE(cache_creation_input_tokens, 0) > 0
THEN COALESCE(actual_cache_creation_cost_usd, 0)
* (COALESCE(cache_creation_input_tokens_1h, 0) * 1.0
/ GREATEST(COALESCE(cache_creation_input_tokens, 0), 1))
ELSE 0
END AS new_acc1h,
COALESCE(actual_cache_creation_cost_usd, 0)
+ COALESCE(actual_cache_read_cost_usd, 0) AS new_accu,
CASE
WHEN COALESCE(cache_creation_input_tokens_5m, 0) > 0
AND COALESCE(cache_creation_input_tokens_1h, 0) = 0
THEN cache_creation_price_per_1m
ELSE NULL
END AS new_cp5m,
CASE
WHEN COALESCE(cache_creation_input_tokens_1h, 0) > 0
AND COALESCE(cache_creation_input_tokens_5m, 0) = 0
THEN cache_creation_price_per_1m
ELSE NULL
END AS new_cp1h,
COALESCE(cache_creation_cost_usd, 0)
+ COALESCE(cache_read_cost_usd, 0) AS new_ccu
FROM usage
WHERE id IN (
SELECT id FROM usage
WHERE
input_output_total_tokens IS DISTINCT FROM
COALESCE(input_output_total_tokens,
COALESCE(input_tokens, 0) + COALESCE(output_tokens, 0))
OR input_context_tokens IS DISTINCT FROM
COALESCE(input_tokens, 0) + COALESCE(cache_read_input_tokens, 0)
OR total_tokens IS DISTINCT FROM (
COALESCE(input_output_total_tokens,
COALESCE(input_tokens, 0) + COALESCE(output_tokens, 0))
+ COALESCE(cache_creation_input_tokens, 0)
+ COALESCE(cache_read_input_tokens, 0)
)
OR cache_creation_cost_usd_5m IS DISTINCT FROM (
CASE
WHEN COALESCE(cache_creation_input_tokens_5m, 0) > 0
AND COALESCE(cache_creation_input_tokens_1h, 0) = 0
THEN COALESCE(cache_creation_cost_usd, 0)
WHEN COALESCE(cache_creation_input_tokens_5m, 0) > 0
AND COALESCE(cache_creation_input_tokens, 0) > 0
THEN COALESCE(cache_creation_cost_usd, 0)
* (COALESCE(cache_creation_input_tokens_5m, 0) * 1.0
/ GREATEST(COALESCE(cache_creation_input_tokens, 0), 1))
ELSE 0
END
)
OR cache_creation_cost_usd_1h IS DISTINCT FROM (
CASE
WHEN COALESCE(cache_creation_input_tokens_1h, 0) > 0
AND COALESCE(cache_creation_input_tokens_5m, 0) = 0
THEN COALESCE(cache_creation_cost_usd, 0)
WHEN COALESCE(cache_creation_input_tokens_1h, 0) > 0
AND COALESCE(cache_creation_input_tokens, 0) > 0
THEN COALESCE(cache_creation_cost_usd, 0)
* (COALESCE(cache_creation_input_tokens_1h, 0) * 1.0
/ GREATEST(COALESCE(cache_creation_input_tokens, 0), 1))
ELSE 0
END
)
OR actual_cache_creation_cost_usd_5m IS DISTINCT FROM (
CASE
WHEN COALESCE(cache_creation_input_tokens_5m, 0) > 0
AND COALESCE(cache_creation_input_tokens_1h, 0) = 0
THEN COALESCE(actual_cache_creation_cost_usd, 0)
WHEN COALESCE(cache_creation_input_tokens_5m, 0) > 0
AND COALESCE(cache_creation_input_tokens, 0) > 0
THEN COALESCE(actual_cache_creation_cost_usd, 0)
* (COALESCE(cache_creation_input_tokens_5m, 0) * 1.0
/ GREATEST(COALESCE(cache_creation_input_tokens, 0), 1))
ELSE 0
END
)
OR actual_cache_creation_cost_usd_1h IS DISTINCT FROM (
CASE
WHEN COALESCE(cache_creation_input_tokens_1h, 0) > 0
AND COALESCE(cache_creation_input_tokens_5m, 0) = 0
THEN COALESCE(actual_cache_creation_cost_usd, 0)
WHEN COALESCE(cache_creation_input_tokens_1h, 0) > 0
AND COALESCE(cache_creation_input_tokens, 0) > 0
THEN COALESCE(actual_cache_creation_cost_usd, 0)
* (COALESCE(cache_creation_input_tokens_1h, 0) * 1.0
/ GREATEST(COALESCE(cache_creation_input_tokens, 0), 1))
ELSE 0
END
)
OR actual_cache_cost_usd IS DISTINCT FROM (
COALESCE(actual_cache_creation_cost_usd, 0)
+ COALESCE(actual_cache_read_cost_usd, 0)
)
OR cache_creation_price_per_1m_5m IS DISTINCT FROM (
CASE
WHEN COALESCE(cache_creation_input_tokens_5m, 0) > 0
AND COALESCE(cache_creation_input_tokens_1h, 0) = 0
THEN cache_creation_price_per_1m
ELSE NULL
END
)
OR cache_creation_price_per_1m_1h IS DISTINCT FROM (
CASE
WHEN COALESCE(cache_creation_input_tokens_1h, 0) > 0
AND COALESCE(cache_creation_input_tokens_5m, 0) = 0
THEN cache_creation_price_per_1m
ELSE NULL
END
)
OR cache_cost_usd IS DISTINCT FROM (
COALESCE(cache_creation_cost_usd, 0) + COALESCE(cache_read_cost_usd, 0)
)
ORDER BY id
LIMIT :batch_size
)
) AS src
WHERE usage.id = src.id
"""
)
_DOWNGRADE_BACKFILL_SQL = sa.text(
"""
UPDATE usage
SET total_tokens = COALESCE(input_output_total_tokens, COALESCE(input_tokens, 0) + COALESCE(output_tokens, 0))
WHERE id IN (
SELECT id
FROM usage
WHERE total_tokens IS DISTINCT FROM
COALESCE(input_output_total_tokens, COALESCE(input_tokens, 0) + COALESCE(output_tokens, 0))
ORDER BY id
LIMIT :batch_size
)
"""
)
def column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
columns = [c["name"] for c in inspector.get_columns(table_name)]
return column_name in columns
def run_backfill_in_batches(sql: sa.TextClause, batch_size: int = BACKFILL_BATCH_SIZE) -> None:
context = op.get_context()
for _ in range(_MAX_BATCHES):
# Commit the preceding schema transaction before each batch so PostgreSQL
# does not keep ALTER TABLE locks for the entire data backfill.
with context.autocommit_block():
rowcount = op.get_bind().execute(sql, {"batch_size": batch_size}).rowcount
if rowcount == 0:
break
else:
raise RuntimeError(
f"Backfill did not converge after {_MAX_BATCHES} batches "
f"(batch_size={batch_size}). Possible infinite loop due to data anomaly."
)
def upgrade() -> None:
if column_exists("usage", "total_tokens") and not column_exists("usage", "input_output_total_tokens"):
with op.batch_alter_table("usage") as batch_op:
batch_op.alter_column(
"total_tokens",
new_column_name="input_output_total_tokens",
existing_type=sa.Integer(),
existing_nullable=True,
)
with op.batch_alter_table("usage") as batch_op:
if not column_exists("usage", "input_context_tokens"):
batch_op.add_column(sa.Column("input_context_tokens", sa.Integer(), nullable=False, server_default="0"))
if not column_exists("usage", "total_tokens"):
batch_op.add_column(sa.Column("total_tokens", sa.Integer(), nullable=False, server_default="0"))
if not column_exists("usage", "cache_creation_cost_usd_5m"):
batch_op.add_column(sa.Column("cache_creation_cost_usd_5m", sa.Numeric(20, 8), nullable=False, server_default="0"))
if not column_exists("usage", "cache_creation_cost_usd_1h"):
batch_op.add_column(sa.Column("cache_creation_cost_usd_1h", sa.Numeric(20, 8), nullable=False, server_default="0"))
if not column_exists("usage", "actual_cache_creation_cost_usd_5m"):
batch_op.add_column(sa.Column("actual_cache_creation_cost_usd_5m", sa.Numeric(20, 8), nullable=False, server_default="0"))
if not column_exists("usage", "actual_cache_creation_cost_usd_1h"):
batch_op.add_column(sa.Column("actual_cache_creation_cost_usd_1h", sa.Numeric(20, 8), nullable=False, server_default="0"))
if not column_exists("usage", "actual_cache_cost_usd"):
batch_op.add_column(sa.Column("actual_cache_cost_usd", sa.Numeric(20, 8), nullable=False, server_default="0"))
if not column_exists("usage", "cache_creation_price_per_1m_5m"):
batch_op.add_column(sa.Column("cache_creation_price_per_1m_5m", sa.Numeric(20, 8), nullable=True))
if not column_exists("usage", "cache_creation_price_per_1m_1h"):
batch_op.add_column(sa.Column("cache_creation_price_per_1m_1h", sa.Numeric(20, 8), nullable=True))
run_backfill_in_batches(_UPGRADE_BACKFILL_SQL)
def downgrade() -> None:
run_backfill_in_batches(_DOWNGRADE_BACKFILL_SQL)
with op.batch_alter_table("usage") as batch_op:
if column_exists("usage", "cache_creation_price_per_1m_1h"):
batch_op.drop_column("cache_creation_price_per_1m_1h")
if column_exists("usage", "cache_creation_price_per_1m_5m"):
batch_op.drop_column("cache_creation_price_per_1m_5m")
if column_exists("usage", "actual_cache_cost_usd"):
batch_op.drop_column("actual_cache_cost_usd")
if column_exists("usage", "actual_cache_creation_cost_usd_1h"):
batch_op.drop_column("actual_cache_creation_cost_usd_1h")
if column_exists("usage", "actual_cache_creation_cost_usd_5m"):
batch_op.drop_column("actual_cache_creation_cost_usd_5m")
if column_exists("usage", "cache_creation_cost_usd_1h"):
batch_op.drop_column("cache_creation_cost_usd_1h")
if column_exists("usage", "cache_creation_cost_usd_5m"):
batch_op.drop_column("cache_creation_cost_usd_5m")
if column_exists("usage", "input_context_tokens"):
batch_op.drop_column("input_context_tokens")
if column_exists("usage", "total_tokens"):
batch_op.drop_column("total_tokens")
if column_exists("usage", "input_output_total_tokens"):
batch_op.alter_column(
"input_output_total_tokens",
new_column_name="total_tokens",
existing_type=sa.Integer(),
existing_nullable=True,
)
+1 -1
View File
@@ -15,7 +15,7 @@
### 用户系统
- **users**: 用户账户管理
- **api_keys**: API 密钥管理
- **user_quotas**: 用户配额管理
- **wallets**: 统一钱包账户(充值余额/赠款余额/无限制模式)
- **user_preferences**: 用户偏好设置
### Provider 三层架构
+218 -28
View File
@@ -3,45 +3,131 @@
#
# 用法:
# 部署/更新: ./deploy.sh (自动检测所有变化)
# 指定 Hub 版本: ./deploy.sh --hub-tag hub-v0.1.0
# 更新 Hub: ./deploy.sh --update-hub
# GitHub 镜像: ./deploy.sh --mirror https://ghfast.top
# 强制重建: ./deploy.sh --rebuild-base
# 强制全部重建: ./deploy.sh --force
set -e
set -euo pipefail
cd "$(dirname "$0")"
# 兼容 docker-compose 和 docker compose
if command -v docker-compose &> /dev/null; then
DC="docker-compose -f docker-compose.build.yml"
USE_LEGACY_COMPOSE=true
else
DC="docker compose -f docker-compose.build.yml"
USE_LEGACY_COMPOSE=false
fi
compose_up() {
if [ "$USE_LEGACY_COMPOSE" = true ]; then
$DC up -d --no-build "$@"
else
$DC up -d --no-build --pull never "$@"
fi
}
# 缓存文件
HASH_FILE=".deps-hash"
CODE_HASH_FILE=".code-hash"
MIGRATION_HASH_FILE=".migration-hash"
# 提取 pyproject.toml 中"会影响运行时依赖安装"的最小指纹(与 CI 保持一致):
# - [build-system] requires / build-backend
# - [project] requires-python / dependencies
# 使用 Python tomllib 解析,不受 TOML 格式变化影响。
pyproject_deps_fingerprint() {
python3 - <<'PY'
import json, pathlib, tomllib
# Hub release 配置
GITHUB_REPO="fawney19/Aether"
HUB_TAG_STATE_FILE=".hub-tag"
data = tomllib.loads(pathlib.Path("pyproject.toml").read_text("utf-8"))
project = data.get("project") or {}
build = data.get("build-system") or {}
usage() {
cat <<'EOF'
Usage: ./deploy.sh [options]
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 []),
Options:
--hub-tag <hub-vX.Y.Z> 指定 Hub Release tag(例如 hub-v0.1.0)
--update-hub 强制刷新 Hub 版本标记(下次构建会重新下载)
--mirror <url> GitHub 下载镜像(例如 https://ghfast.top)
--rebuild-base, -r 仅重建 base 镜像
--force, -f 强制重建全部(hub/base/app)并重启
-h, --help 显示帮助
EOF
}
print(json.dumps(fingerprint, sort_keys=True, separators=(",", ":")))
PY
FORCE_REBUILD_ALL=false
REBUILD_BASE_ONLY=false
FORCE_UPDATE_HUB=false
HUB_TAG="${HUB_TAG:-}"
GITHUB_MIRROR="${GITHUB_MIRROR:-}"
RESOLVED_HUB_TAG=""
while [ $# -gt 0 ]; do
case "$1" in
--hub-tag)
if [ $# -lt 2 ]; then
echo "❌ --hub-tag 需要一个值,例如 hub-v0.1.0"
exit 1
fi
HUB_TAG="$2"
shift 2
;;
--update-hub)
FORCE_UPDATE_HUB=true
shift
;;
--mirror)
if [ $# -lt 2 ]; then
echo "ERROR: --mirror needs a URL, e.g. https://ghfast.top"
exit 1
fi
GITHUB_MIRROR="$2"
shift 2
;;
--rebuild-base|-r)
REBUILD_BASE_ONLY=true
shift
;;
--force|-f)
FORCE_REBUILD_ALL=true
shift
;;
-h|--help)
usage
exit 0
;;
*)
echo "❌ 未知参数: $1"
usage
exit 1
;;
esac
done
if [ -n "$HUB_TAG" ]; then
case "$HUB_TAG" in
hub-v*) ;;
*) echo "❌ --hub-tag 格式应为 hub-vX.Y.Z,例如 hub-v0.1.0"; exit 1 ;;
esac
fi
# 提取 pyproject.toml 中会影响运行时依赖安装的字段指纹(纯 shell,无需 Python)
# 用 sed 提取 dependencies / requires 数组块和单值字段,排序后输出稳定文本
pyproject_deps_fingerprint() {
local file="pyproject.toml"
# 提取 "key = [..." 多行数组块(从 key 行到 ] 行)
extract_array() {
sed -n "/^$1[[:space:]]*=[[:space:]]*\[/,/\]/p" "$file" | grep '"' | sed 's/.*"\(.*\)".*/\1/' | sort
}
# 提取 "key = "value"" 单行值
extract_value() {
grep -m1 "^$1[[:space:]]*=" "$file" 2>/dev/null | sed 's/.*"\(.*\)".*/\1/'
}
{
echo "requires-python=$(extract_value requires-python)"
echo "build-backend=$(extract_value build-backend)"
echo "dependencies:"
extract_array dependencies
echo "build-requires:"
extract_array requires
}
}
# 计算依赖文件的哈希值(包含 Dockerfile.base.local)
@@ -63,6 +149,66 @@ calc_code_hash() {
} | md5sum | cut -d' ' -f1
}
# 获取最新 hub release tag
# 支持 GITHUB_TOKEN 环境变量以避免未认证 API 限流(60 次/小时 -> 5000 次/小时)
get_latest_hub_tag() {
local auth_args=()
if [ -n "${GITHUB_TOKEN:-}" ]; then
auth_args=(-H "Authorization: token ${GITHUB_TOKEN}")
fi
curl -sL "${auth_args[@]}" "https://api.github.com/repos/$GITHUB_REPO/releases" | \
python3 -c "
import json, sys
releases = json.load(sys.stdin)
for r in releases:
tag = r.get('tag_name', '')
if tag.startswith('hub-v') and not r.get('draft') and not r.get('prerelease'):
print(tag)
break
" 2>/dev/null
}
# 解析当前应使用的 Hub release tag(优先使用指定值,否则拉取最新)
resolve_hub_tag() {
local requested_tag="${1:-}"
local latest_tag
if [ -n "$requested_tag" ]; then
echo "$requested_tag"
return 0
fi
latest_tag="$(get_latest_hub_tag || true)"
if [ -n "$latest_tag" ]; then
echo "$latest_tag"
return 0
fi
if [ -f "$HUB_TAG_STATE_FILE" ]; then
echo "⚠️ 无法查询最新 Hub 版本,回退使用本地记录: $(cat "$HUB_TAG_STATE_FILE")" >&2
cat "$HUB_TAG_STATE_FILE"
return 0
fi
echo "❌ 无法获取 Hub Release tag,请检查网络或手动指定 --hub-tag" >&2
exit 1
}
# 确保本次构建的 Hub tag 已解析(默认追踪最新 release,也可通过 --hub-tag 固定版本)
ensure_hub_tag() {
local requested_tag="${1:-}"
RESOLVED_HUB_TAG="$(resolve_hub_tag "$requested_tag")"
if [ -f "$HUB_TAG_STATE_FILE" ] && [ "$(cat "$HUB_TAG_STATE_FILE")" = "$RESOLVED_HUB_TAG" ]; then
echo ">>> Hub 版本未变化: $RESOLVED_HUB_TAG"
return 1
fi
echo "$RESOLVED_HUB_TAG" > "$HUB_TAG_STATE_FILE"
echo ">>> 使用 Hub 版本: $RESOLVED_HUB_TAG"
return 0
}
# 计算迁移文件的哈希值
calc_migration_hash() {
find alembic/versions -name "*.py" -type f 2>/dev/null | sort | xargs cat 2>/dev/null | md5sum | cut -d' ' -f1
@@ -92,6 +238,8 @@ check_code_changed() {
return 0
}
# 检查迁移是否变化
check_migration_changed() {
local current_hash=$(calc_migration_hash)
@@ -112,10 +260,11 @@ save_migration_hash() { calc_migration_hash > "$MIGRATION_HASH_FILE"; }
# 构建基础镜像
build_base() {
echo ">>> Building base image (dependencies)..."
docker build -f Dockerfile.base.local -t aether-base:latest .
docker build --pull=false -f Dockerfile.base.local -t aether-base:latest .
save_deps_hash
}
# 生成版本文件
generate_version_file() {
# 从 git 获取版本号
@@ -137,8 +286,27 @@ EOF
# 构建应用镜像
build_app() {
echo ">>> Building app image (code only)..."
if [ -z "${RESOLVED_HUB_TAG:-}" ]; then
echo ">>> RESOLVED_HUB_TAG 为空,无法构建 app 镜像"
exit 1
fi
echo ">>> Build args: HUB_TAG=$RESOLVED_HUB_TAG"
generate_version_file
docker build -f Dockerfile.app.local -t aether-app:latest .
local token_args=()
if [ -n "${GITHUB_TOKEN:-}" ]; then
token_args=(--build-arg "GITHUB_TOKEN=${GITHUB_TOKEN}")
fi
local mirror_args=()
if [ -n "${GITHUB_MIRROR:-}" ]; then
mirror_args=(--build-arg "GITHUB_MIRROR=${GITHUB_MIRROR}")
fi
docker build --pull=false \
--build-arg HUB_RELEASE_REPO="$GITHUB_REPO" \
--build-arg HUB_TAG="$RESOLVED_HUB_TAG" \
"${token_args[@]}" \
"${mirror_args[@]}" \
-f Dockerfile.app.local \
-t aether-app:latest .
save_code_hash
}
@@ -186,11 +354,15 @@ print('Old version cleared')
}
# 强制全部重建
if [ "$1" = "--force" ] || [ "$1" = "-f" ]; then
if [ "$FORCE_REBUILD_ALL" = true ]; then
echo ">>> Force rebuilding everything..."
if [ "$FORCE_UPDATE_HUB" = true ]; then
rm -f "$HUB_TAG_STATE_FILE"
fi
ensure_hub_tag "$HUB_TAG" || true
build_base
build_app
$DC up -d --force-recreate
compose_up --force-recreate
sleep 3
run_migration
docker image prune -f
@@ -200,19 +372,25 @@ if [ "$1" = "--force" ] || [ "$1" = "-f" ]; then
fi
# 强制重建基础镜像
if [ "$1" = "--rebuild-base" ] || [ "$1" = "-r" ]; then
if [ "$REBUILD_BASE_ONLY" = true ]; then
build_base
echo ">>> Base image rebuilt. Run ./deploy.sh to deploy."
exit 0
fi
# 拉取最新代码
echo ">>> Pulling latest code..."
git pull
# 更新 Hub 版本标记
if [ "$FORCE_UPDATE_HUB" = true ]; then
rm -f "$HUB_TAG_STATE_FILE"
ensure_hub_tag "$HUB_TAG" || true
echo ">>> Hub tag updated: $RESOLVED_HUB_TAG"
echo ">>> Run ./deploy.sh to build app image with the new Hub release."
exit 0
fi
# 标记是否需要重启
NEED_RESTART=false
BASE_REBUILT=false
HUB_UPDATED=false
# 检查基础镜像是否存在,或依赖是否变化
if ! docker image inspect aether-base:latest >/dev/null 2>&1; then
@@ -229,6 +407,14 @@ else
echo ">>> Dependencies unchanged."
fi
# 解析/检查 Hub 版本(构建时由 Dockerfile 从 GitHub Release 下载)
if ensure_hub_tag "$HUB_TAG"; then
HUB_UPDATED=true
NEED_RESTART=true
else
echo ">>> Hub version unchanged."
fi
# 检查代码或迁移是否变化,或者 base 重建了(app 依赖 base)
# 注意:迁移文件打包在镜像中,所以迁移变化也需要重建 app 镜像
MIGRATION_CHANGED=false
@@ -244,6 +430,10 @@ elif [ "$BASE_REBUILT" = true ]; then
echo ">>> Base image rebuilt, rebuilding app image..."
build_app
NEED_RESTART=true
elif [ "$HUB_UPDATED" = true ]; then
echo ">>> Hub version updated, rebuilding app image..."
build_app
NEED_RESTART=true
elif check_code_changed; then
echo ">>> Code changed, rebuilding app image..."
build_app
@@ -265,10 +455,10 @@ fi
# 有变化时重启,或容器未运行时启动
if [ "$NEED_RESTART" = true ]; then
echo ">>> Restarting services..."
$DC up -d
compose_up
elif [ "$CONTAINERS_RUNNING" = false ]; then
echo ">>> Containers not running, starting services..."
$DC up -d
compose_up
else
echo ">>> No changes detected, skipping restart."
fi
+11 -5
View File
@@ -11,10 +11,16 @@ set +a
export DATABASE_URL="postgresql://${DB_USER:-postgres}:${DB_PASSWORD}@${DB_HOST:-localhost}:${DB_PORT:-5432}/${DB_NAME:-aether}"
export REDIS_URL=redis://:${REDIS_PASSWORD}@${REDIS_HOST:-localhost}:${REDIS_PORT:-6379}/0
# 启动 uvicorn(热重载模式)
echo "🚀 启动本地开发服务器..."
echo "📍 后端地址: http://localhost:8084"
echo "📊 数据库: ${DATABASE_URL}"
# 开发环境连接池低配(节省内存)
export DB_POOL_SIZE=${DB_POOL_SIZE:-5}
export DB_MAX_OVERFLOW=${DB_MAX_OVERFLOW:-5}
export HTTP_MAX_CONNECTIONS=${HTTP_MAX_CONNECTIONS:-20}
export HTTP_KEEPALIVE_CONNECTIONS=${HTTP_KEEPALIVE_CONNECTIONS:-5}
# 启动 uvicorn(热重载模式,只监视 src 目录)
echo "=> 启动本地开发服务器..."
echo "=> 后端地址: http://localhost:8084"
echo "=> 数据库: ${DATABASE_URL}"
echo ""
uv run uvicorn src.main:app --reload --port 8084
uv run uvicorn src.main:app --reload --reload-dir src --port 8084
+6 -4
View File
@@ -1,7 +1,10 @@
# Aether 部署配置 - 本地构建
# 使用方法:
# 首次构建 base: docker build -f Dockerfile.base -t aether-base:latest .
# Hub 二进制: cd aether-hub && ./build.sh(默认 binary 模式)供发布;deploy.sh 构建 app 时会自动下载
# Hub 镜像发布: cd aether-hub && ./build.sh --image --tag <tag> --push(可选)
# 启动服务: docker compose -f docker-compose.build.yml up -d --build
# 或使用: ./deploy.sh
services:
postgres:
@@ -26,9 +29,7 @@ services:
redis:
image: redis:7-alpine
container_name: aether-redis
command: redis-server --appendonly yes --requirepass ${REDIS_PASSWORD}
volumes:
- redis_data:/data
command: redis-server --appendonly no --save "" --requirepass ${REDIS_PASSWORD}
ports:
- "${REDIS_PORT:-6379}:6379"
healthcheck:
@@ -42,6 +43,8 @@ services:
build:
context: .
dockerfile: Dockerfile.app.local
args:
GITHUB_TOKEN: ${GITHUB_TOKEN:-}
image: aether-app:latest
container_name: aether-app
env_file:
@@ -67,4 +70,3 @@ services:
volumes:
postgres_data:
redis_data:
+4 -5
View File
@@ -22,9 +22,7 @@ services:
redis:
image: redis:7-alpine
container_name: aether-redis
command: redis-server --appendonly yes --requirepass ${REDIS_PASSWORD}
volumes:
- redis_data:/data
command: redis-server --appendonly no --save "" --requirepass ${REDIS_PASSWORD}
healthcheck:
test: [ "CMD", "redis-cli", "--raw", "incr", "ping" ]
interval: 5s
@@ -33,7 +31,7 @@ services:
restart: unless-stopped
app:
image: ghcr.io/fawney19/aether:latest
image: ${APP_IMAGE:-ghcr.io/fawney19/aether:latest}
container_name: aether-app
env_file:
- .env
@@ -58,4 +56,5 @@ services:
volumes:
postgres_data:
redis_data:
+117
View File
@@ -0,0 +1,117 @@
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 760 420" width="760" height="420" font-family="'Inter', system-ui, -apple-system, sans-serif">
<defs>
<filter id="shadow-sm" x="-10%" y="-10%" width="120%" height="130%">
<feDropShadow dx="0" dy="2" stdDeviation="2" flood-opacity="0.5"/>
</filter>
<filter id="shadow-lg" x="-10%" y="-10%" width="120%" height="130%">
<feDropShadow dx="0" dy="8" stdDeviation="12" flood-opacity="0.7"/>
</filter>
<clipPath id="gateway-clip">
<path d="M 0 0 h 760 v 420 H 0 Z M 74 169 v 54 h 12 v -54 Z" clip-rule="evenodd"/>
</clipPath>
</defs>
<!-- Background -->
<rect width="760" height="420" rx="24" fill="#121212"/>
<!-- === Labels placeholder, drawn at top layer === -->
<!-- Source pills -->
<rect x="138" y="24" width="100" height="36" rx="18" fill="#1a1a1c" stroke="#2a2a2c" filter="url(#shadow-sm)"/>
<circle cx="163" cy="42" r="3" fill="#cc7154"/>
<text x="195" y="46" font-size="11" font-weight="700" fill="#f3f4f6" text-anchor="middle">Claude</text>
<rect x="330" y="24" width="100" height="36" rx="18" fill="#1a1a1c" stroke="#2a2a2c" filter="url(#shadow-sm)"/>
<circle cx="355" cy="42" r="3" fill="#cc7154"/>
<text x="387" y="46" font-size="11" font-weight="700" fill="#f3f4f6" text-anchor="middle">OpenAI</text>
<rect x="522" y="24" width="100" height="36" rx="18" fill="#1a1a1c" stroke="#2a2a2c" filter="url(#shadow-sm)"/>
<circle cx="547" cy="42" r="3" fill="#cc7154"/>
<text x="579" y="46" font-size="11" font-weight="700" fill="#f3f4f6" text-anchor="middle">Gemini</text>
<!-- Output pills -->
<rect x="128" y="360" width="120" height="36" rx="18" fill="#1a1a1c" stroke="#2a2a2c" filter="url(#shadow-sm)"/>
<circle cx="153" cy="378" r="3" fill="#cc7154"/>
<text x="195" y="382" font-size="10" font-weight="700" fill="#f3f4f6" text-anchor="middle">Claude 响应</text>
<rect x="320" y="360" width="120" height="36" rx="18" fill="#1a1a1c" stroke="#2a2a2c" filter="url(#shadow-sm)"/>
<circle cx="345" cy="378" r="3" fill="#cc7154"/>
<text x="387" y="382" font-size="10" font-weight="700" fill="#f3f4f6" text-anchor="middle">OpenAI 响应</text>
<rect x="512" y="360" width="120" height="36" rx="18" fill="#1a1a1c" stroke="#2a2a2c" filter="url(#shadow-sm)"/>
<circle cx="537" cy="378" r="3" fill="#cc7154"/>
<text x="579" y="382" font-size="10" font-weight="700" fill="#f3f4f6" text-anchor="middle">Gemini 响应</text>
<!-- === Layer 2: AETHER GATEWAY box === -->
<rect x="80" y="78" width="600" height="240" rx="24" fill="#1a1a1c" filter="url(#shadow-lg)"/>
<rect x="80" y="78" width="600" height="240" rx="24" fill="none" stroke="rgba(204,113,84,0.4)" clip-path="url(#gateway-clip)"/>
<g transform="translate(80, 176)" font-size="7" font-weight="900" fill="#cc7154" text-anchor="middle">
<text x="0" y="0">A</text><text x="0" y="9">E</text><text x="0" y="18">T</text>
<text x="0" y="27">H</text><text x="0" y="36">E</text><text x="0" y="45">R</text>
</g>
<!-- Dashed box -->
<rect x="180" y="98" width="400" height="112" rx="16" fill="none" stroke="rgba(204,113,84,0.4)" stroke-width="1.5" stroke-dasharray="6 4"/>
<rect x="200" y="108" width="360" height="26" rx="12" fill="none" stroke="#2a2a2c"/>
<text x="380" y="125" font-size="10" font-weight="600" fill="#f3f4f6" text-anchor="middle">统一模型规范 / 协议聚合</text>
<rect x="200" y="140" width="85" height="26" rx="13" fill="rgba(204,113,84,0.12)"/>
<text x="242" y="157" font-size="10" font-weight="600" fill="#f3f4f6" text-anchor="middle">多端鉴权</text>
<rect x="290" y="140" width="85" height="26" rx="13" fill="rgba(204,113,84,0.12)"/>
<text x="332" y="157" font-size="10" font-weight="600" fill="#f3f4f6" text-anchor="middle">配额管控</text>
<rect x="380" y="140" width="85" height="26" rx="13" fill="rgba(204,113,84,0.12)"/>
<text x="422" y="157" font-size="10" font-weight="600" fill="#f3f4f6" text-anchor="middle">全局并发</text>
<rect x="470" y="140" width="85" height="26" rx="13" fill="rgba(204,113,84,0.12)"/>
<text x="512" y="157" font-size="10" font-weight="600" fill="#f3f4f6" text-anchor="middle">缓存亲和</text>
<rect x="200" y="172" width="360" height="28" rx="14" fill="none" stroke="#2a2a2c"/>
<text x="380" y="190" font-size="10" font-weight="600" fill="#f3f4f6" text-anchor="middle">智能调度 / 故障转移</text>
<!-- Engine cards -->
<rect x="103" y="250" width="170" height="28" rx="14" fill="rgba(204,113,84,0.12)" stroke="#2a2a2c"/>
<circle cx="163" cy="264" r="3" fill="#cc7154"/>
<text x="192" y="268" font-size="10" font-weight="700" fill="#f3f4f6" text-anchor="middle">格式转换</text>
<rect x="295" y="250" width="170" height="28" rx="14" fill="rgba(204,113,84,0.12)" stroke="#2a2a2c"/>
<circle cx="355" cy="264" r="3" fill="#cc7154"/>
<text x="384" y="268" font-size="10" font-weight="700" fill="#f3f4f6" text-anchor="middle">反向代理</text>
<rect x="487" y="250" width="170" height="28" rx="14" fill="rgba(204,113,84,0.12)" stroke="#2a2a2c"/>
<circle cx="547" cy="264" r="3" fill="#cc7154"/>
<text x="576" y="268" font-size="10" font-weight="700" fill="#f3f4f6" text-anchor="middle">原生直通</text>
<!-- === Layer 3: Track lines === -->
<g stroke="rgba(255,255,255,0.08)" stroke-width="1.5" fill="none" stroke-linecap="round">
<path id="f-in-1" d="M 188 60 L 188 66 C 188 72, 380 72, 380 78 L 380 98"/>
<path id="f-in-2" d="M 380 60 L 380 98"/>
<path id="f-in-3" d="M 572 60 L 572 66 C 572 72, 380 72, 380 78 L 380 98"/>
<path id="f-sL" d="M 380 210 C 380 230, 188 230, 188 250"/>
<path id="f-sM" d="M 380 210 L 380 250"/>
<path id="f-sR" d="M 380 210 C 380 230, 572 230, 572 250"/>
<path id="f-oL" d="M 188 278 L 188 284 C 188 302, 380 302, 380 318"/>
<path id="f-oM" d="M 380 278 L 380 318"/>
<path id="f-oR" d="M 572 278 L 572 284 C 572 302, 380 302, 380 318"/>
<path id="f-oL2" d="M 380 318 C 380 338, 188 338, 188 360"/>
<path id="f-oM2" d="M 380 318 L 380 360"/>
<path id="f-oR2" d="M 380 318 C 380 338, 572 338, 572 360"/>
</g>
<!-- Animated dots -->
<g>
<circle r="2.5" fill="#cc7154"><animateMotion dur="2.2s" repeatCount="indefinite"><mpath href="#f-in-1"/></animateMotion></circle>
<circle r="2.5" fill="#cc7154"><animateMotion dur="2.2s" repeatCount="indefinite"><mpath href="#f-in-3"/></animateMotion></circle>
<circle r="2.5" fill="#cc7154"><animateMotion dur="1.4s" repeatCount="indefinite"><mpath href="#f-sL"/></animateMotion></circle>
<circle r="2.5" fill="#cc7154"><animateMotion dur="1.4s" repeatCount="indefinite"><mpath href="#f-sR"/></animateMotion></circle>
<circle r="2.5" fill="#cc7154"><animateMotion dur="1.4s" repeatCount="indefinite"><mpath href="#f-oL"/></animateMotion></circle>
<circle r="2.5" fill="#cc7154"><animateMotion dur="1.4s" repeatCount="indefinite"><mpath href="#f-oR"/></animateMotion></circle>
<circle r="2.5" fill="#cc7154"><animateMotion dur="1.2s" repeatCount="indefinite"><mpath href="#f-oL2"/></animateMotion></circle>
<circle r="2.5" fill="#cc7154"><animateMotion dur="1.2s" repeatCount="indefinite"><mpath href="#f-oR2"/></animateMotion></circle>
</g>
<!-- === Layer 4: Junction dots === -->
<circle cx="380" cy="78" r="4" fill="#cc7154" stroke="#1e1e1e"/>
<circle cx="380" cy="210" r="4" fill="#cc7154" stroke="#1e1e1e"/>
<circle cx="380" cy="318" r="4" fill="#cc7154" stroke="#1e1e1e"/>
<!-- === Layer 5: SOURCES / OUTPUTS labels === -->
<g>
<circle cx="355" cy="12" r="2" fill="#6b7280"/>
<text x="365" y="15" font-size="9" font-weight="800" letter-spacing="0.2em" fill="#6b7280">SOURCES</text>
</g>
<g>
<circle cx="353" cy="346" r="2" fill="#6b7280"/>
<text x="363" y="349" font-size="9" font-weight="800" letter-spacing="0.2em" fill="#6b7280">OUTPUTS</text>
</g>
</svg>

After

Width:  |  Height:  |  Size: 7.7 KiB

+118
View File
@@ -0,0 +1,118 @@
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 760 420" width="760" height="420" font-family="'Inter', system-ui, -apple-system, sans-serif">
<defs>
<filter id="shadow-sm" x="-10%" y="-10%" width="120%" height="130%">
<feDropShadow dx="0" dy="2" stdDeviation="3" flood-opacity="0.04"/>
</filter>
<filter id="shadow-lg" x="-10%" y="-10%" width="120%" height="130%">
<feDropShadow dx="0" dy="6" stdDeviation="10" flood-opacity="0.06"/>
</filter>
<!-- Clip path: full area minus AETHER label zone -->
<clipPath id="gateway-clip">
<path d="M 0 0 h 760 v 420 H 0 Z M 74 169 v 54 h 12 v -54 Z" clip-rule="evenodd"/>
</clipPath>
</defs>
<!-- Background -->
<rect width="760" height="420" rx="24" fill="#fcfcfc"/>
<!-- === Labels placeholder, drawn at top layer === -->
<!-- Source pills -->
<rect x="138" y="24" width="100" height="36" rx="18" fill="#ffffff" stroke="#e8e8ec" filter="url(#shadow-sm)"/>
<circle cx="163" cy="42" r="3" fill="#cc7154"/>
<text x="195" y="46" font-size="11" font-weight="700" fill="#1f2937" text-anchor="middle">Claude</text>
<rect x="330" y="24" width="100" height="36" rx="18" fill="#ffffff" stroke="#e8e8ec" filter="url(#shadow-sm)"/>
<circle cx="355" cy="42" r="3" fill="#cc7154"/>
<text x="387" y="46" font-size="11" font-weight="700" fill="#1f2937" text-anchor="middle">OpenAI</text>
<rect x="522" y="24" width="100" height="36" rx="18" fill="#ffffff" stroke="#e8e8ec" filter="url(#shadow-sm)"/>
<circle cx="547" cy="42" r="3" fill="#cc7154"/>
<text x="579" y="46" font-size="11" font-weight="700" fill="#1f2937" text-anchor="middle">Gemini</text>
<!-- Output pills -->
<rect x="128" y="360" width="120" height="36" rx="18" fill="#ffffff" stroke="#e8e8ec" filter="url(#shadow-sm)"/>
<circle cx="153" cy="378" r="3" fill="#cc7154"/>
<text x="195" y="382" font-size="10" font-weight="700" fill="#1f2937" text-anchor="middle">Claude 响应</text>
<rect x="320" y="360" width="120" height="36" rx="18" fill="#ffffff" stroke="#e8e8ec" filter="url(#shadow-sm)"/>
<circle cx="345" cy="378" r="3" fill="#cc7154"/>
<text x="387" y="382" font-size="10" font-weight="700" fill="#1f2937" text-anchor="middle">OpenAI 响应</text>
<rect x="512" y="360" width="120" height="36" rx="18" fill="#ffffff" stroke="#e8e8ec" filter="url(#shadow-sm)"/>
<circle cx="537" cy="378" r="3" fill="#cc7154"/>
<text x="579" y="382" font-size="10" font-weight="700" fill="#1f2937" text-anchor="middle">Gemini 响应</text>
<!-- === Layer 2: AETHER GATEWAY box (z-20) === -->
<rect x="80" y="78" width="600" height="240" rx="24" fill="#ffffff" filter="url(#shadow-lg)"/>
<rect x="80" y="78" width="600" height="240" rx="24" fill="none" stroke="rgba(204,113,84,0.25)" clip-path="url(#gateway-clip)"/>
<g transform="translate(80, 176)" font-size="7" font-weight="900" fill="#cc7154" text-anchor="middle">
<text x="0" y="0">A</text><text x="0" y="9">E</text><text x="0" y="18">T</text>
<text x="0" y="27">H</text><text x="0" y="36">E</text><text x="0" y="45">R</text>
</g>
<!-- Dashed box -->
<rect x="180" y="98" width="400" height="112" rx="16" fill="none" stroke="rgba(204,113,84,0.25)" stroke-width="1.5" stroke-dasharray="6 4"/>
<rect x="200" y="108" width="360" height="26" rx="12" fill="none" stroke="#e8e8ec"/>
<text x="380" y="125" font-size="10" font-weight="600" fill="#1f2937" text-anchor="middle">统一模型规范 / 协议聚合</text>
<rect x="200" y="140" width="85" height="26" rx="13" fill="rgba(204,113,84,0.06)"/>
<text x="242" y="157" font-size="10" font-weight="600" fill="#1f2937" text-anchor="middle">多端鉴权</text>
<rect x="290" y="140" width="85" height="26" rx="13" fill="rgba(204,113,84,0.06)"/>
<text x="332" y="157" font-size="10" font-weight="600" fill="#1f2937" text-anchor="middle">配额管控</text>
<rect x="380" y="140" width="85" height="26" rx="13" fill="rgba(204,113,84,0.06)"/>
<text x="422" y="157" font-size="10" font-weight="600" fill="#1f2937" text-anchor="middle">全局并发</text>
<rect x="470" y="140" width="85" height="26" rx="13" fill="rgba(204,113,84,0.06)"/>
<text x="512" y="157" font-size="10" font-weight="600" fill="#1f2937" text-anchor="middle">缓存亲和</text>
<rect x="200" y="172" width="360" height="28" rx="14" fill="none" stroke="#e8e8ec"/>
<text x="380" y="190" font-size="10" font-weight="600" fill="#1f2937" text-anchor="middle">智能调度 / 故障转移</text>
<!-- Engine cards -->
<rect x="103" y="250" width="170" height="28" rx="14" fill="rgba(204,113,84,0.06)" stroke="#e8e8ec"/>
<circle cx="163" cy="264" r="3" fill="#cc7154"/>
<text x="192" y="268" font-size="10" font-weight="700" fill="#1f2937" text-anchor="middle">格式转换</text>
<rect x="295" y="250" width="170" height="28" rx="14" fill="rgba(204,113,84,0.06)" stroke="#e8e8ec"/>
<circle cx="355" cy="264" r="3" fill="#cc7154"/>
<text x="384" y="268" font-size="10" font-weight="700" fill="#1f2937" text-anchor="middle">反向代理</text>
<rect x="487" y="250" width="170" height="28" rx="14" fill="rgba(204,113,84,0.06)" stroke="#e8e8ec"/>
<circle cx="547" cy="264" r="3" fill="#cc7154"/>
<text x="576" y="268" font-size="10" font-weight="700" fill="#1f2937" text-anchor="middle">原生直通</text>
<!-- === Layer 3: Track lines (z-30, above Gateway fill) === -->
<g stroke="#e5e7eb" stroke-width="1.5" fill="none" stroke-linecap="round">
<path id="f-in-1" d="M 188 60 L 188 66 C 188 72, 380 72, 380 78 L 380 98"/>
<path id="f-in-2" d="M 380 60 L 380 98"/>
<path id="f-in-3" d="M 572 60 L 572 66 C 572 72, 380 72, 380 78 L 380 98"/>
<path id="f-sL" d="M 380 210 C 380 230, 188 230, 188 250"/>
<path id="f-sM" d="M 380 210 L 380 250"/>
<path id="f-sR" d="M 380 210 C 380 230, 572 230, 572 250"/>
<path id="f-oL" d="M 188 278 L 188 284 C 188 302, 380 302, 380 318"/>
<path id="f-oM" d="M 380 278 L 380 318"/>
<path id="f-oR" d="M 572 278 L 572 284 C 572 302, 380 302, 380 318"/>
<path id="f-oL2" d="M 380 318 C 380 338, 188 338, 188 360"/>
<path id="f-oM2" d="M 380 318 L 380 360"/>
<path id="f-oR2" d="M 380 318 C 380 338, 572 338, 572 360"/>
</g>
<!-- Animated dots -->
<g>
<circle r="2.5" fill="#cc7154"><animateMotion dur="2.2s" repeatCount="indefinite"><mpath href="#f-in-1"/></animateMotion></circle>
<circle r="2.5" fill="#cc7154"><animateMotion dur="2.2s" repeatCount="indefinite"><mpath href="#f-in-3"/></animateMotion></circle>
<circle r="2.5" fill="#cc7154"><animateMotion dur="1.4s" repeatCount="indefinite"><mpath href="#f-sL"/></animateMotion></circle>
<circle r="2.5" fill="#cc7154"><animateMotion dur="1.4s" repeatCount="indefinite"><mpath href="#f-sR"/></animateMotion></circle>
<circle r="2.5" fill="#cc7154"><animateMotion dur="1.4s" repeatCount="indefinite"><mpath href="#f-oL"/></animateMotion></circle>
<circle r="2.5" fill="#cc7154"><animateMotion dur="1.4s" repeatCount="indefinite"><mpath href="#f-oR"/></animateMotion></circle>
<circle r="2.5" fill="#cc7154"><animateMotion dur="1.2s" repeatCount="indefinite"><mpath href="#f-oL2"/></animateMotion></circle>
<circle r="2.5" fill="#cc7154"><animateMotion dur="1.2s" repeatCount="indefinite"><mpath href="#f-oR2"/></animateMotion></circle>
</g>
<!-- === Layer 4: Junction dots (z-40, topmost) === -->
<circle cx="380" cy="78" r="4" fill="#cc7154" stroke="#ffffff"/>
<circle cx="380" cy="210" r="4" fill="#cc7154" stroke="#ffffff"/>
<circle cx="380" cy="318" r="4" fill="#cc7154" stroke="#ffffff"/>
<!-- === Layer 5: SOURCES / OUTPUTS labels (topmost) === -->
<g>
<circle cx="355" cy="12" r="2" fill="#8c94a1"/>
<text x="365" y="15" font-size="9" font-weight="800" letter-spacing="0.2em" fill="#8c94a1">SOURCES</text>
</g>
<g>
<circle cx="353" cy="346" r="2" fill="#8c94a1"/>
<text x="363" y="349" font-size="9" font-weight="800" letter-spacing="0.2em" fill="#8c94a1">OUTPUTS</text>
</g>
</svg>

After

Width:  |  Height:  |  Size: 7.8 KiB

-262
View File
@@ -1,262 +0,0 @@
# Aether 请求体规则 (Body Rules) — AI 辅助指南
> 本文档面向 AI 助手。用户在「端点管理 → 请求规则 → + 请求体」中添加规则,AI 需要指导用户在 UI 表单的各个输入框中填写什么内容。
## 系统背景
Aether 是一个 AI API 网关,将用户请求转发给上游 Provider(OpenAI、Claude、Gemini 等)。请求体规则在转发前修改请求体 JSON,每个 Endpoint 可配置多条规则。
## 用户请求体长什么样
规则操作的对象是用户发给 Aether 的 JSON 请求体。根据 Endpoint 的 API 格式不同,结构也不同:
### OpenAI Chat 格式(`openai:chat`)
```json
{
"model": "gpt-4",
"messages": [
{"role": "system", "content": "你是一个助手"},
{"role": "user", "content": "你好"}
],
"temperature": 0.7,
"max_tokens": 1000,
"stream": true,
"top_p": 1,
"frequency_penalty": 0,
"presence_penalty": 0,
"stop": ["\n"],
"tools": [...],
"tool_choice": "auto",
"response_format": {"type": "json_object"},
"metadata": {...}
}
```
### Claude Chat 格式(`claude:chat`)
```json
{
"model": "claude-sonnet-4-20250514",
"messages": [
{"role": "user", "content": "你好"}
],
"system": "你是一个助手",
"max_tokens": 1024,
"temperature": 0.7,
"stream": true,
"top_p": 1,
"top_k": 40,
"stop_sequences": ["###"],
"metadata": {"user_id": "xxx"}
}
```
### Claude CLI 格式(`claude:cli`)
```json
{
"model": "claude-sonnet-4-20250514",
"messages": [
{"role": "user", "content": "帮我重构这段代码"}
],
"system": "你是一个编程助手",
"max_tokens": 16384,
"stream": true
}
```
### Gemini Chat 格式(`gemini:chat`)
```json
{
"contents": [
{"role": "user", "parts": [{"text": "你好"}]}
],
"generationConfig": {
"temperature": 0.7,
"maxOutputTokens": 1024,
"topP": 0.9,
"topK": 40
},
"systemInstruction": {"parts": [{"text": "你是一个助手"}]},
"safetySettings": [...]
}
```
> **注意:** `model` 和 `stream` 是受保护字段,任何规则都无法修改它们。
---
## UI 表单说明
用户在端点管理对话框中点击「+ 请求体」添加规则。每条规则有一个下拉框选择操作类型,然后根据类型显示不同的输入框。
---
### 覆写(set)
**UI 布局:** `[覆写 ▼]` `[字段路径]` `=` `[值]` `[✓]`
| 输入框 | 填什么 | 示例 |
|--------|--------|------|
| 字段路径 | 要设置的字段,用 `.` 分隔层级,`[N]` 访问数组 | `temperature` |
| 值 | **JSON 格式**的值。字符串要加引号,数字直接写 | `0.7` |
值输入框右边有验证图标:绿色勾=JSON合法,红色叉=格式错误。
**用户常见需求 → 怎么填:**
| 用户说 | 字段路径 | 值 |
|--------|---------|-----|
| "固定 temperature 为 0.3" | `temperature` | `0.3` |
| "限制最大输出 500 token" | `max_tokens` | `500` |
| "加一个 metadata 字段标记来源" | `metadata.source` | `"my-app"` |
| "设置 top_p 为 0.9" | `top_p` | `0.9` |
| "添加停止序列" | `stop` | `["\n", "###"]` |
| "设置响应格式为 JSON" | `response_format` | `{"type": "json_object"}` |
| "把第一条消息内容改掉" | `messages[0].content` | `"新的内容"` |
| "设置一个嵌套对象" | `metadata.tracking` | `{"id": "abc", "env": "prod"}` |
| "设置值为空" | `some_field` | `null` |
> **注意:** 字符串值必须加引号!`"hello"` 是字符串,`hello` 是无效 JSON。数字、布尔、null、数组、对象不需要额外引号。
---
### 删除(drop)
**UI 布局:** `[删除 ▼]` `[要删除的字段路径]`
| 输入框 | 填什么 | 示例 |
|--------|--------|------|
| 字段路径 | 要删除的字段路径 | `user_info.ip_address` |
**用户常见需求 → 怎么填:**
| 用户说 | 字段路径 |
|--------|---------|
| "去掉 user 字段" | `user` |
| "删除 metadata 里的 internal_flag" | `metadata.internal_flag` |
| "移除第一条消息" | `messages[0]` |
| "去掉 frequency_penalty" | `frequency_penalty` |
> **注意:** 删除数组元素会导致后续索引前移。如需删除多个数组元素,从后往前删。
---
### 重命名(rename)
**UI 布局:** `[重命名 ▼]` `[原路径]` `→` `[新路径]`
| 输入框 | 填什么 | 示例 |
|--------|--------|------|
| 原路径 | 源字段的路径 | `extra.trace_id` |
| 新路径 | 目标字段的路径(中间层级会自动创建) | `metadata.request_id` |
**用户常见需求 → 怎么填:**
| 用户说 | 原路径 | 新路径 |
|--------|--------|--------|
| "把 max_tokens 改名为 max_completion_tokens" | `max_tokens` | `max_completion_tokens` |
| "把 extra 里的 id 移到 metadata 下" | `extra.id` | `metadata.id` |
---
### 插入(insert)
**UI 布局:** `[插入 ▼]` `[数组路径]` `[位置]` `[值(JSON)]` `[✓]`
| 输入框 | 填什么 | 示例 |
|--------|--------|------|
| 数组路径 | 目标数组的路径(必须是已有的数组) | `messages` |
| 位置 | 插入位置的数字索引,**留空=追加到末尾** | `0`(开头)或留空(末尾) |
| 值 | JSON 格式的元素 | `{"role": "system", "content": "..."}` |
**用户常见需求 → 怎么填:**
| 用户说 | 数组路径 | 位置 | 值 |
|--------|---------|------|-----|
| "在开头加一条 system 消息" | `messages` | `0` | `{"role": "system", "content": "你是一个专业助手"}` |
| "在末尾追加一条消息" | `messages` | (留空) | `{"role": "user", "content": "请用中文回答"}` |
| "在第二条消息前插入" | `messages` | `1` | `{"role": "assistant", "content": "好的"}` |
> **位置说明:** `0`=最前面,`1`=第二个位置,`-1`=倒数第一个前面。留空=追加到最后。
---
### 正则替换(regex_replace)
**UI 布局:** `[正则替换 ▼]` `[字段路径]` `[正则]` `→` `[替换为]` `[ims]` `[✓]`
| 输入框 | 填什么 | 示例 |
|--------|--------|------|
| 字段路径 | 目标字符串字段的路径 | `messages[-1].content` |
| 正则 | 正则表达式(**不需要**加 `/` 包裹) | `1[3-9]\d{9}` |
| 替换为 | 替换成什么,留空=删除匹配内容 | `[手机号已隐藏]` |
| ims | 可选标志,留空=默认 | `i` |
> **重要:** 正则输入框里直接写正则语法即可,**不需要** JSON 转义。`\d` 就写 `\d`,不用写 `\\d`。JSON 转义是保存时系统自动处理的。
**flags 含义:**
| 字母 | 作用 | 什么时候用 |
|------|------|-----------|
| `i` | 忽略大小写 | 匹配 `hello`/`Hello`/`HELLO` |
| `m` | 多行模式 | `^`/`$` 匹配每行而非整个字符串 |
| `s` | dotall | `.` 能匹配换行符 |
大多数情况留空即可。
**用户常见需求 → 怎么填:**
| 用户说 | 字段路径 | 正则 | 替换为 | flags |
|--------|---------|------|--------|-------|
| "隐藏最后一条消息里的手机号" | `messages[-1].content` | `1[3-9]\d{9}` | `[手机号已隐藏]` | |
| "隐藏邮箱地址" | `messages[-1].content` | `[\w.+-]+@[\w.-]+\.\w{2,}` | `[邮箱已隐藏]` | `i` |
| "去掉 HTML 标签" | `messages[-1].content` | `<[^>]+>` | (留空) | |
| "把所有的 foo 替换成 bar" | `messages[-1].content` | `\bfoo\b` | `bar` | |
| "删掉 Markdown 加粗标记" | `messages[-1].content` | `\*\*([^*]+)\*\*` | `\1` | |
| "隐藏身份证号" | `messages[-1].content` | `\d{17}[\dXx]` | `[身份证已隐藏]` | |
| "隐藏 IP 地址" | `messages[-1].content` | `\d{1,3}\.\d{1,3}\.\d{1,3}\.\d{1,3}` | `[IP已隐藏]` | |
---
## 路径语法速查
| 写法 | 含义 |
|------|------|
| `temperature` | 顶层字段 |
| `metadata.source` | 嵌套字段 |
| `messages[0]` | 数组第一个元素 |
| `messages[-1]` | 数组最后一个元素 |
| `messages[0].content` | 第一条消息的 content |
| `messages[-1].content` | 最后一条消息的 content |
| `data[0].items[2]` | 嵌套数组访问 |
| `config\.v1.enabled` | key 名字里有点号时用 `\.` 转义 |
---
## 多条规则示例
用户说:"帮我配置规则:在开头插入一条 system 消息说'用中文回答',固定 temperature 为 0.3,把用户消息里的手机号脱敏"
应指导用户添加 3 条规则:
| # | 操作 | 字段 1 | 字段 2 | 字段 3 | 字段 4 |
|---|------|--------|--------|--------|--------|
| 1 | 插入 | 路径: `messages` | 位置: `0` | 值: `{"role": "system", "content": "请用中文回答所有问题"}` | |
| 2 | 覆写 | 路径: `temperature` | 值: `0.3` | | |
| 3 | 正则替换 | 路径: `messages[-1].content` | 正则: `1[3-9]\d{9}` | 替换为: `[手机号]` | flags: (留空) |
规则按从上到下的顺序执行。
---
## 注意事项
1. **`model` 和 `stream` 不可修改** — 这两个字段由系统管理,写了规则也会被跳过
2. **值必须是合法 JSON** — 覆写/插入的值输入框要求 JSON 格式:字符串加引号 `"text"`,数字直接写 `123`,布尔 `true`/`false`
3. **路径必须指向正确类型** — 插入的路径必须是数组,正则替换的路径必须是字符串
4. **路径不存在时的行为** — 覆写会自动创建中间层级(dict),其他操作遇到不存在的路径会静默跳过
5. **正则在 UI 里直接写** — 不需要 JSON 转义,`\d` 就是 `\d`
6. **正则保存时校验** — 无效的正则表达式会在保存时报错,不会静默通过
+20
View File
@@ -1,6 +1,26 @@
#!/bin/bash
set -e
# Wait for PostgreSQL to be ready
MAX_ATTEMPTS=30
ATTEMPT=0
until python -c "
from sqlalchemy import create_engine, text
import os
engine = create_engine(os.environ['DATABASE_URL'])
with engine.connect() as conn:
conn.execute(text('SELECT 1'))
" 2>/dev/null; do
ATTEMPT=$((ATTEMPT + 1))
if [ "$ATTEMPT" -ge "$MAX_ATTEMPTS" ]; then
echo "Database not ready after $MAX_ATTEMPTS attempts, exiting."
exit 1
fi
echo "Waiting for database... (attempt $ATTEMPT/$MAX_ATTEMPTS)"
sleep 2
done
echo "Database is ready."
echo "Running database migrations..."
alembic upgrade head
+7
View File
@@ -24,6 +24,7 @@
"marked": "^16.0.0",
"otpauth": "^9.5.0",
"pinia": "^3.0.3",
"pinyin-pro": "^3.28.0",
"radix-vue": "^1.9.17",
"tailwind-merge": "^3.3.1",
"three": "^0.180.0",
@@ -4718,6 +4719,12 @@
}
}
},
"node_modules/pinyin-pro": {
"version": "3.28.0",
"resolved": "https://registry.npmmirror.com/pinyin-pro/-/pinyin-pro-3.28.0.tgz",
"integrity": "sha512-mMRty6RisoyYNphJrTo3pnvp3w8OMZBrXm9YSWkxhAfxKj1KZk2y8T2PDIZlDDRsvZ0No+Hz6FI4sZpA6Ey25g==",
"license": "MIT"
},
"node_modules/pirates": {
"version": "4.0.7",
"resolved": "https://registry.npmjs.org/pirates/-/pirates-4.0.7.tgz",
+1
View File
@@ -32,6 +32,7 @@
"marked": "^16.0.0",
"otpauth": "^9.5.0",
"pinia": "^3.0.3",
"pinyin-pro": "^3.28.0",
"radix-vue": "^1.9.17",
"tailwind-merge": "^3.3.1",
"three": "^0.180.0",
+66 -2
View File
@@ -5,12 +5,14 @@
</template>
<script setup lang="ts">
import { onMounted, onErrorCaptured } from 'vue'
import { onMounted, onErrorCaptured, onUnmounted } from 'vue'
import { useAuthStore } from '@/stores/auth'
import ToastContainer from '@/components/ToastContainer.vue'
import ConfirmContainer from '@/components/ConfirmContainer.vue'
import apiClient from '@/api/client'
import apiClient, { AUTH_STATE_CHANGE_EVENT } from '@/api/client'
import { NETWORK_CONFIG, AUTH_CONFIG } from '@/config/constants'
import router from '@/router'
import { hasAuthIdentityChanged } from '@/utils/authToken'
import { log } from '@/utils/logger'
const authStore = useAuthStore()
@@ -86,7 +88,59 @@ if (typeof window !== 'undefined') {
})
}
async function syncExternalAuthState(nextToken: string | null): Promise<void> {
const previousToken = authStore.token
const previousUser = authStore.user
? {
id: authStore.user.id,
role: authStore.user.role,
}
: null
authStore.syncToken()
if (!nextToken) {
if (previousToken || previousUser) {
authStore.applyExternalLogout()
await router.replace('/')
}
return
}
const identityChanged = hasAuthIdentityChanged(previousToken, nextToken, previousUser)
if (!identityChanged && previousUser) {
return
}
const user = await authStore.fetchCurrentUser()
if (!user) {
return
}
if (router.currentRoute.value.path.startsWith('/admin') && user.role !== 'admin') {
await router.replace('/dashboard')
}
}
function handleAuthStorageChange(event: StorageEvent): void {
if (event.key !== 'access_token') {
return
}
syncExternalAuthState(event.newValue).catch((err) => log.error('syncExternalAuthState failed', err))
}
function handleLocalAuthStateChange(event: Event): void {
const authEvent = event as CustomEvent<{ token: string | null }>
syncExternalAuthState(authEvent.detail?.token ?? apiClient.getToken()).catch((err) => log.error('syncExternalAuthState failed', err))
}
onMounted(async () => {
if (typeof window !== 'undefined') {
window.addEventListener('storage', handleAuthStorageChange)
window.addEventListener(AUTH_STATE_CHANGE_EVENT, handleLocalAuthStateChange as (event: Event) => void)
}
// 延迟检查认证状态,让页面先加载
setTimeout(async () => {
try {
@@ -97,4 +151,14 @@ onMounted(async () => {
}
}, AUTH_CONFIG.TOKEN_REFRESH_INTERVAL)
})
onUnmounted(() => {
if (typeof window !== 'undefined') {
window.removeEventListener('storage', handleAuthStorageChange)
window.removeEventListener(
AUTH_STATE_CHANGE_EVENT,
handleLocalAuthStateChange as (event: Event) => void,
)
}
})
</script>

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