Compare commits

...
44 Commits
Author SHA1 Message Date
elky 535ee098c3 fix(auth): reject unsigned admin identity headers 2026-08-18 11:12:17 +08:00
ZheFox b45df89ce4 Merge pull request #730 from zhefox/fix/responses-websocket-current
feat(gateway): add OpenAI Responses WebSocket mode
2026-08-17 21:14:50 +08:00
ZheFox c8118edf36 fix(ws): harden Responses connection lifecycle
Revalidate control policy per turn, isolate downstream credentials, and make planner/turn ownership cancellation-safe.

Preserve opaque protocol events, align configurable timeout semantics, and extend end-to-end security and settlement coverage.
2026-08-17 18:50:29 +08:00
AAEE86 4a0775c4ea style: apply cargo fmt across gateway and aether-ai crates 2026-08-17 14:53:53 +08:00
AAEE86 6fc02dad3e fix(ws): restore redacted PII in provider frames before client delivery
Responses WebSocket 只实现了脱敏的一半:请求侧 mask 之后,provider 事件帧在推给
客户端之前没有还原,于是 session 映射内的占位符以 <AETHER:EMAIL:...> 的形式直接
透给客户端。这里补齐响应侧,语义与 HTTP 路径对齐。

- 还原点是 relay loop 的最后一跳(send_client_message 之前、capture_client_frame
  之前),对应 HTTP 的 restore_sync_response_body / StreamingResponseRestorer 所在
  位置。审计与终态观测继续消费脱敏态事件,只有发往客户端的那一份拷贝被还原。
- 复用 privacy::restore_json_strings(改为 pub(crate))与
  RedactionSession::restore_text,不复制任何还原逻辑:只还原本 session mask 过的
  映射,未映射的占位符原样保留;type / model / id 等协议字段不可能命中 sentinel,
  因此不受影响。批量 {"chunks":[...]} 帧一并递归还原。
- session 生命周期:mask 仍然是 per-turn(slot 依旧每轮新建),但 session 改由连接
  持有,按有界 FIFO 留最近 8 轮。理由是 WS 的会话历史留在上游,continuation 只发
  增量输入,per-turn 释放会漏还原后续响应里回显的更早轮次占位符;HTTP 不会漏,是
  因为它每次重发整段历史、重新 mask 会派生出同一个 sentinel。被挤出窗口的轮次退回
  「占位符原样透传」,不会错误还原成别的值。
- 未命中还原时不改写字节;连接上没有任何 mask session 时(未启用脱敏)连事件 clone
  都不做。

测试:redaction.rs 新增 8 条单测(还原命中/批量帧/未映射占位符原样/未命中不改写/
无 session 不介入/空 session 不留存/审计侧入参不被改写/跨轮还原/窗口有界);
responses_websocket_e2e 新增一条用例,mock 上游回显收到的 input,断言上游只看到
占位符而客户端拿到真实邮箱。
2026-08-17 14:53:46 +08:00
AAEE86 dbf2809bd6 fix(ws): settle the previous attempt before transparent retry replanning
评审第 2 条。配额透明重试原来的顺序是「detach 旧 attempt → 规划并绑定新
attempt → 把旧 attempt 的结算排进队列」。规划因此读到的是旧 attempt 还没投射的
health / adaptive / pool 状态,而且旧 attempt 仍占着自己的 pool key lease——替代
key 的挑选看到的是一把仍被占用的 key,最坏情况下判成「无可用供应商」而放弃一次
本可以成功的重试。

普通的新 turn 早就挡住了这件事:client.rs 在处理 response.create 前调用
await_pending_turn_finalization,注释写的正是「不要让新 turn 基于陈旧的 health /
adaptive / pool 状态规划」。透明重试是同一个问题的另一条入口,漏了这一步。

现在顺序是:detach → 释放准入 → 结算旧 attempt 并等它落地 → 规划/绑定新 attempt。

新增 lifecycle::settle_turn_finalization:与 queue_turn_finalization 的区别只在于
「等」。后者把 handle 挂在连接上让 relay loop 继续跑,用在结算之后不再读取共享
状态的出口;前者用在必须先看到结算结果才能继续的路径上。

顺序用类型固定,而不是靠注释:settle_turn_finalization 返回
PreviousAttemptSettled,retry_active_turn_after_quota_exhaustion 要求这个参数。
凭证只能由 lifecycle 颁发(结算完成,或明确「没有 attempt 要结算」),所以把顺序
写反连编译都过不了。

重试失败路径随之变化:旧 attempt 已经结算,不再 resume 回去。logical turn 仍停在
Replanning,后续分支的 end() / finalize_active_turn 只清 logical turn、不交出
attempt,因此不存在重复结算。结算 outcome 取值不变(两条路径用的都是
terminal_outcome.unwrap_or_else(upstream_closed),而这条分支里 terminal_outcome
必为 Some——usage_limit_error 成立意味着有一个已解析的 error 终态帧)。

代价(都落在「重试失败」这一侧,且只影响已终态 attempt 的报告注解,不影响计费):
- 那条最终转发给客户端的 429 事件不再进旧 attempt 的 client capture;
  provider 侧 capture 早在 observe_upstream_frame 里就记下了。
- 如果转发 429 给客户端也失败,record_client_delivery_aborted 落在一个已经结算的
  attempt 上,成为 no-op。

测试:
- lifecycle:await_turn_finalization_handle 必须「等到落地」而不是「排进队列」
  (C6 依赖的性质);结算完成后规划才读状态的顺序型断言(计数器替身);结算任务
  panic 也必须放行调用方,不能卡死 relay loop。
- turn_state:Replanning 状态下 end() 不再交出第二个 attempt(无重复结算)。
- e2e 新增 provider_quota_exhaustion_transparently_retries_onto_another_key:
  mock 上游首轮只回 Codex 的 429 usage_limit_reached,网关换到第二把 key 重放同一个
  response.create;断言客户端看不到 429、上游被连两次、两次用的不是同一把 key、两个
  attempt 各留一条终态行(429 的那条 + 计费的那条)。已验证它在改动前后都通过——
  它覆盖的是整条路径可用,顺序由上面的单测确定性覆盖。
  夹具随之参数化出 ProviderFixture::CodexKeyPair:透明重试只有 Codex adapter 会
  开启,而 codex 候选要求 auth_type = oauth,所以这个夹具用未过期的 oauth 凭证。
2026-08-17 14:53:40 +08:00
AAEE86 1d3051cb89 refactor(ws): structured terminal observation without SSE text round-trips
评审第 5 条:Responses WebSocket 收到的本来就是结构化协议事件,但为了复用面向
SSE 的 push_line,观测路径要先把每个事件序列化成 data: {json}\n\n,解析器再
decode 回 Value——一次纯粹的往返。这个「伪 SSE」形状是随手拼的,一旦拼装函数
以后被加上换行或分块逻辑,观测结果就会和真实事件悄悄分叉。

aether-ai-formats:
- OpenAIResponsesProviderState::push_line 机械拆成 decode + push_event,
  push_line 现在只做解码。协议状态机一行未动,diff 里除函数签名外只有
  &value → value(value 从拥有改成借用,持有结构化事件的传输不必为了调用它
  先克隆一份)。
- StreamingStandardTerminalObserver::push_event 走 TerminalStreamParser::Standard,
  service tier 的记录方式与 push_line 完全相同。openai:image 的终态状态机按 SSE
  行做增量解析、没有结构化入口,返回 AiSurfaceFinalizeError 让调用方
  disable_with_error 标记 parser_error,而不是静默丢事件、把摘要留成「未观察到
  终态」。ProviderStreamParser 的其余三个格式同样返回 Err:机械拆分随时可做,
  但不建无调用方的接口。

WS 侧:
- 新增 responses/observation.rs 的 ResponsesStructuredTerminalObserver,直接消费
  frame.protocol_events() 借出的事件。包一层的意义是让「不再拼 SSE」成为类型层面
  的事实——这个类型没有任何接受字节的方法,改回 push_line 不可能悄悄发生。
  finish() 里的 Ok(None) / Err → disable_with_error 兜底也一并收进来。
- body capture 不动,仍然是 SSE 形状(data: 开头、\n\n 结尾):
  aether_usage_runtime::report 用 line.strip_prefix("data:") 解析被捕获的 body
  判定 StreamCapturedTerminalState,而它是 stream_report_represents_failure 的一个
  OR 项,换成结构化 JSON 会让终态判定恒为 Missing。capture_sse_event /
  capture_client_frame / websocket_event_as_sse_line 全部保留,原因写在模块文档
  注释里。这一层只换观测,不换捕获。

差分测试(8 个,aether-ai-formats):同一组事件序列分别走 push_line 与
push_event,断言 ExecutionStreamTerminalSummary 完全相等——批量 delta 序列、
completed 带 usage、合法 incomplete、error、response.failed、未知事件、
service tier、缺终态;外加 openai:image 拒绝结构化入口。两条入口不可能有
过滤差异:任何 Value 序列化出来都不会命中 decode_json_data_line 的 empty /
":" / "event:" / [DONE] 四个过滤条件。

turn.rs 里三个既有的 WS 观测测试改走结构化入口;SSE 形状的断言留在 capture 一侧。
验收:crates/aether-usage 零 diff。
2026-08-17 14:53:33 +08:00
AAEE86 59e27524da refactor(gateway): extract transport-neutral execution attempt lifecycle
评审第 4 条:responses/turn.rs 实际复制了一整套 HTTP execution lifecycle——
usage 写入、candidate 状态流转、health/adaptive 效果投射、pool key lease 释放、
body capture、账单失败判定,与 HTTP 的顺序和超时语义只能靠人工对齐。

新增 execution_runtime/attempt_lifecycle.rs,把一次 provider attempt 的记账收成
transport 中立的三段:

  ExecutionAttemptLifecycle::begin        pending usage 行 + Pending candidate
  ExecutionAttemptLifecycle::mark_started usage stream_started + Streaming candidate(幂等)
  ExecutionAttemptLifecycle::settle       终态四段,顺序不可重排:
                                            1 usage terminal(detachable,不可丢)
                                            2 candidate terminal
                                            3 provider 效果 + 超时兜底释放 lease
                                            4 execution report(作废账单不提交)

顺序、5s 分段超时常量、detachable 语义、「每个效果分支都释放 lease」「作废账单
一律不提交 report」这些不变量全部保持原样。

一并上移的辅助设施:
- AttemptStageGuard 取代 await_websocket_lifecycle_stage /
  await_detachable_lifecycle_stage,把「等多久」参数化:WS 用 Bounded(5s),
  HTTP 接线时用 Unbounded 即保持它现在的语义。
- AttemptBodyCapture 取代 append_capture / encode_stream_capture,把
  「缓冲 + 截断标志」两个字段收成一个类型(WS 侧四个字段变两个)。捕获内容
  仍然是 SSE 形状:usage runtime 按 data: 行解析被捕获的 body 来判定
  StreamCapturedTerminalState,换成结构化 JSON 会让终态判定恒为 Missing。
- C2/C3 的结算表本来就不含任何 WS 类型,随之上移。效果表分支与注释逐字未改,
  仅按新位置改名为 AttemptProviderEffect / classify_attempt_provider_effect。
  responses/settlement.rs 只保留 WS 专属的一件事:把 relay loop 的结算信号
  ResponsesWebSocketTurnOutcome 翻译成两个正交事实。

ResponsesProviderAttempt 现在只持有 WS 专有状态:lifecycle 句柄、deadline、
终态观测器、两侧 capture、准入、provider/delivery 事实。plan / trace_id /
report_kind / report_context / candidate 起始时间戳都归 lifecycle。

HTTP 侧不接线:execution_runtime/stream/execution.rs 的
DirectPassthroughFinalizerCore(38 字段)与 failover / oauth 重试 / prefetch 深度
纠缠,无法在「行为等价 + 单 commit 可验证」的前提下改动。逐调用点映射表写在
模块文档注释里作为后续 PR 的接线依据。验收:git diff 对
execution_runtime/stream/ 与 crates/aether-usage 均为零 diff。

新增 6 个测试:效果段超时后仍走兜底 lease 释放、Unbounded 会一直等、detachable
写入在调用方停止等待后仍跑完、settle 四段顺序(计数器替身)、body capture 的
SSE 形状与编码状态(并显式记录默认上限是 usize::MAX,截断分支不可达)、
candidate error_type 映射。
2026-08-17 14:53:25 +08:00
AAEE86 247e7105a2 fix(ws): bill a provider-reached terminal even when client delivery fails
评审第 5 条后半:provider 终态已经到达、只是 gateway 写客户端 socket 失败时,
relay loop 用 client_disconnected() 覆盖了结算信号,于是一条供应商已经完成推理
并消耗了 token 的响应被记成 void billing、candidate 记 Cancelled、不投射供应商
效果、也不提交 execution report。上游成本凭空消失。

结算表只改一行:作废账单的条件从
    provider.cancelled_by_provider() || delivery.is_aborted()
收紧为
    provider.cancelled_by_provider() || (delivery.is_aborted() && !provider.is_terminal())

于是 Terminal{cancelled=false} + delivery Aborted 与 delivery Complete 落在同一侧:
Billed、candidate Success 或 Failed、投射供应商效果、提交 execution report。
状态码随之变成纯 provider 事实(不再把 200 改写成 499);作废分支的 provider
状态码本身就是 499,取值不变。

依据:供应商已经完成推理并消耗 token,客户端还能用 previous_response_id 续取
这条响应。供应商没给出终态时(客户端先走了)仍然作废,这一侧未改。

配套改动:
- connection.rs 写客户端失败处改为 record_client_delivery_aborted(reason) +
  settle_signal_for_client_delivery_failure(terminal_outcome):provider 终态已到达
  就用那条终态作结算信号,不再无条件覆盖。投递失败原因也不再谎称
  「客户端在终态前断开」。
- 投递结果记在 attempt 上而非 logical turn 上:结算按 attempt 进行,且配额透明
  重试时各 attempt 的投递结果彼此独立。
- report_context 新增 websocket_client_delivery="aborted" 与
  websocket_client_delivery_reason,只增字段不改既有字段,便于事后区分
  「客户端拿到了」和「客户端没拿到但已计费」。
- candidate error_type 新增 client_delivery_failed(原先这个场景写的是
  websocket_cancelled)。它排在供应商侧分类之前:这条记录之所以特别正是因为
  内容没送到客户端,供应商侧判定仍由 candidate_status 与 error_message 保留。
- finish_summary 改用作废判定而非「投递失败」判定:provider 终态已到达时摘要
  必须保留真实的 finish_reason 与 usage,否则计费记录会被写坏。

e2e 期望值变化:client_disconnect_mid_turn_still_settles_the_usage_row 改名为
client_disconnect_before_any_provider_output_settles_a_void_row,并补上
「不计费 + status=cancelled + status_code=499」的断言。原用例的 mock 行为是
StallAfterCreated(只发 response.created 就静默),provider 从未给出终态,所以
它走的是未改动的作废一侧;原来的文档注释说「must still be billed」与实际语义
不符,一并纠正。真正被修正的那一行无法在 e2e 里确定性触发——它取决于 relay
loop 的 select! 先观察到上游终态帧还是先观察到已关闭的客户端 socket,是构造性
竞态——因此由 relay 级单测确定性覆盖,e2e 里以注释指向这两个单测。

新增 7 个测试:结算表修正行(并与「投递成功」逐字段对照,只有 candidate 错误
分类不同)、无终态时仍作废、供应商声明取消即使送达也不计费、结算信号选择、
已记录的投递失败不被结算信号覆盖、relay 级「终态到达 + 客户端已关闭 ⇒ Billed /
Success / ProviderSuccess / 已提交 report 且 usage 完整保留」及其镜像、
report_context 只增不改。
2026-08-17 14:53:12 +08:00
AAEE86 dc3743aecf refactor(ws): 拆分 LogicalTurn 与 ProviderAttempt,结算改表驱动
评审第 5 条:一个 ResponsesWebSocketTurn 同时代表 logical turn 和 provider
attempt,finalize() 又用 outcome.cancelled() 一个布尔驱动 billing、candidate
状态和供应商效果,于是 provider 终态已经到达、只是最后一跳写客户端失败时,
供应商事实会被 Cancelled 覆盖掉。

- ResponsesWebSocketTurn → ResponsesProviderAttempt,
  ActiveResponsesWebSocketTurn → ActiveProviderAttempt:类型名字明确它只代表
  一次上游执行,logical turn 由 C1 落地的 LogicalTurn 承担。
- 新增 settlement.rs:AttemptProviderOutcome × AttemptClientDelivery 两个正交
  事实,classify_attempt_settlement 一张表推出 status_code / billing /
  candidate 状态 / candidate 错误分类 / 供应商效果 / 是否提交 execution report。
- attempt 观察到 provider 终态即记录 provider_outcome。结算信号
  ResponsesWebSocketTurnOutcome 只回答「为什么现在结算」:ProviderTerminal 与
  Failure 对 provider 是权威的,Cancelled 只描述客户端/连接层面的停止,不再
  覆盖已观察到的 provider 事实。
- candidate 状态与 candidate 错误分类分开输出:现状存在
  「missing_terminal=true 而记账层判 Success」的组合(report kind 不要求观察到
  终态事件时),会写出 status=Success + error_type=stream_missing_terminal_event,
  这个组合必须原样保留。

classify_responses_websocket_turn_effect 的判定表原样搬入 settlement.rs,分支
和顺序均未改动,两个既有不变量测试随之迁移。

行为等价。结算表当前口径与拆分前完全一致:客户端投递失败仍与「供应商声明取消」
落在同一侧(作废账单、candidate 记 Cancelled、只释放 lease、不提交 execution
report),即使 provider 终态已经到达——这一行由
settlement_table_row_client_delivery_failure_currently_voids_a_reached_terminal
锁住现状,修正它是下一步独立的行为修正。

新增 15 个测试:outcome → 双事实映射表逐行(含 stream_timeout 只在 504 失败一族
成立、provider 终态即使 504 也不投射流式超时)、结算表逐行、投递失败时
forced_error 必须为 None、已观察终态不被 Cancelled 覆盖、以及跨整张表的
「每个分支都释放 pool key lease」「作废账单一律不提交 report」不变量。
2026-08-17 14:53:05 +08:00
AAEE86 1c5ee5228c refactor(ws): 用 ResponsesTurnState 收敛连接 turn 状态
评审第 2 条:BoundResponsesConnection 用 response_in_flight、active_turn、
active_response_create 三个可独立变化的字段编码同一件事,8 种组合里只有 3 种
合法,非法组合只能靠调用点的 if 和「记得同时改另外两个字段」来避免。

三字段合并为一个 ResponsesTurnState:

  Idle                                  没有进行中的 logical turn
  Responding { logical, attempt }        logical 与 attempt 必须同时存在
  Replanning { logical }                 attempt 已取走去结算/重绑,logical 仍在

Replanning 不是新概念:配额透明重试期间现状就处于这个状态,只是靠
Option::take 意外得到。转换只能走 begin / detach_attempt / resume / end,
response_in_flight 与「是否接受新 response.create」都由变体推导。

由此消除的运行时不变量(原来全靠调用点自觉):
- 有 attempt 必有 logical turn
- response_in_flight 与 attempt 同生共死(原来 client 写失败后
  active_turn=None 而 response_in_flight 仍为 true)
- logical turn 结束时必须清 attempt:原来 `active_response_create = None`
  在 connection.rs 里手写 13 处,漏一处就残留;现在只有 end() 一个出口
- 上游绑定返回的连接不再自带 response_in_flight=true 的半成品状态

同时删除 update_response_in_flight:Started 帧把已经是 true 的字段再设一次,
Close 帧因为没有解析出的 frame 而根本不触发,是纯冗余写;它在 Idle 态收到
Started 帧时还会把 response_in_flight 置真,从而永久阻塞后续 response.create。

行为等价。ActiveResponsesWebSocketRequest 改名 LogicalTurn 并随状态机移入
新的 turn_state.rs;状态机对 attempt 类型泛型化,测试用轻量替身驱动同一套
转换逻辑,无需 AppState 或真实 socket。
2026-08-17 14:52:57 +08:00
AAEE86 9d80281b53 fix(ws): route WS planning and continuation through PII redaction 2026-08-17 14:52:49 +08:00
AAEE86 3b036299d4 fix(ws): enforce absolute upstream handshake and initial-message deadlines 2026-08-17 14:52:39 +08:00
AAEE86 f70ae68273 fix(ws): treat max_output_tokens incomplete as legitimate terminal 2026-08-17 14:52:33 +08:00
AAEE86 1353d76e07 feat(frontend): Responses WebSocket 配置与用量展示
provider 表单支持开启 Responses WebSocket;用量列表、状态与详情
区分 WebSocket 请求。
2026-08-17 14:52:25 +08:00
AAEE86 621a528083 test(ws): Responses WebSocket 端到端套件接入 CI
补齐 aether-integration-tests 的 responses_websocket_e2e 集成测试,
并把 CI 的 scenario 任务从 --bins 改为 --bins --tests,否则该套件
不会被执行。
2026-08-17 14:51:24 +08:00
AAEE86 a498875591 feat(gateway): Responses WebSocket 连通性探针
新增 aether-codex-ws-probe 与 aether-openai-responses-ws-probe 两个
二进制,用于在不暴露凭据的前提下验证上游 WebSocket 端点可用性:凭据
只从环境变量读取,不写入日志。公共流程放在
bin/support/responses_ws_probe.rs,各 profile 只负责自己的鉴权与
请求头要求。
2026-08-17 14:51:18 +08:00
AAEE86 71b54070e8 feat(gateway): Codex/OpenAI Responses WebSocket 代理模式
在 /v1/responses 上支持 WebSocket 升级,把客户端帧中继到上游 Codex /
OpenAI Responses WebSocket 端点,同时保持既有的路由、鉴权、配额与用量
语义:

- 路由与准入:control/route/ai.rs 识别 WebSocket 升级请求;
  websocket/ingress.rs 复用 API Key 鉴权、IP 规则与并发许可,并引入
  独立的 WebSocket 连接许可
- 中继:websocket/responses/* 按 connection / session / turn 分层,
  帧解析归一化、socket 写入有界、continuation 保持调度亲和性
- 配额:orchestration/codex_quota_breaker.rs 在账号配额耗尽时熔断并
  自动恢复,不再直接断开客户端连接
- 用量:每个 turn 的终态用量落库,request_metadata 记录
  websocket_mode / websocket_transport,管理端与 usage 视图暴露
  is_websocket
- 管理端:provider 可配置 Responses WebSocket 开关
2026-08-17 14:50:33 +08:00
ZheFox 9a0d346ff3 Merge pull request #728 from zhefox/main
fix(gateway): stop candidate persistence retry storms
2026-08-17 14:28:11 +08:00
ZheFox 32944538e9 fix(gateway): stop candidate persistence retry storms 2026-08-17 13:49:12 +08:00
ZheFox 0b17026eab Merge pull request #726 from zhefox/main
fix(codex): restore upstream model discovery
2026-08-15 20:16:12 +08:00
ZheFox b13d9b9b40 fix(codex): restore upstream model discovery 2026-08-15 19:36:28 +08:00
ZheFox b7fca851b8 Merge pull request #724 from zhefox/main
fix(codex): serve versioned dynamic model catalogs
2026-08-14 19:08:53 +08:00
zhefox 810c3dfe2b fix(codex): serve versioned dynamic model catalogs 2026-08-14 18:41:44 +08:00
elky a1d64e5239 fix(routing): preserve allowlist edits and save state 2026-08-14 11:43:24 +08:00
elky fb33ea57b0 Merge PR #715: decouple routing model overrides 2026-08-14 11:10:21 +08:00
elky 5b0c763086 fix(codex): fence concurrent quota updates 2026-08-14 09:28:07 +08:00
elky f3a12c1008 fix(ai): preserve Codex image edit validation 2026-08-13 11:31:17 +08:00
elky ca35e09eaa Merge pull request #718 from zjm54321/fix/custom-image-edit-json 2026-08-13 11:00:15 +08:00
elky 8cf381b0c3 feat(codex): add OAuth fingerprint convergence 2026-08-13 09:57:17 +08:00
elky 654c4f6978 fix(auth): preserve turnstile while typing email 2026-08-12 16:56:18 +08:00
elky edb8362adc fix(provider): omit default model test temperature 2026-08-12 16:56:18 +08:00
elky 29fa4aed19 perf(gateway): raise default server pool floor 2026-08-12 16:56:18 +08:00
zjm54321 41e93858e1 refactor(ai): simplify image edit serialization 2026-08-11 00:36:33 +08:00
zjm54321 8d918d0459 fix(ai): serialize image edits with images array 2026-08-11 00:30:00 +08:00
ZheFox 3a759fae89 Merge pull request #712 from zhefox/fix/claude-code-conversion-api-key-toggle
fix: restore Claude Code conversion and user API key toggles
2026-08-05 14:53:03 +08:00
zhefox 985ff3c36a test(gateway): align claude_code endpoint reconciliation 2026-08-05 14:34:20 +08:00
zhefox 908d4f2603 style(provider): apply rustfmt to claude_code tests 2026-08-05 13:58:46 +08:00
zhefox 4d67569873 fix(gateway): support claude_code cross-format Claude messages 2026-08-05 13:57:20 +08:00
ZheFox 1aab31a148 Merge pull request #710 from zhefox/main
fix: align Responses compatibility, routing, and model permissions
2026-08-03 19:31:04 +08:00
zhefox aedff9a704 fix(provider): validate mapped model reasoning effort 2026-08-03 19:22:51 +08:00
zhefox 669f4bddc5 fix: align Responses routing and model permissions 2026-08-03 18:48:01 +08:00
zbs 1a4eede34d fix(routing): decouple model overrides from allowed scope 2026-08-03 07:52:49 +08:00
elky 0318808db9 fix(providers): allow transfer limits on creation 2026-07-31 13:35:08 +08:00
251 changed files with 44695 additions and 1982 deletions
+1 -1
View File
@@ -66,7 +66,7 @@ ADMIN_USERNAME=admin123456
# docker compose 下 app 启动前自动执行 pending migration/backfill(默认 true)
# AETHER_GATEWAY_AUTO_PREPARE_DATABASE=true
# PostgreSQL 连接池配置(默认按 CPU 自动计算;正式高并发环境可显式预算)
# PostgreSQL 连接池配置(默认每核 4 条、总池至少 32 条且最多 100 条;多实例部署应显式分配每实例预算)
# AETHER_GATEWAY_DATA_POSTGRES_MIN_CONNECTIONS=12
# AETHER_GATEWAY_DATA_POSTGRES_MAX_CONNECTIONS=80
# AETHER_GATEWAY_MAX_IN_FLIGHT_REQUESTS=2048
+4 -4
View File
@@ -200,13 +200,13 @@ jobs:
RUSTFLAGS: "-C link-arg=-fuse-ld=mold"
run: cargo nextest run -p aether-gateway --lib
- name: Test bin
- name: Test bins
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
RUST_MIN_STACK: "16777216"
RUSTFLAGS: "-C link-arg=-fuse-ld=mold"
run: cargo nextest run -p aether-gateway --bin aether-gateway
run: cargo nextest run -p aether-gateway --bins
- name: Show sccache stats
if: always()
@@ -387,11 +387,11 @@ jobs:
- name: Setup sccache
uses: mozilla-actions/[email protected]
- name: Test scenario binaries
- name: Test scenario binaries and end-to-end suites
env:
RUSTC_WRAPPER: sccache
SCCACHE_GHA_ENABLED: "true"
run: cargo test -p aether-integration-tests --bins
run: cargo test -p aether-integration-tests --bins --tests
- name: Show sccache stats
if: always()
Generated
+4
View File
@@ -325,6 +325,7 @@ dependencies = [
"axum",
"base64 0.22.1",
"bcrypt",
"brotli",
"bytes",
"chrono",
"chrono-tz",
@@ -346,6 +347,7 @@ dependencies = [
"reqwest",
"rsa",
"rustls 0.23.37",
"semver",
"serde",
"serde_json",
"sha1",
@@ -441,6 +443,7 @@ name = "aether-integration-tests"
version = "0.1.0"
dependencies = [
"aether-contracts",
"aether-crypto",
"aether-data",
"aether-data-contracts",
"aether-gateway",
@@ -457,6 +460,7 @@ dependencies = [
"sqlx",
"tokio",
"tokio-tungstenite 0.28.0",
"uuid",
]
[[package]]
+2 -1
View File
@@ -106,6 +106,7 @@ async-trait = "0.1"
axum = "0.8"
base64 = "0.22"
bcrypt = "0.16"
brotli = "8"
bytes = "1"
cbc = "0.1"
chrono = { version = "0.4", features = ["serde"] }
@@ -135,7 +136,7 @@ tokio = { version = "1", features = ["macros", "net", "rt-multi-thread", "signal
tokio-util = { version = "0.7", features = ["codec", "io-util"] }
tracing = "0.1"
tracing-subscriber = { version = "0.3", features = ["env-filter", "json"] }
uuid = { version = "1", features = ["serde", "v4", "v5"] }
uuid = { version = "1", features = ["serde", "v4", "v5", "v7"] }
webpki-roots = "0.26"
wreq = { version = "6.0.0-rc.28", default-features = false, features = ["json", "stream", "socks", "webpki-roots", "ws"] }
wreq-util = "3.0.0-rc.10"
+3 -1
View File
@@ -137,12 +137,14 @@ Aether Tunnel 是配套的正向代理节点,部署在海外 VPS 上,为墙
- Embeddings: [OpenAI compatible `POST /v1/embeddings`](docs/api/embeddings.md)
- Rerank: [OpenAI/Jina compatible `POST /v1/rerank`](docs/api/rerank.md)
- Responses WebSocket mode: [protocol and Aether behavior](docs/WebSocket-Mode.md)
- WebSocket probes: [Codex](docs/operations/codex-responses-websocket-probe.md) · [OpenAI Responses](docs/operations/openai-responses-websocket-probe.md)
## 环境变量
- `APP_PORT`:`aether-gateway` 唯一监听端口,固定绑定 `0.0.0.0:${APP_PORT}`
- `DATABASE_URL`:数据库连接串;SQLite 例如 `sqlite:///opt/aether/data/aether.db`,Postgres 例如 `postgresql://postgres:aether@postgres:5432/aether`
- `AETHER_GATEWAY_DATA_POSTGRES_MIN_CONNECTIONS` / `AETHER_GATEWAY_DATA_POSTGRES_MAX_CONNECTIONS`:数据库连接池手动覆盖值;未配置时会自动推导,SQLite 固定 `1/1`,Postgres/MySQL 按 CPU 核心数计算并默认封顶 `100`
- `AETHER_GATEWAY_DATA_POSTGRES_MIN_CONNECTIONS` / `AETHER_GATEWAY_DATA_POSTGRES_MAX_CONNECTIONS`:数据库连接池手动覆盖值;未配置时 SQLite 固定 `1/1`,Postgres/MySQL 按每核 `4` 条自动推导,总池范围为 `32-100`。该预算按进程计算,多实例部署应按数据库连接上限显式分配
- `AETHER_GATEWAY_MAX_IN_FLIGHT_REQUESTS`:单实例请求并发上限;未配置时按 CPU 自动推导(基础范围 `512-65536`),低文件描述符预算时会进一步下调
- `AETHER_GATEWAY_REQUEST_BODY_BUFFER_BUDGET_MB`:单实例同时读取和解压请求体的加权内存预算,默认 `256MB`
- `AETHER_GATEWAY_REQUEST_BODY_READ_TIMEOUT_MS`:请求体完整读取超时,默认 `120000ms`
+2
View File
@@ -52,6 +52,7 @@ async-trait.workspace = true
axum = { version = "0.8", features = ["ws"] }
base64.workspace = true
bcrypt.workspace = true
brotli.workspace = true
bytes.workspace = true
chrono.workspace = true
chrono-tz.workspace = true
@@ -75,6 +76,7 @@ rsa = "0.9.10"
rustls.workspace = true
serde.workspace = true
serde_json.workspace = true
semver.workspace = true
sha1 = "0.10"
sha2 = { workspace = true, features = ["oid"] }
socket2.workspace = true
@@ -69,6 +69,9 @@ pub(crate) use aether_ai_formats::api::{
};
pub(crate) use aether_ai_formats::protocol::stream::CanonicalUsage as StreamingCanonicalUsage;
pub(crate) use aether_ai_formats::CODEX_RESPONSES_LITE_HEADER;
/// Codex client identity headers re-exported for out-of-crate probe binaries,
/// which must reach `aether_ai_formats` through this seam.
pub use aether_ai_formats::{CODEX_CLIENT_ORIGINATOR, CODEX_CLIENT_USER_AGENT};
pub(crate) fn parse_direct_request_body(
parts: &http::request::Parts,
+7 -5
View File
@@ -52,17 +52,19 @@ pub(crate) use self::planner::{
build_standard_family_sync_plan_and_reports, build_standard_stream_plan_from_decision,
build_standard_sync_plan_from_decision, candidate_auth_channel_skip_reason,
codex_model_capabilities_for_transport, extract_pool_sticky_session_token,
maybe_build_stream_decision_payload, maybe_build_stream_plan_payload,
maybe_build_sync_decision_payload, maybe_build_sync_plan_payload,
planner_is_matching_stream_request, provider_key_pool_score_id, provider_key_pool_score_scope,
read_candidate_transport_snapshot, record_local_runtime_candidate_skip_reason,
maybe_build_responses_websocket_decision, maybe_build_stream_decision_payload,
maybe_build_stream_plan_payload, maybe_build_sync_decision_payload,
maybe_build_sync_plan_payload, planner_is_matching_stream_request, provider_key_pool_score_id,
provider_key_pool_score_scope, read_candidate_transport_snapshot,
record_local_runtime_candidate_skip_reason, resolve_provider_chat_pii_redaction,
resolve_tunnel_scheduler_affinity_context, resolve_upstream_is_stream_for_provider,
set_local_openai_chat_execution_exhausted_diagnostic,
set_local_openai_image_execution_exhausted_diagnostic, validate_final_openai_provider_request,
CandidateFailureDiagnostic, CandidateFailureDiagnosticKind, EligibleLocalExecutionCandidate,
GatewayAuthApiKeySnapshot, GatewayProviderTransportSnapshot, LocalExecutionAttemptSource,
LocalExecutionCandidateKind, LocalResolvedOAuthRequestAuth, PlannerAppState,
SkippedLocalExecutionCandidate,
ResponsesWebSocketBodyNormalization, ResponsesWebSocketDecision,
ResponsesWebSocketPinnedCandidate, SkippedLocalExecutionCandidate,
};
pub(crate) use self::pure::*;
pub(crate) use self::response_history::{
@@ -188,9 +188,14 @@ pub(crate) async fn scheduler_ordering_config_for_routing_policy(
state: PlannerAppState<'_>,
routing_policy: Option<&ResolvedRoutingPolicy>,
) -> SchedulerOrderingConfig {
let system_config = read_scheduler_ordering_config_or_default(state).await;
match routing_policy {
Some(policy) => scheduler_ordering_config_from_routing_policy(policy),
None => read_scheduler_ordering_config_or_default(state).await,
Some(policy) => {
let mut config = scheduler_ordering_config_from_routing_policy(policy);
config.keep_priority_on_conversion |= system_config.keep_priority_on_conversion;
config
}
None => system_config,
}
}
@@ -390,6 +395,43 @@ mod tests {
assert_eq!(overlaid.key_global_priority_for_format, Some(2));
}
#[tokio::test]
async fn routing_policy_inherits_global_conversion_priority_override() {
let data_state = GatewayDataState::default().with_system_config_values_for_tests([(
"keep_priority_on_conversion".to_string(),
json!(true),
)]);
let state = AppState::new()
.expect("state should build")
.with_data_state_for_tests(data_state);
let policy = aether_routing_core::ResolvedRoutingPolicy {
group_id: Some("group-1".to_string()),
group_version: Some(1),
selection_source: "system_default".to_string(),
requested_model: "gpt-5.4-mini".to_string(),
resolved_model: "gpt-5.4-mini".to_string(),
priority_mode: aether_routing_core::RoutingSetPriorityMode::Provider,
scheduling_mode: aether_routing_core::RoutingSchedulingMode::FixedOrder,
keep_priority_on_conversion: false,
ranking_overlay: Default::default(),
mutation_plan: Default::default(),
pool_policy_overrides: Default::default(),
matched_rules: Vec::new(),
};
let ordering = super::scheduler_ordering_config_for_routing_policy(
PlannerAppState::new(&state),
Some(&policy),
)
.await;
assert_eq!(
ordering.scheduling_mode,
crate::scheduler::config::SchedulerSchedulingMode::FixedOrder
);
assert!(ordering.keep_priority_on_conversion);
}
#[test]
fn routing_policy_uses_pool_priority_for_pool_group_global_key_slot() {
let mut candidate = sample_candidate("endpoint-1", "representative-key");
@@ -2540,4 +2540,146 @@ mod tests {
vec!["provider-openai-responses-regular"]
);
}
#[tokio::test]
async fn fixed_order_prefers_codex_responses_when_conversion_keeps_priority() {
let mut codex = standard_candidate_row("provider-codex", "openai:responses", 0);
codex.provider_type = "codex".to_string();
codex.key_auth_type = "oauth".to_string();
codex.global_model_name = "gpt-5.4-mini".to_string();
codex.model_provider_model_name = "gpt-5.4-mini".to_string();
let mut custom_chat = standard_candidate_row("provider-custom", "openai:chat", 10);
custom_chat.global_model_name = "gpt-5.4-mini".to_string();
custom_chat.model_provider_model_name = "gpt-5.4-mini".to_string();
let candidate_repository: Arc<dyn MinimalCandidateSelectionReadRepository> =
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed([
codex.clone(),
custom_chat.clone(),
]));
let catalog_items = [
provider_catalog_for_standard_row(&codex, false),
provider_catalog_for_standard_row(&custom_chat, false),
];
let provider_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
catalog_items
.iter()
.map(|(provider, _, _)| provider.clone())
.collect(),
catalog_items
.iter()
.map(|(_, endpoint, _)| endpoint.clone())
.collect(),
catalog_items
.iter()
.map(|(_, _, key)| key.clone())
.collect(),
));
let data_state =
GatewayDataState::with_provider_catalog_and_minimal_candidate_selection_for_tests(
provider_repository,
candidate_repository,
)
.with_encryption_key_for_tests("development-key")
.with_system_config_values_for_tests([
(
"scheduling_mode".to_string(),
serde_json::json!("fixed_order"),
),
(
"keep_priority_on_conversion".to_string(),
serde_json::json!(true),
),
]);
let app = AppState::new()
.expect("gateway state should build")
.with_data_state_for_tests(data_state);
let auth_snapshot = unrestricted_auth_snapshot();
let model_directive_policy =
crate::system_features::ModelDirectivePolicySnapshot::load(&app).await;
let routing_policy = ResolvedRoutingPolicy {
group_id: Some("routing-group-codex-first".to_string()),
group_version: Some(1),
selection_source: "test".to_string(),
requested_model: "gpt-5.4-mini".to_string(),
resolved_model: "gpt-5.4-mini".to_string(),
priority_mode: aether_routing_core::RoutingSetPriorityMode::Provider,
scheduling_mode: aether_routing_core::RoutingSchedulingMode::FixedOrder,
keep_priority_on_conversion: false,
ranking_overlay: Default::default(),
mutation_plan: Default::default(),
pool_policy_overrides: Default::default(),
matched_rules: Vec::new(),
};
let mut cursor = LocalCandidatePreselectionPageCursor::new(
PlannerAppState::new(&app),
&model_directive_policy,
"openai:chat",
"gpt-5.4-mini",
None,
false,
None,
&auth_snapshot,
Some(&routing_policy),
None,
None,
true,
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
false,
None,
)
.await;
let first_page = cursor
.next_page()
.await
.expect("preselection should succeed")
.expect("Codex and custom candidates should share the priority page");
assert!(
first_page.skipped_candidates.is_empty(),
"priority page unexpectedly skipped candidates: {:?}",
first_page
.skipped_candidates
.iter()
.map(|candidate| {
(
candidate.candidate.provider_id.as_str(),
candidate.skip_reason,
)
})
.collect::<Vec<_>>()
);
assert_eq!(
first_page
.candidates
.iter()
.map(|candidate| candidate.provider_id.as_str())
.collect::<Vec<_>>(),
vec!["provider-custom", "provider-codex"]
);
let (ranked, skipped) =
super::super::candidate_resolution::resolve_and_rank_logical_local_execution_candidates(
PlannerAppState::new(&app),
first_page.candidates,
"openai:chat",
Some("gpt-5.4-mini"),
Some(&auth_snapshot),
None,
None,
Some(&routing_policy),
None,
None,
aether_ai_serving::AiCandidateResolutionMode::Standard,
)
.await;
assert!(skipped.is_empty());
assert_eq!(
ranked
.iter()
.map(|candidate| candidate.candidate.provider_id.as_str())
.collect::<Vec<_>>(),
vec!["provider-codex", "provider-custom"]
);
}
}
@@ -55,6 +55,7 @@ pub(crate) struct LocalRequestedModelDecisionInput {
pub(crate) client_surface: Option<ClientSurface>,
pub(crate) gateway_credential_carrier: Option<GatewayCredentialCarrier>,
pub(crate) client_session_affinity: Option<ClientSessionAffinity>,
pub(crate) original_client_session_id: Option<String>,
pub(crate) routing_policy: Option<ResolvedRoutingPolicy>,
pub(crate) routing_trace_seed: Option<RoutingDecisionTrace>,
pub(crate) routing_context: Option<LocalRoutingRequestContext>,
@@ -150,6 +151,12 @@ pub(crate) fn apply_provider_request_routing_policy_to_decision(
provider_api_format.as_str(),
);
}
apply_codex_oauth_fingerprint_convergence_to_decision(
input,
decision,
transport,
provider_api_format.as_str(),
);
return Ok(());
};
let provider_body_rules = decision
@@ -207,6 +214,12 @@ pub(crate) fn apply_provider_request_routing_policy_to_decision(
provider_api_format.as_str(),
);
}
apply_codex_oauth_fingerprint_convergence_to_decision(
input,
decision,
transport,
provider_api_format.as_str(),
);
return Ok(());
}
if original_provider_request_body.is_none() && !policy.mutation_plan.body_patch.is_empty() {
@@ -305,13 +318,40 @@ pub(crate) fn apply_provider_request_routing_policy_to_decision(
if original_provider_request_body.is_some() {
decision.provider_request_body = Some(provider_request_body);
}
apply_codex_oauth_fingerprint_convergence_to_decision(
input,
decision,
transport,
provider_api_format.as_str(),
);
update_report_context_provider_request_mutation(decision, &policy);
Ok(())
}
fn apply_codex_oauth_fingerprint_convergence_to_decision(
input: &LocalRequestedModelDecisionInput,
decision: &mut AiExecutionDecision,
transport: Option<&GatewayProviderTransportSnapshot>,
provider_api_format: &str,
) {
let (Some(transport), Some(provider_request_body)) =
(transport, decision.provider_request_body.as_mut())
else {
return;
};
crate::ai_serving::transport::apply_codex_oauth_fingerprint_convergence(
transport,
provider_api_format,
input.original_client_session_id.as_deref(),
&mut decision.provider_request_headers,
provider_request_body,
);
}
struct GatewayAuthenticatedDecisionInputPort<'a> {
state: PlannerAppState<'a>,
now_unix_secs: u64,
auth_snapshot_override: Option<GatewayAuthApiKeySnapshot>,
model_directive_policy: &'a crate::system_features::ModelDirectivePolicySnapshot,
model_directive_base_model: Option<String>,
}
@@ -328,6 +368,17 @@ impl AiAuthenticatedDecisionInputPort for GatewayAuthenticatedDecisionInputPort<
&self,
auth_context: &Self::AuthContext,
) -> Result<Option<Self::AuthSnapshot>, Self::Error> {
if let Some(snapshot) = self.auth_snapshot_override.as_ref() {
if snapshot.user_id != auth_context.user_id
|| snapshot.api_key_id != auth_context.api_key_id
{
return Err(GatewayError::Internal(
"WebSocket auth snapshot identity does not match its control decision"
.to_string(),
));
}
return Ok(Some(snapshot.clone()));
}
self.state
.read_auth_api_key_snapshot(
&auth_context.user_id,
@@ -383,6 +434,7 @@ pub(crate) fn build_local_requested_model_decision_input(
client_surface: None,
gateway_credential_carrier: None,
client_session_affinity: None,
original_client_session_id: None,
routing_policy: None,
routing_trace_seed: None,
routing_context: None,
@@ -397,6 +449,8 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input(
body_json: &Value,
client_api_format: &str,
) -> Result<(), GatewayError> {
input.original_client_session_id = routing_header_value_str(&parts.headers, "session-id")
.or_else(|| routing_header_value_str(&parts.headers, "session_id"));
let explicit_group = routing_header_value_str(&parts.headers, ROUTING_GROUP_HEADER);
let selected_group = match state.routing_group_read_repository() {
Some(repository) => {
@@ -708,6 +762,27 @@ pub(crate) async fn resolve_local_authenticated_decision_input(
requested_model_api_format: Option<&str>,
explicit_required_capabilities: Option<&serde_json::Value>,
model_directive_policy: &crate::system_features::ModelDirectivePolicySnapshot,
) -> Result<Option<ResolvedLocalDecisionAuthInput>, GatewayError> {
resolve_local_authenticated_decision_input_with_snapshot(
state,
auth_context,
None,
requested_model,
requested_model_api_format,
explicit_required_capabilities,
model_directive_policy,
)
.await
}
pub(crate) async fn resolve_local_authenticated_decision_input_with_snapshot(
state: &AppState,
auth_context: ExecutionRuntimeAuthContext,
auth_snapshot_override: Option<GatewayAuthApiKeySnapshot>,
requested_model: Option<&str>,
requested_model_api_format: Option<&str>,
explicit_required_capabilities: Option<&serde_json::Value>,
model_directive_policy: &crate::system_features::ModelDirectivePolicySnapshot,
) -> Result<Option<ResolvedLocalDecisionAuthInput>, GatewayError> {
let model_directive_base_model = match (requested_model, requested_model_api_format) {
(Some(model), Some(api_format)) => model_directive_policy
@@ -719,6 +794,7 @@ pub(crate) async fn resolve_local_authenticated_decision_input(
let port = GatewayAuthenticatedDecisionInputPort {
state: PlannerAppState::new(state),
now_unix_secs: current_unix_secs(),
auth_snapshot_override,
model_directive_policy,
model_directive_base_model,
};
@@ -1024,6 +1100,52 @@ mod tests {
}
}
#[tokio::test]
async fn explicit_auth_snapshot_override_does_not_fall_back_to_the_planner_cache() {
// AppState::new has no auth snapshot repository. Without the explicit
// override this resolver returns None; a WebSocket strong snapshot must
// therefore be the exact value used to build the planner input.
let state = AppState::new().expect("test state should build");
let mut strong_snapshot = sample_auth_snapshot();
strong_snapshot.api_key_allowed_models = Some(vec!["gpt-live-only".to_string()]);
let resolved = resolve_local_authenticated_decision_input_with_snapshot(
&state,
sample_auth_context(),
Some(strong_snapshot.clone()),
Some("gpt-live-only"),
Some("openai:responses"),
None,
&Default::default(),
)
.await
.expect("snapshot override should resolve")
.expect("the explicit snapshot should replace the missing cached value");
assert_eq!(resolved.auth_snapshot, strong_snapshot);
}
#[tokio::test]
async fn explicit_auth_snapshot_override_rejects_an_identity_mismatch() {
let state = AppState::new().expect("test state should build");
let mut wrong_snapshot = sample_auth_snapshot();
wrong_snapshot.api_key_id = "another-key".to_string();
let error = resolve_local_authenticated_decision_input_with_snapshot(
&state,
sample_auth_context(),
Some(wrong_snapshot),
Some("gpt-live-only"),
Some("openai:responses"),
None,
&Default::default(),
)
.await
.expect_err("a snapshot for another API key must never be injected");
assert!(matches!(error, GatewayError::Internal(message) if message.contains("identity")));
}
#[tokio::test]
async fn explicit_routing_attachment_authorizes_and_caches_per_principal() {
let repository = Arc::new(InMemoryRoutingGroupRepository::default());
@@ -1154,6 +1276,7 @@ mod tests {
client_surface: None,
gateway_credential_carrier: None,
client_session_affinity: None,
original_client_session_id: None,
routing_policy: None,
routing_trace_seed: None,
model_directive_policy: Default::default(),
@@ -1312,6 +1435,57 @@ mod tests {
}
}
fn sample_codex_fingerprint_transport() -> GatewayProviderTransportSnapshot {
let mut transport = sample_codex_transport_with_card();
transport.provider.config = Some(json!({
"codex": {"fingerprint_convergence_enabled": true}
}));
transport.endpoint.api_format = "openai:responses".to_string();
transport.endpoint.endpoint_kind = Some("responses".to_string());
transport.key.api_formats = Some(vec!["openai:responses".to_string()]);
transport.key.decrypted_auth_config =
Some(json!({"account_id": "account-codex-1"}).to_string());
transport
}
fn sample_codex_fingerprint_decision() -> AiExecutionDecision {
let prompt_cache_key = "172c39e6-c0a0-5a70-8b63-e0f8e0d185a3";
let mut decision = sample_decision();
decision.provider_type = Some("codex".to_string());
decision.provider_api_format = Some("openai:responses".to_string());
decision.client_api_format = Some("openai:responses".to_string());
decision.provider_request_headers.extend([
("session-id".to_string(), "spoofed-session".to_string()),
("thread-id".to_string(), "spoofed-thread".to_string()),
(
"x-codex-turn-metadata".to_string(),
json!({
"installation_id": "spoofed-installation",
"session_id": "spoofed-session",
"thread_source": "cli"
})
.to_string(),
),
]);
decision.provider_request_body = Some(json!({
"model": "gpt-5",
"input": [],
"metadata": {},
"prompt_cache_key": prompt_cache_key,
"client_metadata": {
"session_id": "spoofed-session",
"thread_id": "spoofed-thread",
"caller": "sdk",
"x-codex-turn-metadata": json!({
"installation_id": "spoofed-installation",
"session_id": "spoofed-session",
"sandbox": "workspace-write"
}).to_string()
}
}));
decision
}
fn set_provider_request_rules(
input: &mut LocalRequestedModelDecisionInput,
allowed_models: &[&str],
@@ -1351,6 +1525,7 @@ mod tests {
client_surface: None,
gateway_credential_carrier: None,
client_session_affinity: None,
original_client_session_id: None,
routing_policy: None,
routing_trace_seed: None,
model_directive_policy: Default::default(),
@@ -1420,6 +1595,7 @@ mod tests {
client_surface: None,
gateway_credential_carrier: None,
client_session_affinity: None,
original_client_session_id: None,
routing_policy: None,
routing_trace_seed: None,
routing_context: None,
@@ -1488,6 +1664,124 @@ mod tests {
);
}
#[test]
fn codex_fingerprint_convergence_runs_at_every_provider_routing_success_exit() {
let transport = sample_codex_fingerprint_transport();
let mut no_context = sample_decision_input();
no_context.routing_context = None;
let mut empty_mutation = sample_decision_input();
empty_mutation
.routing_context
.as_mut()
.expect("routing context")
.group_config_json = json!({
"allowed_models": ["gpt-5"],
"rules": []
});
let mut with_mutation = sample_decision_input();
for input in [&mut no_context, &mut empty_mutation, &mut with_mutation] {
input.original_client_session_id = Some("client-session-1".to_string());
}
let mut stable_identity = None;
let mut turn_ids = std::collections::BTreeSet::new();
for (exit_name, input) in [
("no_context", no_context),
("empty_mutation", empty_mutation),
("with_mutation", with_mutation),
] {
let mut decision = sample_codex_fingerprint_decision();
apply_provider_request_routing_policy_to_decision(
&input,
&mut decision,
Some(&transport),
)
.unwrap_or_else(|error| panic!("{exit_name} should converge: {error:?}"));
let session_id = decision.provider_request_headers["session-id"].clone();
let thread_id = decision.provider_request_headers["thread-id"].clone();
let installation_id =
decision.provider_request_headers["x-codex-installation-id"].clone();
let window_id = decision.provider_request_headers["x-codex-window-id"].clone();
assert_eq!(decision.provider_request_headers["session_id"], session_id);
assert_eq!(
decision.provider_request_headers["x-client-request-id"],
thread_id
);
assert_eq!(window_id, format!("{thread_id}:0"));
assert_eq!(
uuid::Uuid::parse_str(&session_id)
.expect("session UUID")
.get_version_num(),
4
);
assert_eq!(
uuid::Uuid::parse_str(&thread_id)
.expect("thread UUID")
.get_version_num(),
4
);
let body = decision
.provider_request_body
.as_ref()
.expect("request body");
assert_eq!(
body["prompt_cache_key"],
"172c39e6-c0a0-5a70-8b63-e0f8e0d185a3"
);
assert_eq!(body["client_metadata"]["session_id"], session_id);
assert_eq!(body["client_metadata"]["thread_id"], thread_id);
assert_eq!(body["client_metadata"]["caller"], "sdk");
assert_eq!(
body["client_metadata"]["x-codex-installation-id"],
installation_id
);
assert_eq!(body["client_metadata"]["x-codex-window-id"], window_id);
let header_metadata: Value =
serde_json::from_str(&decision.provider_request_headers["x-codex-turn-metadata"])
.expect("header turn metadata");
let body_metadata: Value = serde_json::from_str(
body["client_metadata"]["x-codex-turn-metadata"]
.as_str()
.expect("embedded turn metadata"),
)
.expect("embedded turn metadata JSON");
assert_eq!(header_metadata["thread_source"], "cli");
assert_eq!(body_metadata["sandbox"], "workspace-write");
assert_eq!(
header_metadata["turn_id"],
body["client_metadata"]["turn_id"]
);
assert_eq!(body_metadata["turn_id"], body["client_metadata"]["turn_id"]);
assert_eq!(
header_metadata["turn_started_at_unix_ms"],
body_metadata["turn_started_at_unix_ms"]
);
let turn_id = body["client_metadata"]["turn_id"]
.as_str()
.expect("turn ID")
.to_string();
assert_eq!(
uuid::Uuid::parse_str(&turn_id)
.expect("turn UUID")
.get_version_num(),
7
);
turn_ids.insert(turn_id);
let identity = (installation_id, session_id, thread_id);
if let Some(expected) = stable_identity.as_ref() {
assert_eq!(&identity, expected, "identity changed at {exit_name}");
} else {
stable_identity = Some(identity);
}
}
assert_eq!(turn_ids.len(), 3, "each request needs a fresh turn ID");
}
#[test]
fn provider_request_routing_policy_cannot_restore_credentials_or_aether_internal_headers() {
for header_name in [
@@ -49,6 +49,7 @@ pub(crate) use self::plan_builders::{
pub(crate) use self::pool_scores::{
build_provider_key_pool_score_upsert, provider_key_pool_score_id, provider_key_pool_score_scope,
};
pub(crate) use self::redaction::resolve_provider_chat_pii_redaction;
pub(crate) use self::request_gzip::resolve_transport_request_encoding_policy;
pub(crate) use self::route::is_matching_stream_request as planner_is_matching_stream_request;
pub(crate) use self::runtime_miss::{
@@ -80,8 +81,10 @@ pub(crate) use self::standard::{
build_local_stream_plan_and_reports as build_standard_family_stream_plan_and_reports,
build_local_sync_attempt_source as build_standard_family_sync_attempt_source,
build_local_sync_plan_and_reports as build_standard_family_sync_plan_and_reports,
codex_model_capabilities_for_transport, set_local_openai_chat_execution_exhausted_diagnostic,
validate_final_openai_provider_request,
codex_model_capabilities_for_transport, maybe_build_responses_websocket_decision,
set_local_openai_chat_execution_exhausted_diagnostic, validate_final_openai_provider_request,
ResponsesWebSocketBodyNormalization, ResponsesWebSocketDecision,
ResponsesWebSocketPinnedCandidate,
};
pub(crate) use self::state::{
GatewayAuthApiKeySnapshot, GatewayProviderTransportSnapshot, LocalResolvedOAuthRequestAuth,
@@ -167,20 +167,21 @@ pub(crate) async fn materialize_local_same_format_provider_candidate_attempts(
.collect(),
LocalCandidateResolutionMode::Standard,
|eligible| {
let provider_api_format = eligible.provider_api_format.clone();
let (execution_strategy, conversion_mode) = ai_local_execution_contract_for_formats(
spec_metadata.api_format,
spec_metadata.api_format,
&provider_api_format,
);
Some(build_local_execution_candidate_contract_metadata(
LocalExecutionCandidateMetadataParts {
eligible,
provider_api_format: spec_metadata.api_format,
provider_api_format: provider_api_format.as_str(),
client_api_format: spec_metadata.api_format,
extra_fields: serde_json::Map::new(),
},
execution_strategy,
conversion_mode,
spec_metadata.api_format,
provider_api_format.as_str(),
))
},
|mut skipped_candidate| {
@@ -273,20 +274,21 @@ pub(crate) async fn build_local_same_format_provider_candidate_attempt_source<'a
.collect(),
LocalCandidateResolutionMode::Standard,
|eligible| {
let provider_api_format = eligible.provider_api_format.clone();
let (execution_strategy, conversion_mode) = ai_local_execution_contract_for_formats(
spec_metadata.api_format,
spec_metadata.api_format,
&provider_api_format,
);
Some(build_local_execution_candidate_contract_metadata(
LocalExecutionCandidateMetadataParts {
eligible,
provider_api_format: spec_metadata.api_format,
provider_api_format: provider_api_format.as_str(),
client_api_format: spec_metadata.api_format,
extra_fields: serde_json::Map::new(),
},
execution_strategy,
conversion_mode,
spec_metadata.api_format,
provider_api_format.as_str(),
))
},
|mut skipped_candidate| {
@@ -55,8 +55,6 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
..
} = &attempt;
let candidate = &eligible.candidate;
let (execution_strategy, conversion_mode) =
ai_local_execution_contract_for_formats(spec_metadata.api_format, spec_metadata.api_format);
let Some(resolved) = resolve_local_same_format_provider_candidate_payload_parts(
state, parts, trace_id, body_json, input, &attempt, spec,
)
@@ -164,6 +162,10 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
}
}
let provider_api_format = resolved.provider_api_format.clone();
let (execution_strategy, conversion_mode) = ai_local_execution_contract_for_formats(
spec_metadata.api_format,
provider_api_format.as_str(),
);
let effective_headers = input.effective_headers(&parts.headers);
let report_context = append_local_failover_policy_to_value(
append_execution_contract_fields_to_value(
@@ -207,7 +209,10 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
.unwrap_or(false),
upstream_is_stream: resolved.upstream_is_stream,
has_envelope: resolved.is_kiro || resolved.is_antigravity || resolved.is_gemini_cli,
needs_conversion: false,
needs_conversion: matches!(
conversion_mode,
crate::ai_serving::ConversionMode::Bidirectional
),
extra_fields,
}),
execution_strategy,
@@ -377,6 +377,7 @@ mod tests {
client_surface: None,
gateway_credential_carrier: None,
client_session_affinity: None,
original_client_session_id: None,
routing_policy: None,
routing_trace_seed: None,
routing_context: None,
@@ -1314,21 +1314,19 @@ async fn resolve_local_gemini_image_to_openai_image_candidate_payload_parts(
.as_object_mut()?
.insert("stream".to_string(), Value::Bool(true));
}
provider_request_body = project_openai_image_api_request_body(
&provider_request_body,
&prepared_candidate.mapped_model,
converted.operation,
crate::image_capabilities::openai_image_provider_max_generation_count_for_model(
transport.provider.provider_type.as_str(),
Some(prepared_candidate.mapped_model.as_str()),
),
)?;
if is_codex {
provider_request_body = project_codex_openai_image_api_request_body(
provider_request_body = if is_codex {
project_codex_openai_image_api_request_body(&provider_request_body, converted.operation)?
} else {
project_openai_image_api_request_body(
&provider_request_body,
&prepared_candidate.mapped_model,
converted.operation,
)?;
}
crate::image_capabilities::openai_image_provider_max_generation_count_for_model(
transport.provider.provider_type.as_str(),
Some(prepared_candidate.mapped_model.as_str()),
),
)?
};
let request_path = match converted.operation {
OpenAiImageOperation::Generate => "/v1/images/generations",
OpenAiImageOperation::Edit => "/v1/images/edits",
@@ -42,13 +42,15 @@ pub(crate) use self::openai::{
build_local_openai_responses_sync_attempt_source_for_kind,
build_local_openai_responses_sync_plan_and_reports_for_kind, copy_request_number_field,
copy_request_number_field_as, map_openai_reasoning_effort_to_claude_output,
map_openai_reasoning_effort_to_gemini_budget, maybe_build_stream_local_decision_payload,
map_openai_reasoning_effort_to_gemini_budget, maybe_build_responses_websocket_decision,
maybe_build_stream_local_decision_payload,
maybe_build_stream_local_openai_responses_decision_payload,
maybe_build_sync_local_decision_payload,
maybe_build_sync_local_openai_embedding_decision_payload,
maybe_build_sync_local_openai_responses_decision_payload, parse_openai_stop_sequences,
resolve_openai_chat_max_tokens, set_local_openai_chat_execution_exhausted_diagnostic,
value_as_u64,
value_as_u64, ResponsesWebSocketBodyNormalization, ResponsesWebSocketDecision,
ResponsesWebSocketPinnedCandidate,
};
pub(crate) use crate::ai_serving::normalize_standard_request_to_openai_chat_request;
pub(crate) use crate::ai_serving::{
@@ -352,7 +352,7 @@ fn final_openai_provider_contract_uses_the_mapped_model_for_reasoning() {
false,
)
.is_some());
assert!(build_local_openai_responses_request_body(
let remapped = build_local_openai_responses_request_body(
&alias,
"gpt-5.4",
false,
@@ -364,14 +364,15 @@ fn final_openai_provider_contract_uses_the_mapped_model_for_reasoning() {
&http::HeaderMap::new(),
false,
)
.is_none());
.expect("explicit reasoning effort should pass through to the mapped model");
assert_eq!(remapped["reasoning"]["effort"], "max");
let minimal = json!({
"model": "deployment-alias",
"messages": [{"role": "user", "content": "hello"}],
"reasoning_effort": "minimal"
});
assert!(build_local_openai_chat_request_body(
let minimal = build_local_openai_chat_request_body(
&minimal,
"gpt-5.6-terra",
false,
@@ -380,7 +381,8 @@ fn final_openai_provider_contract_uses_the_mapped_model_for_reasoning() {
&http::HeaderMap::new(),
false,
)
.is_none());
.expect("explicit chat reasoning effort should be validated by the upstream");
assert_eq!(minimal["reasoning_effort"], "minimal");
let opaque_mapping = json!({
"model": "gpt-5.6-sol-max",
@@ -426,7 +428,7 @@ fn final_openai_provider_contract_validates_body_rule_output() {
let model_override = json!([
{"action":"set","path":"model","value":"gpt-5.4"}
]);
assert!(build_local_openai_responses_request_body(
let provider_request = build_local_openai_responses_request_body(
&body,
"gpt-5.6-sol",
false,
@@ -438,7 +440,9 @@ fn final_openai_provider_contract_validates_body_rule_output() {
&http::HeaderMap::new(),
false,
)
.is_none());
.expect("body rule output should preserve explicit reasoning effort");
assert_eq!(provider_request["model"], "gpt-5.4");
assert_eq!(provider_request["reasoning"]["effort"], "max");
let cache_override = json!([
{"action":"set","path":"prompt_cache_options.ttl","value":"1h"}
@@ -1417,38 +1417,20 @@ async fn resolve_openai_chat_to_openai_image_payload_parts(
return Ok(None);
};
if !is_chatgpt_web {
let Some(projected) = project_openai_image_api_request_body(
&provider_request_body,
&prepared_candidate.mapped_model,
operation,
crate::image_capabilities::openai_image_provider_max_generation_count_for_model(
transport.provider.provider_type.as_str(),
Some(prepared_candidate.mapped_model.as_str()),
),
) else {
mark_skipped_local_openai_chat_candidate_with_extra_data(
state,
input,
trace_id,
candidate,
candidate_index,
candidate_id,
"provider_request_body_build_failed",
request_body_build_failure_extra_data(
body_json,
"openai:chat",
provider_api_format,
let projected = if is_codex {
project_codex_openai_image_api_request_body(&provider_request_body, operation)
} else {
project_openai_image_api_request_body(
&provider_request_body,
&prepared_candidate.mapped_model,
operation,
crate::image_capabilities::openai_image_provider_max_generation_count_for_model(
transport.provider.provider_type.as_str(),
Some(prepared_candidate.mapped_model.as_str()),
),
)
.await;
return Ok(None);
};
provider_request_body = projected;
}
if is_codex {
let Some(projected) =
project_codex_openai_image_api_request_body(&provider_request_body, operation)
else {
let Some(projected) = projected else {
mark_skipped_local_openai_chat_candidate_with_extra_data(
state,
input,
@@ -2195,6 +2177,7 @@ mod tests {
client_surface: None,
gateway_credential_carrier: None,
client_session_affinity: None,
original_client_session_id: None,
routing_policy: None,
routing_trace_seed: None,
routing_context: None,
@@ -23,6 +23,8 @@ pub(crate) use responses::{
build_local_openai_responses_stream_plan_and_reports_for_kind,
build_local_openai_responses_sync_attempt_source_for_kind,
build_local_openai_responses_sync_plan_and_reports_for_kind,
maybe_build_responses_websocket_decision,
maybe_build_stream_local_openai_responses_decision_payload,
maybe_build_sync_local_openai_responses_decision_payload,
maybe_build_sync_local_openai_responses_decision_payload, ResponsesWebSocketBodyNormalization,
ResponsesWebSocketDecision, ResponsesWebSocketPinnedCandidate,
};
@@ -9,7 +9,9 @@ pub(super) use self::payload::maybe_build_local_openai_responses_decision_payloa
pub(super) use self::support::{
build_local_openai_responses_candidate_attempt_source,
materialize_local_openai_responses_candidate_attempts,
resolve_local_openai_responses_decision_input, LocalOpenAiResponsesCandidateAttempt,
LocalOpenAiResponsesCandidateAttemptSource, LocalOpenAiResponsesDecisionInput,
resolve_local_openai_responses_decision_input,
resolve_local_openai_responses_decision_input_with_snapshot,
LocalOpenAiResponsesCandidateAttempt, LocalOpenAiResponsesCandidateAttemptSource,
LocalOpenAiResponsesDecisionInput,
};
pub(super) use crate::ai_serving::LocalOpenAiResponsesSpec;
@@ -1395,7 +1395,10 @@ async fn resolve_openai_responses_to_openai_image_payload_parts(
return None;
};
let operation = openai_image_operation_from_summary(&image_request_summary)?;
if !is_chatgpt_web {
if is_codex {
provider_request_body =
project_codex_openai_image_api_request_body(&provider_request_body, operation)?;
} else if !is_chatgpt_web {
provider_request_body = project_openai_image_api_request_body(
&provider_request_body,
&prepared_candidate.mapped_model,
@@ -1406,10 +1409,6 @@ async fn resolve_openai_responses_to_openai_image_payload_parts(
),
)?;
}
if is_codex {
provider_request_body =
project_codex_openai_image_api_request_body(&provider_request_body, operation)?;
}
let upstream_url = if is_chatgpt_web {
chatgpt_web_image_internal_url(&transport.endpoint.base_url)
@@ -22,6 +22,7 @@ use crate::ai_serving::planner::common::extract_standard_requested_model;
use crate::ai_serving::planner::decision_input::{
attach_routing_policy_to_local_requested_model_input,
build_local_requested_model_decision_input, resolve_local_authenticated_decision_input,
resolve_local_authenticated_decision_input_with_snapshot,
};
use crate::ai_serving::planner::materialization_policy::{
build_local_candidate_persistence_policy, LocalCandidatePersistencePolicyKind,
@@ -32,7 +33,8 @@ use crate::ai_serving::planner::CandidateFailureDiagnostic;
use crate::ai_serving::{
ai_local_execution_contract_for_formats, extract_pool_sticky_session_token,
openai_responses_request_operation, resolve_local_decision_execution_runtime_auth_context,
ExecutionRuntimeAuthContext, GatewayControlDecision, PlannerAppState,
ExecutionRuntimeAuthContext, GatewayAuthApiKeySnapshot, GatewayControlDecision,
PlannerAppState,
};
use crate::client_session_affinity::client_session_affinity_from_parts;
use crate::{AppState, GatewayError};
@@ -51,6 +53,21 @@ pub(crate) async fn resolve_local_openai_responses_decision_input(
decision: &GatewayControlDecision,
body_json: &serde_json::Value,
plan_kind: &str,
) -> Result<Option<LocalOpenAiResponsesDecisionInput>, GatewayError> {
resolve_local_openai_responses_decision_input_with_snapshot(
state, parts, trace_id, decision, body_json, plan_kind, None,
)
.await
}
pub(crate) async fn resolve_local_openai_responses_decision_input_with_snapshot(
state: &AppState,
parts: &http::request::Parts,
trace_id: &str,
decision: &GatewayControlDecision,
body_json: &serde_json::Value,
plan_kind: &str,
auth_snapshot_override: Option<&GatewayAuthApiKeySnapshot>,
) -> Result<Option<LocalOpenAiResponsesDecisionInput>, GatewayError> {
let Some(auth_context) = resolve_local_decision_execution_runtime_auth_context(decision) else {
warn!(
@@ -87,16 +104,28 @@ pub(crate) async fn resolve_local_openai_responses_decision_input(
return Ok(None);
};
let resolved_input = match resolve_local_authenticated_decision_input(
state,
auth_context.clone(),
Some(requested_model.as_str()),
decision.auth_endpoint_signature.as_deref(),
None,
&decision.model_directive_policy,
)
.await
{
let resolved_input = match if let Some(auth_snapshot) = auth_snapshot_override {
resolve_local_authenticated_decision_input_with_snapshot(
state,
auth_context.clone(),
Some(auth_snapshot.clone()),
Some(requested_model.as_str()),
decision.auth_endpoint_signature.as_deref(),
None,
&decision.model_directive_policy,
)
.await
} else {
resolve_local_authenticated_decision_input(
state,
auth_context.clone(),
Some(requested_model.as_str()),
decision.auth_endpoint_signature.as_deref(),
None,
&decision.model_directive_policy,
)
.await
} {
Ok(Some(resolved_input)) => resolved_input,
Ok(None) => {
warn!(
@@ -1,6 +1,59 @@
use crate::ai_serving::planner::common::endpoint_config_forces_body_stream_field;
use crate::ai_serving::planner::plan_builders::{AiStreamAttempt, AiSyncAttempt};
use crate::ai_serving::planner::spec_metadata::local_openai_responses_spec_metadata;
use crate::ai_serving::planner::standard::codex::codex_model_capabilities_for_transport;
use crate::ai_serving::planner::standard::normalize::build_local_openai_responses_request_body_with_codex_model_capabilities;
use crate::ai_serving::GatewayControlDecision;
use crate::orchestration::{
codex_quota_breaker_blocks_candidate, log_codex_quota_breaker_check_failure,
responses_websocket_adapter, ResponsesWebSocketAdapter,
};
use crate::{AiExecutionDecision, AppState, GatewayError};
use aether_runtime_state::RuntimeLockLease;
use std::collections::BTreeSet;
/// Releases a scheduler pool-key lease if WebSocket planning is cancelled
/// after candidate selection but before ownership reaches the turn lifecycle.
struct ResponsesWebSocketPlanningLeaseGuard {
state: AppState,
lease: Option<RuntimeLockLease>,
}
impl ResponsesWebSocketPlanningLeaseGuard {
fn new(state: &AppState, lease: Option<&RuntimeLockLease>) -> Self {
Self {
state: state.clone(),
lease: lease.cloned(),
}
}
async fn release(mut self) {
// Keep the lease armed across the await. If the owner task is aborted
// or reaches its hard deadline while the runtime backend is stalled,
// Drop can still hand cleanup to a detached owner.
if release_responses_websocket_planning_lease(&self.state, self.lease.as_ref()).await {
self.lease = None;
}
}
fn disarm(&mut self) {
self.lease = None;
}
}
impl Drop for ResponsesWebSocketPlanningLeaseGuard {
fn drop(&mut self) {
let Some(lease) = self.lease.take() else {
return;
};
let state = self.state.clone();
if let Ok(handle) = tokio::runtime::Handle::try_current() {
handle.spawn(async move {
let _ = release_responses_websocket_planning_lease(&state, Some(&lease)).await;
});
}
}
}
mod decision;
mod plans;
@@ -9,6 +62,7 @@ use self::decision::{
build_local_openai_responses_candidate_attempt_source,
maybe_build_local_openai_responses_decision_payload_for_candidate,
resolve_local_openai_responses_decision_input,
resolve_local_openai_responses_decision_input_with_snapshot,
};
use self::plans::{
build_local_stream_attempt_source, build_local_stream_plan_and_reports,
@@ -165,3 +219,373 @@ pub(crate) async fn maybe_build_stream_local_openai_responses_decision_payload(
Ok(None)
}
/// One eligible upstream plus the adapter that is allowed to speak to it.
///
/// The adapter is selected from the provider-scoped capability before the
/// decision leaves the planner. This prevents a public Responses socket from
/// choosing an arbitrary provider protocol after scheduling has completed.
pub(crate) struct ResponsesWebSocketDecision {
pub(crate) execution: AiExecutionDecision,
pub(crate) adapter: ResponsesWebSocketAdapter,
pub(crate) normalization: ResponsesWebSocketBodyNormalization,
}
/// The scheduler identity a continuation is allowed to reuse.
///
/// A `previous_response_id` chain cannot move to another provider connection,
/// but it still has to pass the current scheduler runtime checks on every
/// turn. The planner uses this identity as a filter rather than selecting an
/// arbitrary eligible replacement.
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct ResponsesWebSocketPinnedCandidate {
provider_id: String,
endpoint_id: String,
key_id: String,
}
impl ResponsesWebSocketPinnedCandidate {
pub(crate) fn from_decision(decision: &AiExecutionDecision) -> Option<Self> {
Some(Self {
provider_id: non_empty_decision_identity(decision.provider_id.as_deref())?,
endpoint_id: non_empty_decision_identity(decision.endpoint_id.as_deref())?,
key_id: non_empty_decision_identity(decision.key_id.as_deref())?,
})
}
fn matches(
&self,
candidate: &aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate,
) -> bool {
candidate.provider_id == self.provider_id
&& candidate.endpoint_id == self.endpoint_id
&& candidate.key_id == self.key_id
}
}
fn non_empty_decision_identity(value: Option<&str>) -> Option<String> {
value
.map(str::trim)
.filter(|value| !value.is_empty())
.map(str::to_string)
}
/// Everything needed to re-run provider-body normalization for the candidate a
/// socket is already bound to.
///
/// A continuation turn (`previous_response_id` on the bound upstream) cannot
/// re-enter the planner, because planning selects a candidate and a different
/// key would break the response chain. Without this, such turns reached the
/// provider with only their `model` rewritten — skipping model directives,
/// endpoint body rules, and the Codex body contract that turn 1 received.
///
/// This value holds cloned scalars and JSON only: no candidate, no pool key
/// lease, no `AppState`. It cannot influence selection.
#[derive(Debug, Clone)]
pub(crate) struct ResponsesWebSocketBodyNormalization {
provider_type: String,
provider_api_format: String,
client_api_format: String,
mapped_model: String,
requested_model: String,
upstream_is_stream: bool,
force_body_stream_field: bool,
body_rules: Option<serde_json::Value>,
request_headers: http::HeaderMap,
codex_model_capabilities: Option<crate::ai_serving::CodexResponsesModelCapabilities>,
model_directive_patch: Option<serde_json::Value>,
}
impl ResponsesWebSocketBodyNormalization {
/// Builds a normalizer for a plain `openai:responses` upstream with no
/// endpoint body rules, directives or Codex capabilities, so relay tests can
/// construct a bound connection without standing up a provider snapshot.
#[cfg(test)]
pub(crate) fn for_tests(mapped_model: &str) -> Self {
Self {
provider_type: "openai".to_string(),
provider_api_format: "openai:responses".to_string(),
client_api_format: "openai:responses".to_string(),
mapped_model: mapped_model.to_string(),
requested_model: mapped_model.to_string(),
upstream_is_stream: true,
force_body_stream_field: false,
body_rules: None,
request_headers: http::HeaderMap::new(),
codex_model_capabilities: None,
model_directive_patch: None,
}
}
#[cfg(test)]
pub(crate) fn with_provider_type_for_tests(mut self, provider_type: &str) -> Self {
self.provider_type = provider_type.to_string();
self
}
#[cfg(test)]
pub(crate) fn with_model_directive_patch_for_tests(mut self, patch: serde_json::Value) -> Self {
self.model_directive_patch = Some(patch);
self
}
/// Applies the same body transformations the planner applied on the turn
/// that bound this upstream.
///
/// Mirrors the same-format branch of
/// `resolve_local_openai_responses_candidate_payload_parts`. The
/// cross-format, Kiro, Windsurf and Antigravity branches are unreachable
/// here: the WebSocket planner only returns candidates whose provider API
/// format is `openai:responses`.
///
/// Returns `None` when normalization fails, leaving the caller to fall back
/// to the unnormalized event — a continuation cannot re-select a candidate,
/// so failing the turn outright would be worse than sending it as-is.
pub(crate) fn normalize_response_create(
&self,
client_event: &serde_json::Value,
) -> Option<serde_json::Value> {
use crate::ai_serving::planner::common::{
enforce_provider_body_stream_policy, request_requires_body_stream_field,
};
let source_model = client_event
.get("model")
.and_then(serde_json::Value::as_str)
.unwrap_or(self.requested_model.as_str());
let require_body_stream_field =
request_requires_body_stream_field(client_event, self.force_body_stream_field);
let mut body = build_local_openai_responses_request_body_with_codex_model_capabilities(
client_event,
&self.mapped_model,
self.upstream_is_stream,
self.force_body_stream_field,
self.provider_type.as_str(),
self.provider_api_format.as_str(),
self.body_rules.as_ref(),
&self.request_headers,
self.codex_model_capabilities.as_ref(),
false,
)?;
if let Some(patch) = self.model_directive_patch.as_ref() {
crate::ai_serving::apply_model_directive_mapping_patch(&mut body, patch);
// The patch is a deep merge and may reintroduce `stream`.
enforce_provider_body_stream_policy(
&mut body,
self.provider_api_format.as_str(),
self.upstream_is_stream,
require_body_stream_field,
);
}
crate::ai_serving::finalize_openai_provider_request_with_codex_model_capabilities(
&mut body,
crate::ai_serving::OpenAiProviderRequestFinalization {
source_api_format: self.client_api_format.as_str(),
provider_api_format: self.provider_api_format.as_str(),
provider_type: self.provider_type.as_str(),
provider_model: self.mapped_model.as_str(),
source_model,
body_rules: self.body_rules.as_ref(),
upstream_is_stream: self.upstream_is_stream,
require_body_stream_field,
},
self.codex_model_capabilities.as_ref(),
)
.ok()?;
Some(body)
}
}
/// Builds one upstream decision for a Responses WebSocket turn. The session
/// reuses this decision for same-model turns and invokes the planner again when
/// a later `response.create` changes the public model.
pub(crate) async fn maybe_build_responses_websocket_decision(
state: &AppState,
parts: &http::request::Parts,
trace_id: &str,
decision: &GatewayControlDecision,
auth_snapshot: Option<&crate::ai_serving::GatewayAuthApiKeySnapshot>,
body_json: &serde_json::Value,
excluded_key_ids: Option<&BTreeSet<String>>,
excluded_codex_account_ids: Option<&BTreeSet<String>>,
pinned_candidate: Option<&ResponsesWebSocketPinnedCandidate>,
) -> Result<Option<ResponsesWebSocketDecision>, GatewayError> {
let Some(spec) = resolve_stream_spec(crate::ai_serving::OPENAI_RESPONSES_STREAM_PLAN_KIND)
else {
return Ok(None);
};
let Some(input) = resolve_local_openai_responses_decision_input_with_snapshot(
state,
parts,
trace_id,
decision,
body_json,
spec.decision_kind,
auth_snapshot,
)
.await?
else {
return Ok(None);
};
let body_json = input.effective_body_json(body_json);
let (mut source, _) = build_local_openai_responses_candidate_attempt_source(
state, trace_id, &input, body_json, spec,
)
.await?;
while let Some(attempt) = source.next_attempt().await? {
// `next_attempt` may return with a distributed pool-key lease. Arm a
// guard before the first await so owner-task timeout/cancellation
// cannot strand that lease until its server-side TTL expires.
let mut planning_lease = ResponsesWebSocketPlanningLeaseGuard::new(
state,
attempt.eligible.orchestration.pool_key_lease.as_ref(),
);
if pinned_candidate.is_some_and(|pinned| !pinned.matches(&attempt.eligible.candidate)) {
planning_lease.release().await;
continue;
}
if excluded_key_ids
.is_some_and(|key_ids| key_ids.contains(attempt.eligible.candidate.key_id.as_str()))
{
planning_lease.release().await;
continue;
}
let Some(adapter) = responses_websocket_adapter(
&attempt.eligible.transport.provider.provider_type,
attempt.eligible.transport.provider.config.as_ref(),
) else {
planning_lease.release().await;
continue;
};
// Captured before `attempt` is consumed so a later continuation turn can
// reproduce this candidate's body normalization without re-planning.
let transport = std::sync::Arc::clone(&attempt.eligible.transport);
let candidate_provider_api_format = attempt.eligible.provider_api_format.clone();
let payload = match maybe_build_local_openai_responses_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
.await
{
Ok(Some(payload)) => payload,
Ok(None) => {
planning_lease.release().await;
continue;
}
Err(error) => {
planning_lease.release().await;
return Err(error);
}
};
if payload
.provider_type
.as_deref()
.is_some_and(|value| value.trim().eq_ignore_ascii_case("codex"))
&& crate::orchestration::codex_account_id_from_headers(
&payload.provider_request_headers,
)
.is_some_and(|account_id| {
excluded_codex_account_ids
.is_some_and(|account_ids| account_ids.contains(account_id))
})
{
planning_lease.release().await;
continue;
}
match codex_quota_breaker_blocks_candidate(
state,
payload.provider_type.as_deref(),
payload.key_id.as_deref(),
&payload.provider_request_headers,
)
.await
{
Ok(true) => {
planning_lease.release().await;
continue;
}
Ok(false) => {}
Err(error) => log_codex_quota_breaker_check_failure(&error),
}
if payload
.provider_type
.as_deref()
.is_some_and(|value| adapter.supports_provider_type(value))
&& payload.provider_api_format.as_deref().is_some_and(|value| {
crate::ai_serving::normalize_api_format_alias(value) == "openai:responses"
})
{
let mapped_model = payload.mapped_model.clone().unwrap_or_default();
let source_model = body_json
.get("model")
.and_then(serde_json::Value::as_str)
.unwrap_or(input.requested_model.as_str());
let normalization = ResponsesWebSocketBodyNormalization {
provider_type: transport.provider.provider_type.clone(),
provider_api_format: candidate_provider_api_format.clone(),
client_api_format: local_openai_responses_spec_metadata(spec)
.api_format
.to_string(),
requested_model: input.requested_model.clone(),
upstream_is_stream: payload.upstream_is_stream,
force_body_stream_field: endpoint_config_forces_body_stream_field(
transport.endpoint.config.as_ref(),
),
body_rules: transport.endpoint.body_rules.clone(),
request_headers: input.effective_headers(&parts.headers).clone(),
codex_model_capabilities: codex_model_capabilities_for_transport(
&transport,
candidate_provider_api_format.as_str(),
mapped_model.as_str(),
source_model,
),
model_directive_patch: input
.model_directive_policy
.resolve_reasoning(
candidate_provider_api_format.as_str(),
Some(&input.requested_model),
)
.mapping_patch_for_mapped_model(mapped_model.as_str())
.ok()
.flatten(),
mapped_model,
};
let decision = ResponsesWebSocketDecision {
execution: payload,
adapter,
normalization,
};
// The decision report context now carries the lease identity. The
// WebSocket ownership layer takes over before any further await.
planning_lease.disarm();
return Ok(Some(decision));
}
planning_lease.release().await;
}
Ok(None)
}
async fn release_responses_websocket_planning_lease(
state: &AppState,
lease: Option<&RuntimeLockLease>,
) -> bool {
let Some(lease) = lease else {
return true;
};
match crate::handlers::shared::provider_pool::release_admin_provider_pool_key_lease(
state.runtime_state.as_ref(),
lease,
)
.await
{
Ok(_) => true,
Err(error) => {
tracing::warn!(
error = ?error,
"gateway Responses WebSocket planner failed to release an unused pool key lease"
);
false
}
}
}
@@ -82,10 +82,10 @@ pub(crate) use aether_ai_formats::api::{
parse_codex_auth_identity, parse_direct_request_body, parse_model_directive,
parse_model_directive_with_suffixes, parse_openai_stop_sequences,
parse_openai_tool_result_content, prepare_local_success_response_parts,
prepare_local_success_response_parts_owned, project_codex_openai_image_api_request_body,
project_openai_image_api_request_body, provider_adaptation_allows_sync_finalize_envelope,
provider_adaptation_anchor_api_format, provider_adaptation_descriptor_for_envelope,
provider_adaptation_descriptor_for_provider_type,
prepare_local_success_response_parts_owned, project_codex_catalog_model_card,
project_codex_openai_image_api_request_body, project_openai_image_api_request_body,
provider_adaptation_allows_sync_finalize_envelope, provider_adaptation_anchor_api_format,
provider_adaptation_descriptor_for_envelope, provider_adaptation_descriptor_for_provider_type,
provider_adaptation_requires_eventstream_accept,
provider_adaptation_should_unwrap_stream_envelope,
provider_private_response_allows_sync_finalize, record_converted_response_history,
@@ -175,6 +175,7 @@ pub(crate) use aether_ai_formats::{
is_rerank_api_format, openai_responses_request_operation,
openai_responses_synthetic_reasoning_item_id,
strip_incompatible_openai_responses_reasoning_items, ApiOperation, ClientSurface,
CODEX_CLIENT_VERSION,
};
pub(crate) fn plan_kind_matches_api_operation(
@@ -59,8 +59,9 @@ pub(crate) mod windsurf {
}
pub(crate) use aether_provider_transport::{
append_transport_diagnostics_to_value, apply_local_auth_config_header_overrides,
apply_local_body_rules, apply_local_body_rules_with_request_headers, apply_local_header_rules,
append_transport_diagnostics_to_value, apply_codex_oauth_fingerprint_convergence,
apply_local_auth_config_header_overrides, apply_local_body_rules,
apply_local_body_rules_with_request_headers, apply_local_header_rules,
apply_local_header_rules_with_request_headers, apply_standard_provider_request_body_rules,
apply_standard_provider_request_body_rules_with_request_headers,
apply_transport_request_body_semantics, body_rules_are_locally_supported,
+11 -3
View File
@@ -1,13 +1,17 @@
use axum::body::Body;
use axum::extract::Request;
use axum::http::{header, HeaderValue, Response, StatusCode};
use axum::routing::{any, post};
use axum::routing::{any, get, post};
use axum::Router;
use super::{aliyun, claude, doubao, gemini, jina, openai};
use crate::api::response::build_local_http_error_response_with_request_path;
use crate::headers::extract_or_generate_trace_id;
use crate::{handlers::proxy::proxy_request, state::AppState, GatewayError};
use crate::{
handlers::proxy::{proxy_request, responses_websocket},
state::AppState,
GatewayError,
};
// Router registration patterns live here so AI public ingress has a single mount registry.
// They intentionally stay separate from manifest-facing route inventories in constants.rs,
@@ -51,7 +55,11 @@ const AI_ANY_ROUTE_PATTERNS: &[&str] = &[
pub(crate) fn mount_ai_routes(mut router: Router<AppState>) -> Router<AppState> {
for path in AI_POST_ROUTE_PATTERNS {
router = router.route(path, post(proxy_request));
router = if *path == "/v1/responses" {
router.route(path, get(responses_websocket).post(proxy_request))
} else {
router.route(path, post(proxy_request))
};
}
for path in CLAUDE_POST_ROUTE_PATTERNS {
router = router.route(
+34
View File
@@ -50,12 +50,40 @@ pub(crate) async fn health(State(state): State<AppState>) -> impl IntoResponse {
"rejected": snapshot.rejected,
})
});
let websocket_connection_concurrency =
state
.websocket_connection_concurrency_snapshot()
.map(|snapshot| {
json!({
"limit": snapshot.limit,
"in_flight": snapshot.in_flight,
"available_permits": snapshot.available_permits,
"high_watermark": snapshot.high_watermark,
"rejected": snapshot.rejected,
})
});
let distributed_websocket_connection_concurrency = state
.distributed_websocket_connection_concurrency_snapshot()
.await
.ok()
.flatten()
.map(|snapshot| {
json!({
"limit": snapshot.limit,
"in_flight": snapshot.in_flight,
"available_permits": snapshot.available_permits,
"high_watermark": snapshot.high_watermark,
"rejected": snapshot.rejected,
})
});
Json(json!({
"status": "ok",
"component": "aether-gateway",
"control_api_enabled": true,
"request_concurrency": request_concurrency,
"distributed_request_concurrency": distributed_request_concurrency,
"websocket_connection_concurrency": websocket_connection_concurrency,
"distributed_websocket_connection_concurrency": distributed_websocket_connection_concurrency,
}))
}
@@ -113,6 +141,12 @@ pub(crate) async fn frontdoor_manifest(State(state): State<AppState>) -> impl In
"execution_runtime_configured": state.execution_runtime_configured(),
"request_concurrency_enabled": state.request_concurrency_snapshot().is_some(),
"distributed_request_concurrency_enabled": state.distributed_request_gate.is_some(),
"websocket_connection_concurrency_enabled": state
.websocket_connection_concurrency_snapshot()
.is_some(),
"distributed_websocket_connection_concurrency_enabled": state
.distributed_websocket_connection_gate
.is_some(),
"frontdoor_cors_enabled": cors_enabled,
"frontdoor_cors_allow_credentials": cors_allow_credentials,
"frontdoor_cors_allowed_origins": cors_allowed_origins,
+3 -15
View File
@@ -85,29 +85,17 @@ async fn read_last_backup_slot(app: &AppState) -> Result<Option<String>, Gateway
#[cfg(test)]
mod tests {
use crate::backup::schedule::{BackupSchedule, BackupScheduleUnit};
use crate::task_runtime::{task_definition, TASK_KEY_SYSTEM_S3_BACKUP};
#[test]
fn backup_worker_skips_already_recorded_slot() {
let schedule = BackupSchedule {
unit: BackupScheduleUnit::Days,
interval: 1,
minute: 0,
hour: 3,
weekday: 1,
month_day: 1,
};
let now = chrono::DateTime::parse_from_rfc3339("2026-05-24T03:00:30+08:00")
.unwrap()
.with_timezone(&chrono::Utc);
let slot = schedule.due_slot(now).expect("slot should be due");
let slot = "days:2026-05-23T19:00:00Z";
assert!(super::should_start_scheduled_backup(
Some("days:2026-05-22T19:00:00Z"),
&slot
slot
));
assert!(!super::should_start_scheduled_backup(Some(&slot), &slot));
assert!(!super::should_start_scheduled_backup(Some(slot), slot));
}
#[test]
@@ -0,0 +1,124 @@
//! Credential-safe compatibility probe for the Codex Responses WebSocket path.
//!
//! This binary preserves the established Codex CLI and environment contract.
//! The common Responses WebSocket flow lives in `support/responses_ws_probe`;
//! this profile owns only Codex authentication and header requirements.
#[path = "support/responses_ws_probe.rs"]
mod responses_ws_probe;
use aether_gateway::{CODEX_CLIENT_ORIGINATOR, CODEX_CLIENT_USER_AGENT};
use clap::Parser;
use http::header::{AUTHORIZATION, USER_AGENT};
use http::{HeaderMap, HeaderName, HeaderValue};
use responses_ws_probe::{
bearer_authorization_value, required_env, resolve_probe_url, run_profile_probe, turn_timeout,
ProbeArgs, ProbeConfig, ProbeFailure, ResponsesWebSocketProbeProfile,
};
const ACCESS_TOKEN_ENV: &str = "AETHER_CODEX_WS_PROBE_ACCESS_TOKEN";
const ACCOUNT_ID_ENV: &str = "AETHER_CODEX_WS_PROBE_ACCOUNT_ID";
const MODEL_ENV: &str = "AETHER_CODEX_WS_PROBE_MODEL";
const URL_ENV: &str = "AETHER_CODEX_WS_PROBE_URL";
#[derive(Parser)]
#[command(
name = "aether-codex-ws-probe",
about = "Verify a Codex Responses WebSocket endpoint without exposing credentials"
)]
struct Args {
/// WebSocket endpoint. If omitted, AETHER_CODEX_WS_PROBE_URL is used.
#[arg(long)]
url: Option<String>,
/// Per-turn receive timeout in seconds.
#[arg(long, default_value_t = 20, value_parser = clap::value_parser!(u64).range(1..=120))]
timeout_secs: u64,
}
impl From<Args> for ProbeArgs {
fn from(args: Args) -> Self {
Self {
url: args.url,
timeout_secs: args.timeout_secs,
}
}
}
struct CodexResponsesProbeProfile;
impl ResponsesWebSocketProbeProfile for CodexResponsesProbeProfile {
fn build_config(args: &ProbeArgs) -> Result<ProbeConfig, ProbeFailure> {
let url = resolve_probe_url(args, URL_ENV, None)?;
let access_token = required_env(ACCESS_TOKEN_ENV)?;
let account_id = required_env(ACCOUNT_ID_ENV)?;
let model = required_env(MODEL_ENV)?;
Ok(ProbeConfig::new(
url,
model,
turn_timeout(args),
handshake_headers(&access_token, &account_id)?,
Self::sent_header_names(),
))
}
fn sent_header_names() -> Vec<&'static str> {
vec![
"authorization",
"chatgpt-account-id",
"user-agent",
"originator",
]
}
}
fn handshake_headers(access_token: &str, account_id: &str) -> Result<HeaderMap, ProbeFailure> {
let account_id =
HeaderValue::from_str(account_id).map_err(|_| ProbeFailure::MissingConfiguration)?;
let mut headers = HeaderMap::new();
headers.insert(AUTHORIZATION, bearer_authorization_value(access_token)?);
headers.insert(HeaderName::from_static("chatgpt-account-id"), account_id);
headers.insert(
USER_AGENT,
HeaderValue::from_static(CODEX_CLIENT_USER_AGENT),
);
headers.insert(
HeaderName::from_static("originator"),
HeaderValue::from_static(CODEX_CLIENT_ORIGINATOR),
);
Ok(headers)
}
#[tokio::main]
async fn main() {
let exit_code = run_profile_probe::<CodexResponsesProbeProfile>(Args::parse().into()).await;
if exit_code != 0 {
std::process::exit(exit_code);
}
}
#[cfg(test)]
mod tests {
use http::header::{AUTHORIZATION, USER_AGENT};
use super::{handshake_headers, CodexResponsesProbeProfile, ResponsesWebSocketProbeProfile};
#[test]
fn codex_profile_keeps_its_required_handshake_headers() {
let headers =
handshake_headers("test-token", "test-account").expect("headers should build");
assert!(headers.contains_key(AUTHORIZATION));
assert!(headers.contains_key("chatgpt-account-id"));
assert!(headers.contains_key(USER_AGENT));
assert!(headers.contains_key("originator"));
assert_eq!(
CodexResponsesProbeProfile::sent_header_names(),
vec![
"authorization",
"chatgpt-account-id",
"user-agent",
"originator",
]
);
}
}
@@ -0,0 +1,104 @@
//! Credential-safe compatibility probe for the official OpenAI Responses
//! WebSocket endpoint.
//!
//! This profile uses standard API-key Bearer authentication and shares the
//! protocol flow with the Codex probe without inheriting Codex-specific
//! account headers or quota assumptions.
#[path = "support/responses_ws_probe.rs"]
mod responses_ws_probe;
use clap::Parser;
use http::header::AUTHORIZATION;
use http::HeaderMap;
use responses_ws_probe::{
bearer_authorization_value, required_env, resolve_probe_url, run_profile_probe, turn_timeout,
ProbeArgs, ProbeConfig, ProbeFailure, ResponsesWebSocketProbeProfile,
};
const API_KEY_ENV: &str = "AETHER_OPENAI_WS_PROBE_API_KEY";
const MODEL_ENV: &str = "AETHER_OPENAI_WS_PROBE_MODEL";
const URL_ENV: &str = "AETHER_OPENAI_WS_PROBE_URL";
const DEFAULT_URL: &str = "wss://api.openai.com/v1/responses";
#[derive(Parser)]
#[command(
name = "aether-openai-responses-ws-probe",
about = "Verify an OpenAI Responses WebSocket endpoint without exposing credentials"
)]
struct Args {
/// WebSocket endpoint. If omitted, AETHER_OPENAI_WS_PROBE_URL or the
/// official OpenAI endpoint is used.
#[arg(long)]
url: Option<String>,
/// Per-turn receive timeout in seconds.
#[arg(long, default_value_t = 20, value_parser = clap::value_parser!(u64).range(1..=120))]
timeout_secs: u64,
}
impl From<Args> for ProbeArgs {
fn from(args: Args) -> Self {
Self {
url: args.url,
timeout_secs: args.timeout_secs,
}
}
}
struct OpenAiResponsesProbeProfile;
impl ResponsesWebSocketProbeProfile for OpenAiResponsesProbeProfile {
fn build_config(args: &ProbeArgs) -> Result<ProbeConfig, ProbeFailure> {
let url = resolve_probe_url(args, URL_ENV, Some(DEFAULT_URL))?;
let api_key = required_env(API_KEY_ENV)?;
let model = required_env(MODEL_ENV)?;
let mut headers = HeaderMap::new();
headers.insert(AUTHORIZATION, bearer_authorization_value(&api_key)?);
Ok(ProbeConfig::new(
url,
model,
turn_timeout(args),
headers,
Self::sent_header_names(),
))
}
fn sent_header_names() -> Vec<&'static str> {
vec!["authorization"]
}
}
#[tokio::main]
async fn main() {
let exit_code = run_profile_probe::<OpenAiResponsesProbeProfile>(Args::parse().into()).await;
if exit_code != 0 {
std::process::exit(exit_code);
}
}
#[cfg(test)]
mod tests {
use http::header::AUTHORIZATION;
use super::{
bearer_authorization_value, responses_ws_probe::parse_probe_url,
OpenAiResponsesProbeProfile, ResponsesWebSocketProbeProfile, DEFAULT_URL,
};
#[test]
fn openai_profile_exposes_only_standard_bearer_authentication() {
let authorization = bearer_authorization_value("test-key").expect("header should build");
assert_eq!(authorization.to_str().ok(), Some("Bearer test-key"));
assert_eq!(
OpenAiResponsesProbeProfile::sent_header_names(),
vec![AUTHORIZATION.as_str()]
);
}
#[test]
fn openai_profile_uses_the_official_responses_websocket_endpoint_by_default() {
let url = parse_probe_url(DEFAULT_URL).expect("default OpenAI endpoint should be valid");
assert_eq!(url.as_str(), DEFAULT_URL);
}
}
@@ -0,0 +1,567 @@
//! Shared, credential-safe engine for Responses WebSocket compatibility probes.
//!
//! Provider profiles own their environment variables and handshake headers.
//! This module owns the common Responses WebSocket contract: two sequential
//! `response.create` warmups, continuation with `previous_response_id`, safe
//! event observation, and a redacted JSON report.
use std::env;
use std::time::{Duration, Instant};
use http::{HeaderMap, HeaderValue};
use serde::Serialize;
use serde_json::{json, Value};
use url::Url;
use wreq::ws::message::Message as WreqWsMessage;
const MAX_FRAME_SIZE: usize = 1 << 20;
const MAX_EVENTS_PER_TURN: usize = 16;
pub(crate) struct ProbeArgs {
pub(crate) url: Option<String>,
pub(crate) timeout_secs: u64,
}
pub(crate) struct ProbeConfig {
url: Url,
model: String,
turn_timeout: Duration,
handshake_headers: HeaderMap,
sent_header_names: Vec<&'static str>,
}
impl ProbeConfig {
pub(crate) fn new(
url: Url,
model: String,
turn_timeout: Duration,
handshake_headers: HeaderMap,
sent_header_names: Vec<&'static str>,
) -> Self {
Self {
url,
model,
turn_timeout,
handshake_headers,
sent_header_names,
}
}
}
/// A profile retains provider-specific authentication and configuration while
/// reusing one Responses protocol probe engine.
pub(crate) trait ResponsesWebSocketProbeProfile {
fn build_config(args: &ProbeArgs) -> Result<ProbeConfig, ProbeFailure>;
fn sent_header_names() -> Vec<&'static str>;
}
#[derive(Debug, Clone, Copy)]
pub(crate) enum ProbeFailure {
MissingConfiguration,
InvalidEndpoint,
ClientBuild,
Handshake,
Upgrade,
Send,
ReceiveTimeout,
Receive,
RemoteError,
MissingResponseId,
UnexpectedFrame,
}
impl ProbeFailure {
const fn code(self) -> &'static str {
match self {
Self::MissingConfiguration => "missing_configuration",
Self::InvalidEndpoint => "invalid_endpoint",
Self::ClientBuild => "client_build_failed",
Self::Handshake => "handshake_failed",
Self::Upgrade => "upgrade_failed",
Self::Send => "send_failed",
Self::ReceiveTimeout => "receive_timeout",
Self::Receive => "receive_failed",
Self::RemoteError => "upstream_error_event",
Self::MissingResponseId => "response_id_not_observed",
Self::UnexpectedFrame => "unexpected_frame",
}
}
}
#[derive(Serialize)]
struct ProbeReport {
status: &'static str,
target_host: Option<String>,
handshake_status: Option<u16>,
sent_header_names: Vec<&'static str>,
received_header_names: Vec<String>,
observed_event_types: Vec<String>,
continuation_confirmed: bool,
elapsed_ms: u64,
error: Option<&'static str>,
}
impl ProbeReport {
fn failed(
config: Option<&ProbeConfig>,
sent_header_names: Vec<&'static str>,
started_at: Instant,
error: ProbeFailure,
) -> Self {
Self {
status: "failed",
target_host: config.and_then(target_host),
handshake_status: None,
sent_header_names,
received_header_names: Vec::new(),
observed_event_types: Vec::new(),
continuation_confirmed: false,
elapsed_ms: started_at.elapsed().as_millis() as u64,
error: Some(error.code()),
}
}
}
/// Runs a profile and returns the process exit code after emitting exactly one
/// credential-safe JSON report.
pub(crate) async fn run_profile_probe<P: ResponsesWebSocketProbeProfile>(args: ProbeArgs) -> i32 {
let started_at = Instant::now();
let config = match P::build_config(&args) {
Ok(config) => config,
Err(error) => {
print_report(&ProbeReport::failed(
None,
P::sent_header_names(),
started_at,
error,
));
return 2;
}
};
match run_probe(&config, started_at).await {
Ok(report) => {
print_report(&report);
0
}
Err(error) => {
print_report(&ProbeReport::failed(
Some(&config),
config.sent_header_names.clone(),
started_at,
error,
));
1
}
}
}
pub(crate) fn required_env(name: &str) -> Result<String, ProbeFailure> {
env::var(name)
.ok()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
.ok_or(ProbeFailure::MissingConfiguration)
}
pub(crate) fn resolve_probe_url(
args: &ProbeArgs,
url_env: &str,
default_url: Option<&str>,
) -> Result<Url, ProbeFailure> {
let raw_url = args
.url
.as_deref()
.map(str::to_owned)
.or_else(|| env::var(url_env).ok())
.or_else(|| default_url.map(str::to_owned));
let Some(raw_url) = raw_url
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
else {
return Err(ProbeFailure::MissingConfiguration);
};
parse_probe_url(raw_url)
}
pub(crate) fn parse_probe_url(raw: &str) -> Result<Url, ProbeFailure> {
let url = Url::parse(raw).map_err(|_| ProbeFailure::InvalidEndpoint)?;
if !matches!(url.scheme(), "ws" | "wss")
|| url.host_str().is_none()
|| !url.username().is_empty()
|| url.password().is_some()
|| url.query().is_some()
|| url.fragment().is_some()
{
return Err(ProbeFailure::InvalidEndpoint);
}
Ok(url)
}
pub(crate) fn bearer_authorization_value(token: &str) -> Result<HeaderValue, ProbeFailure> {
HeaderValue::from_str(format!("Bearer {token}").as_str())
.map_err(|_| ProbeFailure::MissingConfiguration)
}
pub(crate) const fn turn_timeout(args: &ProbeArgs) -> Duration {
Duration::from_secs(args.timeout_secs)
}
async fn run_probe(config: &ProbeConfig, started_at: Instant) -> Result<ProbeReport, ProbeFailure> {
let client = wreq::Client::builder()
.connect_timeout(config.turn_timeout)
.timeout(config.turn_timeout)
.build()
.map_err(|_| ProbeFailure::ClientBuild)?;
let response = client
.websocket(config.url.as_str())
.headers(config.handshake_headers.clone())
.max_frame_size(MAX_FRAME_SIZE)
.max_message_size(MAX_FRAME_SIZE)
.send()
.await
.map_err(|_| ProbeFailure::Handshake)?;
let handshake_status = response.status().as_u16();
let received_header_names = response
.headers()
.keys()
.map(|name| name.as_str().to_string())
.collect();
let mut socket = response
.into_websocket()
.await
.map_err(|_| ProbeFailure::Upgrade)?;
let mut observed_event_types = Vec::new();
send_warmup(&mut socket, &config.model, None).await?;
let first_response_id =
receive_completed_response_id(&mut socket, config.turn_timeout, &mut observed_event_types)
.await?;
send_warmup(&mut socket, &config.model, Some(&first_response_id)).await?;
let _second_response_id =
receive_completed_response_id(&mut socket, config.turn_timeout, &mut observed_event_types)
.await?;
Ok(ProbeReport {
status: "passed",
target_host: target_host(config),
handshake_status: Some(handshake_status),
sent_header_names: config.sent_header_names.clone(),
received_header_names,
observed_event_types,
continuation_confirmed: true,
elapsed_ms: started_at.elapsed().as_millis() as u64,
error: None,
})
}
fn target_host(config: &ProbeConfig) -> Option<String> {
config.url.host_str().map(|host| match config.url.port() {
Some(port) => format!("{host}:{port}"),
None => host.to_string(),
})
}
async fn send_warmup(
socket: &mut wreq::ws::WebSocket,
model: &str,
previous_response_id: Option<&str>,
) -> Result<(), ProbeFailure> {
let mut event = json!({
"type": "response.create",
"model": model,
"store": false,
"generate": false,
"input": [],
"tools": [],
});
if let Some(previous_response_id) = previous_response_id {
event["previous_response_id"] = Value::String(previous_response_id.to_string());
}
socket
.send(WreqWsMessage::text(event.to_string()))
.await
.map_err(|_| ProbeFailure::Send)
}
async fn receive_completed_response_id(
socket: &mut wreq::ws::WebSocket,
timeout: Duration,
observed_event_types: &mut Vec<String>,
) -> Result<String, ProbeFailure> {
let mut response_id = None;
for _ in 0..MAX_EVENTS_PER_TURN {
let message = tokio::time::timeout(timeout, socket.recv())
.await
.map_err(|_| ProbeFailure::ReceiveTimeout)?
.ok_or(ProbeFailure::MissingResponseId)?
.map_err(|_| ProbeFailure::Receive)?;
match message {
WreqWsMessage::Text(text) => {
let event: Value = serde_json::from_str(text.as_str())
.map_err(|_| ProbeFailure::UnexpectedFrame)?;
let event_type = event
.get("type")
.and_then(Value::as_str)
.map(safe_event_label)
.unwrap_or_else(|| "unknown".to_string());
let is_remote_error = event_type == "error";
let is_completed = event_type == "response.completed";
observed_event_types.push(event_type);
if is_remote_error {
return Err(ProbeFailure::RemoteError);
}
if let Some(observed_response_id) = event
.pointer("/response/id")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
{
response_id = Some(observed_response_id.to_string());
}
if is_completed {
return response_id.ok_or(ProbeFailure::MissingResponseId);
}
}
WreqWsMessage::Ping(_) | WreqWsMessage::Pong(_) => continue,
WreqWsMessage::Close(_) => return Err(ProbeFailure::MissingResponseId),
_ => return Err(ProbeFailure::UnexpectedFrame),
}
}
Err(ProbeFailure::MissingResponseId)
}
fn safe_event_label(value: &str) -> String {
let trimmed = value.trim();
if trimmed.is_empty()
|| trimmed.len() > 80
|| !trimmed
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'.' | b'_' | b'-'))
{
return "unknown".to_string();
}
trimmed.to_string()
}
fn print_report(report: &ProbeReport) {
match serde_json::to_string(report) {
Ok(json) => println!("{json}"),
Err(_) => println!("{{\"status\":\"failed\",\"error\":\"report_serialization_failed\"}}"),
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::time::{Duration, Instant};
use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
use axum::extract::State;
use axum::http::header::AUTHORIZATION;
use axum::http::{HeaderMap, HeaderValue};
use axum::response::IntoResponse;
use axum::routing::get;
use axum::Router;
use futures_util::{SinkExt, StreamExt};
use serde_json::Value;
use tokio::sync::{oneshot, Mutex};
use super::{parse_probe_url, run_probe, ProbeConfig};
#[derive(Default)]
struct MockState {
observed: Mutex<Option<oneshot::Sender<ObservedClientMessages>>>,
}
struct ObservedClientMessages {
authorization_present: bool,
profile_header_present: bool,
second_before_first_completion: bool,
first: Value,
second: Value,
}
#[tokio::test]
async fn probe_confirms_sequential_response_continuation_without_exposing_values() {
let (url, observed, server) = spawn_mock_server().await;
let mut headers = HeaderMap::new();
headers.insert(
AUTHORIZATION,
HeaderValue::from_static("Bearer test-token-that-must-not-be-reported"),
);
headers.insert(
"x-aether-probe-profile",
HeaderValue::from_static("test-profile-id"),
);
let config = ProbeConfig::new(
parse_probe_url(url.as_str()).expect("mock URL should be valid"),
"gpt-test".to_string(),
Duration::from_secs(2),
headers,
vec!["authorization", "x-aether-probe-profile"],
);
let report = run_probe(&config, Instant::now())
.await
.expect("probe should complete against mock server");
let client_messages = observed.await.expect("mock should observe client messages");
server.abort();
assert_eq!(report.status, "passed");
assert!(report.continuation_confirmed);
assert!(report
.observed_event_types
.contains(&"response.created".to_string()));
assert!(report
.observed_event_types
.contains(&"response.completed".to_string()));
assert!(client_messages.authorization_present);
assert!(client_messages.profile_header_present);
assert!(!client_messages.second_before_first_completion);
assert_eq!(client_messages.first["type"], "response.create");
assert_eq!(client_messages.first["generate"], false);
assert_eq!(client_messages.first["store"], false);
assert_eq!(client_messages.second["previous_response_id"], "resp-first");
let report_json = serde_json::to_string(&report).expect("report should serialize");
assert!(!report_json.contains("test-token-that-must-not-be-reported"));
assert!(!report_json.contains("test-profile-id"));
assert!(!report_json.contains("resp-first"));
}
#[test]
fn probe_url_rejects_credentials_and_query_strings() {
assert!(parse_probe_url("wss://example.test/v1/responses").is_ok());
assert!(parse_probe_url("https://example.test/v1/responses").is_err());
assert!(parse_probe_url("wss://[email protected]/v1/responses").is_err());
assert!(parse_probe_url("wss://example.test/v1/responses?token=secret").is_err());
}
async fn spawn_mock_server() -> (
String,
oneshot::Receiver<ObservedClientMessages>,
tokio::task::JoinHandle<()>,
) {
let (observed_tx, observed_rx) = oneshot::channel();
let state = Arc::new(MockState {
observed: Mutex::new(Some(observed_tx)),
});
let app = Router::new()
.route("/v1/responses", get(mock_websocket))
.with_state(state);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("mock listener should bind");
let address = listener
.local_addr()
.expect("mock listener should expose address");
let server = tokio::spawn(async move {
axum::serve(listener, app)
.await
.expect("mock server should run");
});
(format!("ws://{address}/v1/responses"), observed_rx, server)
}
async fn mock_websocket(
ws: WebSocketUpgrade,
State(state): State<Arc<MockState>>,
headers: HeaderMap,
) -> impl IntoResponse {
let authorization_present = headers
.get(AUTHORIZATION)
.and_then(|value| value.to_str().ok())
.is_some_and(|value| value.starts_with("Bearer "));
let profile_header_present = headers.contains_key("x-aether-probe-profile");
ws.on_upgrade(move |socket| async move {
serve_mock_socket(socket, state, authorization_present, profile_header_present).await;
})
}
async fn serve_mock_socket(
socket: WebSocket,
state: Arc<MockState>,
authorization_present: bool,
profile_header_present: bool,
) {
let (mut sender, mut receiver) = socket.split();
let first = receive_json(&mut receiver).await;
let _ = sender
.send(Message::Text(
serde_json::json!({
"type": "response.created",
"response": {"id": "resp-first"}
})
.to_string()
.into(),
))
.await;
let early_second = tokio::select! {
message = receiver.next() => Some(message),
_ = tokio::time::sleep(Duration::from_millis(50)) => None,
};
let second_before_first_completion = early_second.is_some();
let _ = sender
.send(Message::Text(
serde_json::json!({
"type": "response.completed",
"response": {"id": "resp-first", "status": "completed"}
})
.to_string()
.into(),
))
.await;
let second = match early_second {
Some(Some(Ok(Message::Text(text)))) => {
serde_json::from_str(text.as_str()).expect("early client message should be JSON")
}
Some(Some(Ok(_))) => panic!("expected text continuation message"),
Some(Some(Err(error))) => panic!("client message should be valid: {error}"),
Some(None) => panic!("client closed before continuation"),
None => receive_json(&mut receiver).await,
};
let _ = sender
.send(Message::Text(
serde_json::json!({
"type": "response.created",
"response": {"id": "resp-second"}
})
.to_string()
.into(),
))
.await;
let _ = sender
.send(Message::Text(
serde_json::json!({
"type": "response.completed",
"response": {"id": "resp-second", "status": "completed"}
})
.to_string()
.into(),
))
.await;
if let Some(observed) = state.observed.lock().await.take() {
let _ = observed.send(ObservedClientMessages {
authorization_present,
profile_header_present,
second_before_first_completion,
first,
second,
});
}
}
async fn receive_json(receiver: &mut futures_util::stream::SplitStream<WebSocket>) -> Value {
let message = receiver
.next()
.await
.expect("client should send a message")
.expect("client message should be valid");
let Message::Text(text) = message else {
panic!("expected text message");
};
serde_json::from_str(text.as_str()).expect("client message should be JSON")
}
}
-10
View File
@@ -1274,7 +1274,6 @@ mod tests {
let first_cache = Arc::clone(&cache);
let first_key = key.clone();
let first_calls = Arc::clone(&calls);
let first_started = Instant::now();
let first = tokio::spawn(async move {
first_cache
.get_or_load_once_stale_while_refreshing::<(), _, _>(
@@ -1292,7 +1291,6 @@ mod tests {
});
let follower_cache = Arc::clone(&cache);
let follower_started = Instant::now();
let follower_calls = Arc::clone(&calls);
let follower = tokio::spawn(async move {
follower_cache
@@ -1310,15 +1308,7 @@ mod tests {
});
assert_eq!(first.await.unwrap().unwrap(), Some(1));
assert!(
first_started.elapsed() < Duration::from_millis(80),
"stale value should not wait for request-path refresh"
);
assert_eq!(follower.await.unwrap().unwrap(), Some(1));
assert!(
follower_started.elapsed() < Duration::from_millis(80),
"follower should return stale value without waiting for refresh"
);
assert_eq!(calls.load(Ordering::Acquire), 0);
}
@@ -177,6 +177,19 @@ fn extract_trusted_auth_headers(headers: &http::HeaderMap) -> Option<GatewayTrus
})
}
#[cfg(not(test))]
pub(super) fn extract_trusted_admin_headers(
_headers: &http::HeaderMap,
) -> Option<GatewayTrustedAdminHeaders> {
// The public gateway has no authenticated upstream that is allowed to
// assert an administrator principal. `x-aether-gateway` is also emitted
// on public responses, so it cannot serve as proof that these headers were
// produced by a trusted hop. Production requests must authenticate with a
// real admin session or management bearer token instead.
None
}
#[cfg(test)]
pub(super) fn extract_trusted_admin_headers(
headers: &http::HeaderMap,
) -> Option<GatewayTrustedAdminHeaders> {
+3 -2
View File
@@ -11,8 +11,9 @@ pub(crate) use gate::{
should_buffer_request_for_local_auth, trusted_auth_local_rejection, GatewayLocalAuthRejection,
};
pub(crate) use resolution::{
refresh_execution_runtime_auth_context, resolve_execution_runtime_auth_context,
GatewayAdminPrincipalContext, GatewayControlAuthContext,
refresh_execution_runtime_auth_context, refresh_execution_runtime_auth_context_with_snapshot,
resolve_execution_runtime_auth_context, GatewayAdminPrincipalContext,
GatewayControlAuthContext,
};
pub(super) use resolution::{resolve_control_decision_auth, ControlDecisionAuthResolution};
pub(crate) use types::GatewayCredentialCarrier;
@@ -725,20 +725,47 @@ pub(crate) async fn refresh_execution_runtime_auth_context(
auth_context: GatewayControlAuthContext,
auth_endpoint_signature: Option<&str>,
) -> Result<GatewayControlAuthContext, GatewayError> {
refresh_execution_runtime_auth_context_with_snapshot(
state,
auth_context,
auth_endpoint_signature,
)
.await
.map(|(auth_context, _)| auth_context)
}
/// Strongly refreshes the long-lived execution authorization context and
/// returns the exact API-key snapshot that produced it.
///
/// WebSocket turns need both values: using the refreshed context for RPM and
/// balance checks while letting the planner independently read its normal
/// cache can authorize a different provider/model snapshot for up to the cache
/// TTL. Ordinary HTTP callers keep using [`refresh_execution_runtime_auth_context`].
pub(crate) async fn refresh_execution_runtime_auth_context_with_snapshot(
state: &AppState,
auth_context: GatewayControlAuthContext,
auth_endpoint_signature: Option<&str>,
) -> Result<
(
GatewayControlAuthContext,
Option<crate::ai_serving::GatewayAuthApiKeySnapshot>,
),
GatewayError,
> {
if auth_context.local_rejection.is_some() || !auth_context.access_allowed {
return Ok(auth_context);
return Ok((auth_context, None));
}
let Some(auth_endpoint_signature) = auth_endpoint_signature
.map(str::trim)
.filter(|value| !value.is_empty())
else {
return Ok(auth_context);
return Ok((auth_context, None));
};
if !state.has_auth_api_key_reader()
|| auth_context.user_id.trim().is_empty()
|| auth_context.api_key_id.trim().is_empty()
{
return Ok(auth_context);
return Ok((auth_context, None));
}
let snapshot = {
@@ -758,19 +785,20 @@ pub(crate) async fn refresh_execution_runtime_auth_context(
denied.access_allowed = false;
denied.local_rejection = Some(GatewayLocalAuthRejection::InvalidApiKey);
denied.balance_remaining = None;
return Ok(denied);
return Ok((denied, None));
};
let wallet_access = resolve_wallet_auth_gate_uncached(state, &snapshot).await?;
Ok(build_data_backed_auth_context(
let refreshed = build_data_backed_auth_context(
state,
snapshot,
snapshot.clone(),
auth_endpoint_signature,
Some(true),
auth_context.balance_remaining,
wallet_access,
)
.await)
.await;
Ok((refreshed, Some(snapshot)))
}
fn put_cached_auth_context(
+5 -4
View File
@@ -9,10 +9,11 @@ mod route;
pub(crate) use auth::{
execution_plan_balance_capacity_rejection, extract_requested_model,
refresh_execution_runtime_auth_context, request_model_local_rejection,
resolve_execution_runtime_auth_context, should_buffer_request_for_local_auth,
trusted_auth_local_rejection, GatewayAdminPrincipalContext, GatewayControlAuthContext,
GatewayCredentialCarrier, GatewayLocalAuthRejection,
refresh_execution_runtime_auth_context, refresh_execution_runtime_auth_context_with_snapshot,
request_model_local_rejection, resolve_execution_runtime_auth_context,
should_buffer_request_for_local_auth, trusted_auth_local_rejection,
GatewayAdminPrincipalContext, GatewayControlAuthContext, GatewayCredentialCarrier,
GatewayLocalAuthRejection,
};
pub(crate) use execute::{allows_control_execute_emergency, maybe_execute_via_control};
pub(crate) use management_token_permissions::{
+51 -1
View File
@@ -35,7 +35,10 @@ pub(super) fn classify_ai_public_route(
"openai:rerank",
true,
))
} else if method == http::Method::POST
} else if (method == http::Method::POST
|| (method == http::Method::GET
&& normalized_path == "/v1/responses"
&& is_websocket_upgrade_request(headers)))
&& matches!(normalized_path, "/v1/responses" | "/v1/responses/compact")
{
if normalized_path.ends_with("/compact") {
@@ -199,6 +202,24 @@ fn claude_request_auth_channel(headers: &http::HeaderMap) -> &'static str {
}
}
fn is_websocket_upgrade_request(headers: &http::HeaderMap) -> bool {
let has_upgrade_connection = headers
.get(http::header::CONNECTION)
.and_then(|value| value.to_str().ok())
.is_some_and(|value| {
value
.split(',')
.map(str::trim)
.any(|value| value.eq_ignore_ascii_case("upgrade"))
});
let has_websocket_upgrade = headers
.get(http::header::UPGRADE)
.and_then(|value| value.to_str().ok())
.is_some_and(|value| value.eq_ignore_ascii_case("websocket"));
has_upgrade_connection && has_websocket_upgrade
}
fn is_gemini_operation_method(method: &http::Method, normalized_path: &str) -> bool {
method == http::Method::GET
|| (method == http::Method::POST && normalized_path.ends_with(":cancel"))
@@ -242,3 +263,32 @@ fn classify_antigravity_v1internal_route(
execution_runtime_candidate,
))
}
#[cfg(test)]
mod tests {
use axum::http::header::{CONNECTION, UPGRADE};
use axum::http::{HeaderMap, HeaderValue, Method};
use super::classify_ai_public_route;
#[test]
fn classifies_websocket_upgrade_on_responses_route() {
let mut headers = HeaderMap::new();
headers.insert(CONNECTION, HeaderValue::from_static("keep-alive, Upgrade"));
headers.insert(UPGRADE, HeaderValue::from_static("websocket"));
let route = classify_ai_public_route(&Method::GET, "/v1/responses", &headers)
.expect("Responses WebSocket should be an AI public route");
assert_eq!(route.route_class, "ai_public");
assert_eq!(route.route_family, "openai");
assert_eq!(route.route_kind, "responses");
assert_eq!(route.auth_endpoint_signature, "openai:responses");
}
#[test]
fn does_not_classify_plain_get_as_responses_websocket() {
assert!(
classify_ai_public_route(&Method::GET, "/v1/responses", &HeaderMap::new()).is_none()
);
}
}
+1 -2
View File
@@ -247,8 +247,7 @@ pub(super) fn detect_public_models_auth_signature(uri: &Uri, headers: &http::Hea
let has_codex_client_version = uri.path() == "/v1/models"
&& uri.query().is_some_and(|query| {
url::form_urlencoded::parse(query.as_bytes())
.any(|(key, value)| key == "client_version" && !value.trim().is_empty())
url::form_urlencoded::parse(query.as_bytes()).any(|(key, _)| key == "client_version")
});
if has_codex_client_version {
return "openai:responses".to_string();
@@ -39,7 +39,7 @@ fn classifies_codex_models_list_with_responses_auth_signature() {
}
#[test]
fn empty_codex_client_version_keeps_standard_openai_models_signature() {
fn empty_codex_client_version_uses_responses_signature_for_bounded_fallback() {
let headers = headers(&[("authorization", "Bearer sk-test")]);
let uri: Uri = "/v1/models?client_version="
.parse()
@@ -49,7 +49,7 @@ fn empty_codex_client_version_keeps_standard_openai_models_signature() {
assert_eq!(
decision.auth_endpoint_signature.as_deref(),
Some("openai:chat")
Some("openai:responses")
);
}
+29
View File
@@ -195,6 +195,35 @@ mod tests {
assert_eq!(background.pool.min_connections, 1);
}
#[test]
fn runtime_pool_split_gives_default_small_server_more_foreground_capacity() {
let config = GatewayDataConfig::from_database_config(
SqlDatabaseConfig::new(
DatabaseDriver::Postgres,
"postgres://localhost/aether",
SqlPoolConfig {
min_connections: 4,
max_connections: 32,
..SqlPoolConfig::default()
},
)
.expect("database config should be valid"),
);
let (foreground, background) = config.split_runtime_pools_with_background_max(None);
let foreground = foreground.database().expect("foreground database");
let background = background
.expect("background database config")
.database()
.expect("background database")
.clone();
assert_eq!(foreground.pool.max_connections, 26);
assert_eq!(background.pool.max_connections, 6);
assert_eq!(foreground.pool.min_connections, 4);
assert_eq!(background.pool.min_connections, 1);
}
#[test]
fn runtime_pool_split_can_be_disabled_or_degrade_for_single_connection() {
let mut database = SqlDatabaseConfig::sqlite_default();
+32 -8
View File
@@ -1,14 +1,14 @@
use super::{
ApiKeyLastUsedDelta, DataLayerError, GatewayDataState, GeminiFileMappingListQuery,
GeminiFileMappingStats, ProviderCatalogKeyAdaptiveStateUpdate,
ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListQuery,
ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthRuntimeStateCasUpdate,
ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate,
PublicHealthStatusCount, PublicHealthTimelineBucket, StoredGeminiFileMapping,
StoredGeminiFileMappingListPage, StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage,
StoredProviderCatalogKeyStats, StoredProviderCatalogProvider, StoredRequestCandidate,
UpsertGeminiFileMappingRecord, UpsertRequestCandidateRecord,
ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyHealthStateUpdate,
ProviderCatalogKeyListQuery, ProviderCatalogKeyOAuthCredentialCasDelete,
ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate,
ProviderCatalogKeyStatusSnapshotUpdate, PublicHealthStatusCount, PublicHealthTimelineBucket,
StoredGeminiFileMapping, StoredGeminiFileMappingListPage, StoredProviderCatalogEndpoint,
StoredProviderCatalogKey, StoredProviderCatalogKeyMaintenanceSummary,
StoredProviderCatalogKeyPage, StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
StoredRequestCandidate, UpsertGeminiFileMappingRecord, UpsertRequestCandidateRecord,
};
impl GatewayDataState {
@@ -282,6 +282,16 @@ impl GatewayDataState {
}
}
pub(crate) async fn list_provider_catalog_keys_by_ids_strong(
&self,
key_ids: &[String],
) -> Result<Vec<StoredProviderCatalogKey>, DataLayerError> {
match &self.provider_catalog_reader {
Some(repository) => repository.list_keys_by_ids_strong(key_ids).await,
None => Ok(Vec::new()),
}
}
pub(crate) async fn list_provider_catalog_keys_by_provider_ids(
&self,
provider_ids: &[String],
@@ -561,6 +571,20 @@ impl GatewayDataState {
Ok(updated)
}
pub(crate) async fn compare_and_update_provider_catalog_key_admin_state(
&self,
update: &ProviderCatalogKeyAdminCasUpdate,
) -> Result<bool, DataLayerError> {
let updated = match &self.provider_catalog_writer {
Some(repository) => repository.compare_and_update_key_admin_state(update).await,
None => Ok(false),
}?;
// Clear on both success and conflict so a retry cannot reuse the stale
// credential snapshot that lost the CAS.
self.clear_provider_catalog_cache();
Ok(updated)
}
pub(crate) async fn update_provider_catalog_keys(
&self,
keys: &[StoredProviderCatalogKey],
+7 -7
View File
@@ -121,13 +121,13 @@ use aether_data_contracts::repository::pool_scores::{
UpsertPoolMemberScore,
};
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyAdaptiveStateUpdate, ProviderCatalogKeyHealthStateUpdate,
ProviderCatalogKeyListQuery, ProviderCatalogKeyOAuthCredentialCasDelete,
ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate,
ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogReadRepository,
ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage,
StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
ProviderCatalogKeyAdaptiveStateUpdate, ProviderCatalogKeyAdminCasUpdate,
ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListQuery,
ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthRuntimeStateCasUpdate,
ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate,
ProviderCatalogReadRepository, ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint,
StoredProviderCatalogKey, StoredProviderCatalogKeyMaintenanceSummary,
StoredProviderCatalogKeyPage, StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
};
use aether_data_contracts::repository::quota::{
ProviderQuotaReadRepository, ProviderQuotaWriteRepository, StoredProviderQuotaSnapshot,
@@ -334,6 +334,13 @@ impl ProviderCatalogReadRepository for CachedProviderCatalogReadRepository {
}
}
async fn list_keys_by_ids_strong(
&self,
key_ids: &[String],
) -> Result<Vec<StoredProviderCatalogKey>, DataLayerError> {
self.inner.list_keys_by_ids_strong(key_ids).await
}
async fn list_keys_by_provider_ids(
&self,
provider_ids: &[String],
@@ -538,6 +545,7 @@ fn normalize_ids(ids: &[String]) -> Vec<String> {
mod tests {
use super::*;
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use aether_data_contracts::repository::provider_catalog::ProviderCatalogWriteRepository;
fn cache() -> CachedProviderCatalogReadRepository {
CachedProviderCatalogReadRepository::new(Arc::new(
@@ -555,6 +563,55 @@ mod tests {
.expect("provider should be valid")
}
#[tokio::test]
async fn provider_catalog_strong_key_read_bypasses_fresh_cached_generation() {
let old_metadata = serde_json::json!({
"codex": {"credential_generation": "old"}
});
let mut key = StoredProviderCatalogKey::new(
"key-1".to_string(),
"provider-1".to_string(),
"key-1".to_string(),
"oauth".to_string(),
None,
true,
)
.expect("key should be valid");
key.upstream_metadata = Some(old_metadata.clone());
let inner = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider("provider-1")],
Vec::new(),
vec![key],
));
let cache = CachedProviderCatalogReadRepository::new(inner.clone());
let key_ids = vec!["key-1".to_string()];
let first = cache
.list_keys_by_ids(&key_ids)
.await
.expect("initial key read should succeed");
assert_eq!(first[0].upstream_metadata.as_ref(), Some(&old_metadata));
let new_metadata = serde_json::json!({
"codex": {"credential_generation": "new"}
});
assert!(inner
.upsert_key_upstream_metadata_namespace("key-1", "codex", &new_metadata["codex"], None,)
.await
.expect("inner metadata update should succeed"));
let cached = cache
.list_keys_by_ids(&key_ids)
.await
.expect("cached key read should succeed");
assert_eq!(cached[0].upstream_metadata.as_ref(), Some(&old_metadata));
let strong = cache
.list_keys_by_ids_strong(&key_ids)
.await
.expect("strong key read should succeed");
assert_eq!(strong[0].upstream_metadata.as_ref(), Some(&new_metadata));
}
#[tokio::test]
async fn provider_catalog_follower_observes_completion_before_first_poll() {
let cache = cache();
@@ -0,0 +1,62 @@
//! Shared admission helpers for local upstream execution.
//!
//! The stream candidate loop and long-lived WebSocket turns both need to
//! participate in the same gateway-wide upstream execution gate. Keep the
//! provider abstraction here so tests can supply an isolated gate while
//! production callers use `AppState` directly.
use std::time::Duration;
use aether_runtime::{ConcurrencyGate, ConcurrencyPermit};
use tokio::time::timeout;
use crate::stage_metrics::observe_gateway_stage_ms;
use crate::{AppState, GatewayError};
pub(crate) const UPSTREAM_EXECUTION_GATE_NAME: &str = "gateway_upstream_execution";
pub(crate) trait UpstreamExecutionGateProvider {
fn upstream_execution_gate(&self) -> Option<&ConcurrencyGate>;
fn upstream_execution_gate_queue_budget(&self) -> Duration;
}
impl UpstreamExecutionGateProvider for AppState {
fn upstream_execution_gate(&self) -> Option<&ConcurrencyGate> {
self.upstream_execution_gate.as_deref()
}
fn upstream_execution_gate_queue_budget(&self) -> Duration {
self.frontdoor_runtime_guards.internal_gate_queue_budget
}
}
/// Acquires the shared gateway-wide upstream execution permit.
///
/// A missing gate is an intentional configuration (unlimited), so callers
/// receive `Ok(None)`. Saturation keeps the existing candidate-level
/// `AdmissionTimeout` contract used by the HTTP stream path.
pub(crate) async fn acquire_upstream_execution_gate(
state: &(impl UpstreamExecutionGateProvider + ?Sized),
trace_id: &str,
) -> Result<Option<ConcurrencyPermit>, GatewayError> {
let Some(gate) = state.upstream_execution_gate() else {
return Ok(None);
};
let budget = state.upstream_execution_gate_queue_budget();
let gate_wait_started_at = std::time::Instant::now();
match timeout(budget, gate.acquire()).await {
Ok(Ok(permit)) => {
observe_gateway_stage_ms(
"upstream_execution_gate_wait",
gate_wait_started_at.elapsed().as_millis() as u64,
);
Ok(Some(permit))
}
Ok(Err(err)) => Err(GatewayError::Internal(err.to_string())),
Err(_) => Err(GatewayError::AdmissionTimeout {
trace_id: trace_id.to_string(),
gate: UPSTREAM_EXECUTION_GATE_NAME,
queue_budget_ms: budget.as_millis() as u64,
}),
}
}
File diff suppressed because it is too large Load Diff
@@ -2573,6 +2573,7 @@ fn json_execution_result(
candidate_id: plan.candidate_id.clone(),
status_code,
headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]),
response_observation: None,
body: Some(ResponseBody {
json_body: Some(body),
body_bytes_b64: None,
@@ -2615,6 +2616,7 @@ fn bytes_execution_result(
candidate_id: plan.candidate_id.clone(),
status_code,
headers,
response_observation: None,
body: Some(ResponseBody {
json_body: None,
body_bytes_b64: Some(base64::engine::general_purpose::STANDARD.encode(body)),
@@ -2637,6 +2639,7 @@ fn execution_result_frame_stream(
payload: StreamFramePayload::Headers {
status_code: result.status_code,
headers: result.headers.clone(),
response_observation: result.response_observation.clone(),
},
},
StreamFrame {
@@ -480,6 +480,7 @@ mod tests {
candidate_id: None,
status_code: 502,
headers: Default::default(),
response_observation: None,
body: None,
telemetry: None,
error: None,
@@ -519,6 +520,7 @@ mod tests {
candidate_id: None,
status_code: 502,
headers: Default::default(),
response_observation: None,
body: None,
telemetry: None,
error: None,
@@ -583,6 +585,7 @@ mod tests {
candidate_id: None,
status_code: 429,
headers: Default::default(),
response_observation: None,
body: None,
telemetry: None,
error: None,
@@ -614,6 +617,7 @@ mod tests {
candidate_id: None,
status_code: 401,
headers: Default::default(),
response_observation: None,
body: None,
telemetry: None,
error: None,
@@ -645,6 +649,7 @@ mod tests {
candidate_id: None,
status_code: 502,
headers: Default::default(),
response_observation: None,
body: None,
telemetry: None,
error: Some(ExecutionError {
@@ -708,6 +713,7 @@ mod tests {
candidate_id: None,
status_code: 404,
headers: Default::default(),
response_observation: None,
body: None,
telemetry: None,
error: None,
@@ -900,6 +906,7 @@ mod tests {
candidate_id: None,
status_code: 200,
headers: Default::default(),
response_observation: None,
body: None,
telemetry: None,
error: None,
@@ -1022,6 +1029,7 @@ mod tests {
candidate_id: None,
status_code: 429,
headers: Default::default(),
response_observation: None,
body: None,
telemetry: None,
error: None,
@@ -1068,6 +1076,7 @@ mod tests {
candidate_id: None,
status_code: 429,
headers: Default::default(),
response_observation: None,
body: None,
telemetry: None,
error: None,
@@ -1174,6 +1183,7 @@ mod tests {
candidate_id: None,
status_code: 200,
headers: Default::default(),
response_observation: None,
body: None,
telemetry: None,
error: None,
@@ -1211,6 +1221,7 @@ mod tests {
candidate_id: None,
status_code: 400,
headers: Default::default(),
response_observation: None,
body: None,
telemetry: None,
error: None,
@@ -1259,6 +1270,7 @@ mod tests {
candidate_id: None,
status_code: 429,
headers: Default::default(),
response_observation: None,
body: None,
telemetry: None,
error: None,
@@ -841,6 +841,7 @@ fn encode_grok_headers_frame(
payload: StreamFramePayload::Headers {
status_code,
headers,
response_observation: None,
},
})
}
@@ -2157,6 +2158,7 @@ fn grok_execution_result(
candidate_id: plan.candidate_id.clone(),
status_code,
headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]),
response_observation: None,
body: Some(ResponseBody {
json_body: Some(body_json),
body_bytes_b64: None,
@@ -2220,6 +2222,7 @@ fn grok_collected_frame_stream(
"application/json".to_string()
},
)]),
response_observation: None,
},
},
StreamFrame {
@@ -279,6 +279,7 @@ fn raw_response_frame_stream(
payload: StreamFramePayload::Headers {
status_code,
headers,
response_observation: None,
},
},
StreamFrame {
@@ -1449,6 +1450,7 @@ mod tests {
candidate_id: None,
status_code: 200,
headers: BTreeMap::new(),
response_observation: None,
body: Some(aether_contracts::ResponseBody {
json_body: Some(json!({
"jsonrpc": "2.0",
@@ -3,6 +3,8 @@ use std::collections::BTreeMap;
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
pub(crate) mod admission;
pub(crate) mod attempt_lifecycle;
mod chatgpt_web_image;
mod constants;
mod fallback;
@@ -23,6 +25,9 @@ pub(crate) mod transport;
mod transport_failure;
mod windsurf;
pub(crate) use self::admission::{
acquire_upstream_execution_gate, UpstreamExecutionGateProvider, UPSTREAM_EXECUTION_GATE_NAME,
};
pub(crate) use self::chatgpt_web_image::maybe_execute_chatgpt_web_image_sync;
pub(crate) use self::constants::{
MAX_ERROR_BODY_BYTES, MAX_STREAM_PREFETCH_BYTES, MAX_STREAM_PREFETCH_FRAMES,
@@ -2,10 +2,10 @@ use aether_contracts::ExecutionPlan;
use tracing::warn;
use crate::orchestration::{
oauth_status_may_be_invalid as status_may_be_oauth_invalid,
local_failover_error_message, oauth_status_may_be_invalid as status_may_be_oauth_invalid,
oauth_status_proves_access_token_invalid as status_proves_access_token_invalid,
};
use crate::state::AgentIdentityAuthConfigFence;
use crate::state::{AgentIdentityAuthConfigFence, CodexRuntimeOAuthObservation};
use crate::{provider_transport::LocalOAuthRefreshError, AppState};
pub(crate) async fn refresh_oauth_plan_auth_for_retry(
@@ -14,6 +14,9 @@ pub(crate) async fn refresh_oauth_plan_auth_for_retry(
status_code: u16,
response_text: Option<&str>,
trace_id: &str,
report_context: Option<&serde_json::Value>,
request_started_at_unix_ms: Option<u64>,
request_order_id: Option<&str>,
) -> bool {
if !status_may_be_oauth_invalid(status_code, response_text) {
return false;
@@ -109,15 +112,49 @@ pub(crate) async fn refresh_oauth_plan_auth_for_retry(
body_excerpt,
..
}) if matches!(refresh_status_code, 400 | 401 | 403) => {
if let Err(err) = state
.persist_local_oauth_refresh_failure_state(
&transport,
refresh_status_code,
body_excerpt.as_str(),
access_token_invalid_proven,
)
.await
{
let observed_credential_generation =
report_context_string(report_context, "codex_credential_generation");
let runtime_invalid_message = local_failover_error_message(response_text);
let runtime_invalid_reason =
aether_admin::provider::quota::codex_runtime_invalid_reason(
status_code,
runtime_invalid_message.as_deref(),
);
let persist_result = match (request_started_at_unix_ms, request_order_id) {
(Some(request_started_at_unix_ms), Some(request_order_id))
if transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case("codex") =>
{
state
.persist_local_oauth_refresh_failure_state_observed(
&transport,
refresh_status_code,
body_excerpt.as_str(),
access_token_invalid_proven,
CodexRuntimeOAuthObservation {
request_started_at_unix_ms,
request_order_id,
observed_credential_generation,
runtime_invalid_reason: runtime_invalid_reason.as_deref(),
},
)
.await
}
_ => {
state
.persist_local_oauth_refresh_failure_state(
&transport,
refresh_status_code,
body_excerpt.as_str(),
access_token_invalid_proven,
)
.await
}
};
if let Err(err) = persist_result {
warn!(
event_name = "local_oauth_retry_refresh_failure_persist_failed",
log_type = "ops",
@@ -161,6 +198,17 @@ pub(crate) async fn refresh_oauth_plan_auth_for_retry(
}
}
fn report_context_string<'a>(
report_context: Option<&'a serde_json::Value>,
field: &str,
) -> Option<&'a str> {
report_context
.and_then(|context| context.get(field))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
}
fn execution_plan_authorization(plan: &ExecutionPlan) -> Option<&str> {
plan.headers
.iter()
@@ -209,6 +257,7 @@ mod tests {
use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY};
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyOAuthCredentialFence,
ProviderCatalogReadRepository, ProviderCatalogWriteRepository,
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
@@ -312,7 +361,7 @@ mod tests {
}
#[tokio::test]
async fn auto_removes_codex_key_after_request_proven_terminal_refresh_failure() {
async fn retains_codex_key_after_request_proven_terminal_refresh_failure() {
let token_hits = Arc::new(Mutex::new(0usize));
let token_hits_clone = Arc::clone(&token_hits);
let token_server = Router::new().route(
@@ -458,16 +507,34 @@ mod tests {
401,
Some(r#"{"error":"oauth_token_invalid"}"#),
"trace-oauth-retry",
None,
Some(1_000),
Some("01900000-0000-7000-8000-000000000010"),
)
.await;
assert!(!retried);
assert_eq!(*token_hits.lock().expect("mutex should lock"), 1);
let keys = provider_catalog_repository
let stored_key = provider_catalog_repository
.list_keys_by_ids(&["key-codex-oauth-retry".to_string()])
.await
.expect("keys should read");
assert!(keys.is_empty());
.expect("keys should read")
.into_iter()
.next()
.expect("request-scoped refresh failure should retain the key");
let invalid_reason = stored_key
.oauth_invalid_reason
.as_deref()
.expect("combined invalid reason should persist");
assert!(invalid_reason.contains("[OAUTH_EXPIRED]"));
assert!(invalid_reason.contains("[REFRESH_FAILED]"));
assert_eq!(
stored_key
.upstream_metadata
.as_ref()
.and_then(|metadata| metadata.pointer("/codex/oauth_state_request_id")),
Some(&json!("01900000-0000-7000-8000-000000000010"))
);
token_handle.abort();
}
@@ -619,6 +686,9 @@ mod tests {
401,
Some(r#"{"error":"invalid_token"}"#),
"trace-claude-oauth-fence-first",
None,
None,
None,
)
.await
);
@@ -647,6 +717,9 @@ mod tests {
401,
Some(r#"{"error":"invalid_token"}"#),
"trace-claude-oauth-fence-stale",
None,
None,
None,
)
.await
);
@@ -665,6 +738,7 @@ mod tests {
.expect("Claude key should load")
.pop()
.expect("Claude key should exist");
let expected_admin_replacement = admin_replacement.clone();
admin_replacement.encrypted_api_key = Some(
encrypt_python_fernet_plaintext(
DEVELOPMENT_ENCRYPTION_KEY,
@@ -673,10 +747,23 @@ mod tests {
.expect("admin access token should encrypt"),
);
admin_replacement.expires_at_unix_secs = Some(4_102_444_800);
provider_catalog_repository
.update_key(&admin_replacement)
assert!(provider_catalog_repository
.compare_and_update_key_admin_state(&ProviderCatalogKeyAdminCasUpdate {
expected_encrypted_auth_config: expected_admin_replacement
.encrypted_auth_config
.clone(),
expected_credential: ProviderCatalogKeyOAuthCredentialFence {
encrypted_api_key: expected_admin_replacement.encrypted_api_key.clone(),
auth_type: expected_admin_replacement.auth_type.clone(),
provider_id: expected_admin_replacement.provider_id.clone(),
provider_type: "claude_code".to_string(),
},
key: admin_replacement,
codex_rotation: None,
reset_oauth_runtime: true,
})
.await
.expect("admin replacement should persist");
.expect("admin replacement CAS should run"));
let admin_result = state
.force_local_oauth_refresh_entry(&stale_transport)
@@ -10,6 +10,10 @@ use crate::{AppState, GatewayError};
const RESPONSE_HEADER_RULES_KEY: &str = "response_header_rules";
const RESPONSE_HEADER_RULES_CAMEL_KEY: &str = "responseHeaderRules";
const PROVIDER_RESPONSE_HEADERS_CONTEXT_KEY: &str = "provider_response_headers";
const PROVIDER_REQUEST_STARTED_AT_UNIX_MS_CONTEXT_KEY: &str = "provider_request_started_at_unix_ms";
const PROVIDER_REQUEST_ORDER_ID_CONTEXT_KEY: &str = "provider_request_order_id";
const PROVIDER_RESPONSE_HEADERS_OBSERVED_AT_UNIX_MS_CONTEXT_KEY: &str =
"provider_response_headers_observed_at_unix_ms";
const RESPONSE_HEADER_RULE_PROTECTED_KEYS: &[&str] = &["content-length"];
const RESPONSE_HEADER_RULES_CACHE_TTL: Duration = Duration::from_secs(5);
@@ -98,6 +102,9 @@ pub(crate) async fn apply_endpoint_response_header_rules(
pub(crate) fn attach_provider_response_headers_to_report_context(
report_context: Option<Value>,
provider_headers: &BTreeMap<String, String>,
provider_request_started_at_unix_ms: u64,
provider_response_headers_observed_at_unix_ms: u64,
provider_request_order_id: &str,
) -> Option<Value> {
let provider_headers = serde_json::to_value(provider_headers).ok()?;
let mut object = match report_context {
@@ -105,9 +112,99 @@ pub(crate) fn attach_provider_response_headers_to_report_context(
Some(other) => Map::from_iter([("seed".to_string(), other)]),
None => Map::new(),
};
object.insert(
PROVIDER_RESPONSE_HEADERS_CONTEXT_KEY.to_string(),
provider_headers,
);
let observation_is_absent = !object.contains_key(PROVIDER_RESPONSE_HEADERS_CONTEXT_KEY)
&& !object.contains_key(PROVIDER_REQUEST_STARTED_AT_UNIX_MS_CONTEXT_KEY)
&& !object.contains_key(PROVIDER_RESPONSE_HEADERS_OBSERVED_AT_UNIX_MS_CONTEXT_KEY)
&& !object.contains_key(PROVIDER_REQUEST_ORDER_ID_CONTEXT_KEY);
if observation_is_absent {
object.insert(
PROVIDER_RESPONSE_HEADERS_CONTEXT_KEY.to_string(),
provider_headers,
);
object.insert(
PROVIDER_REQUEST_STARTED_AT_UNIX_MS_CONTEXT_KEY.to_string(),
Value::from(provider_request_started_at_unix_ms),
);
object.insert(
PROVIDER_RESPONSE_HEADERS_OBSERVED_AT_UNIX_MS_CONTEXT_KEY.to_string(),
Value::from(provider_response_headers_observed_at_unix_ms),
);
object.insert(
PROVIDER_REQUEST_ORDER_ID_CONTEXT_KEY.to_string(),
Value::from(provider_request_order_id),
);
}
Some(Value::Object(object))
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn provider_response_observation_is_first_write_wins() {
let first_headers =
BTreeMap::from([("x-codex-primary-used-percent".to_string(), "10".to_string())]);
let second_headers =
BTreeMap::from([("x-codex-primary-used-percent".to_string(), "20".to_string())]);
let report_context = attach_provider_response_headers_to_report_context(
Some(json!("seed-value")),
&first_headers,
100,
200,
"observation-1",
);
let report_context = attach_provider_response_headers_to_report_context(
report_context,
&second_headers,
300,
400,
"observation-2",
)
.expect("report context should exist");
assert_eq!(report_context["seed"], json!("seed-value"));
assert_eq!(
report_context["provider_response_headers"]["x-codex-primary-used-percent"],
json!("10")
);
assert_eq!(
report_context["provider_request_started_at_unix_ms"],
json!(100)
);
assert_eq!(
report_context["provider_response_headers_observed_at_unix_ms"],
json!(200)
);
assert_eq!(
report_context["provider_request_order_id"],
json!("observation-1")
);
}
#[test]
fn provider_response_observation_does_not_complete_a_partial_triplet() {
let report_context = attach_provider_response_headers_to_report_context(
Some(json!({"provider_response_headers": {"x-existing": "1"}})),
&BTreeMap::from([("x-new".to_string(), "2".to_string())]),
300,
400,
"observation-2",
)
.expect("report context should exist");
assert_eq!(
report_context["provider_response_headers"]["x-existing"],
json!("1")
);
assert!(report_context
.get("provider_request_started_at_unix_ms")
.is_none());
assert!(report_context
.get("provider_response_headers_observed_at_unix_ms")
.is_none());
assert!(report_context.get("provider_request_order_id").is_none());
}
}
@@ -11,8 +11,8 @@ use std::time::{Duration, Instant};
use aether_ai_serving::{AiAttemptExecutionOutcome, AiAttemptRetryScope};
use aether_contracts::{
ExecutionPlan, ExecutionStreamTerminalSummary, ExecutionTelemetry, StandardizedUsage,
StreamFrame, StreamFramePayload,
ExecutionPlan, ExecutionResponseObservation, ExecutionStreamTerminalSummary,
ExecutionTelemetry, StandardizedUsage, StreamFrame, StreamFramePayload,
};
use aether_data_contracts::repository::candidates::{
RequestCandidateStatus, UpsertRequestCandidateRecord,
@@ -112,12 +112,13 @@ use crate::execution_runtime::{
use crate::log_ids::short_request_id;
use crate::orchestration::{
apply_local_execution_effect, build_local_error_flow_metadata, classify_failure_disposition,
cyber_continue_failover_enabled, trace_upstream_response_body, with_error_flow_report_context,
cyber_continue_failover_enabled, spawn_local_oauth_success_effect,
trace_upstream_response_body, with_error_flow_report_context,
with_upstream_response_report_context, FailureDisposition, FailureTokenAction,
LocalAdaptiveRateLimitEffect, LocalAdaptiveSuccessEffect, LocalAttemptFailureEffect,
LocalExecutionEffect, LocalExecutionEffectContext, LocalFailoverAnalysis,
LocalHealthFailureEffect, LocalHealthSuccessEffect, LocalOAuthInvalidationEffect,
LocalPoolErrorEffect,
LocalOAuthSuccessEffect, LocalPoolErrorEffect,
};
use crate::provider_pool_demand::{
acquire_provider_pool_in_flight_guard, ProviderPoolInFlightGuard,
@@ -1249,6 +1250,9 @@ async fn execute_in_process_stream_with_oauth_retry(
retry_status_code,
response_text.as_deref(),
trace_id,
report_context,
Some(execution.response_observation.request_started_at_unix_ms),
Some(&execution.response_observation.request_order_id),
)
.await
{
@@ -2818,6 +2822,7 @@ async fn execute_stream_from_direct_passthrough(
stream_precommit_committed: _,
response,
started_at: upstream_started_at,
response_observation,
stream_first_byte_timeout,
upstream_target_permit,
} = execution;
@@ -2834,8 +2839,23 @@ async fn execute_stream_from_direct_passthrough(
let request_id = plan.request_id.clone();
let candidate_id = plan.candidate_id.clone();
let request_id_for_log = short_request_id(request_id.as_str());
let mut report_context =
attach_provider_response_headers_to_report_context(report_context, &headers);
let mut report_context = attach_provider_response_headers_to_report_context(
report_context,
&headers,
response_observation.request_started_at_unix_ms,
response_observation.response_headers_observed_at_unix_ms,
&response_observation.request_order_id,
);
spawn_local_oauth_success_effect(
state.clone(),
&plan,
report_context.as_ref(),
LocalOAuthSuccessEffect {
status_code,
request_started_at_unix_ms: Some(response_observation.request_started_at_unix_ms),
request_order_id: Some(&response_observation.request_order_id),
},
);
if status_code == 200 {
seed_kiro_simulated_cache_enabled(state, &plan, &mut report_context).await;
if kiro_simulated_cache_enabled_from_report_context(report_context.as_ref()) {
@@ -3819,6 +3839,7 @@ async fn execute_execution_runtime_stream_inner(
provider_pool_in_flight_guard.take(),
retry_scope_out.as_deref_mut(),
retry_fallback_out.as_deref_mut(),
None,
)
.await;
}
@@ -3891,6 +3912,7 @@ async fn execute_execution_runtime_stream_inner(
provider_pool_in_flight_guard.take(),
retry_scope_out.as_deref_mut(),
retry_fallback_out.as_deref_mut(),
None,
)
.await;
}
@@ -3963,6 +3985,7 @@ async fn execute_execution_runtime_stream_inner(
provider_pool_in_flight_guard.take(),
retry_scope_out.as_deref_mut(),
retry_fallback_out.as_deref_mut(),
None,
)
.await;
}
@@ -4035,6 +4058,7 @@ async fn execute_execution_runtime_stream_inner(
provider_pool_in_flight_guard.take(),
retry_scope_out.as_deref_mut(),
retry_fallback_out.as_deref_mut(),
None,
)
.await;
}
@@ -4192,6 +4216,15 @@ async fn execute_execution_runtime_stream_inner(
record_stream_pending_lifecycle(state, seed, &mut stage_trace).await;
lifecycle_pending_recorded = true;
}
let report_context = attach_provider_response_headers_to_report_context(
report_context,
&execution.headers,
execution.response_observation.request_started_at_unix_ms,
execution
.response_observation
.response_headers_observed_at_unix_ms,
&execution.response_observation.request_order_id,
);
let stream_precommit_committed = execution.stream_precommit_committed;
let frame_stream = build_direct_execution_frame_stream(execution).boxed();
return execute_stream_from_frame_stream_with_retry_scope(
@@ -4211,6 +4244,7 @@ async fn execute_execution_runtime_stream_inner(
provider_pool_in_flight_guard.take(),
retry_scope_out,
retry_fallback_out,
None,
)
.await;
}
@@ -4327,6 +4361,15 @@ async fn execute_execution_runtime_stream_inner(
record_stream_pending_lifecycle(state, seed, &mut stage_trace).await;
lifecycle_pending_recorded = true;
}
let report_context = attach_provider_response_headers_to_report_context(
report_context,
&execution.headers,
execution.response_observation.request_started_at_unix_ms,
execution
.response_observation
.response_headers_observed_at_unix_ms,
&execution.response_observation.request_order_id,
);
let stream_precommit_committed = execution.stream_precommit_committed;
let frame_stream = build_direct_execution_frame_stream(execution).boxed();
return execute_stream_from_frame_stream_with_retry_scope(
@@ -4346,10 +4389,13 @@ async fn execute_execution_runtime_stream_inner(
provider_pool_in_flight_guard.take(),
retry_scope_out.as_deref_mut(),
retry_fallback_out.as_deref_mut(),
None,
)
.await;
}
let remote_request_started_at_unix_ms = current_request_candidate_unix_ms();
let remote_request_order_id = uuid::Uuid::now_v7().to_string();
let response = match post_stream_plan_to_remote_execution_runtime(
state,
remote_execution_runtime_base_url,
@@ -4431,6 +4477,12 @@ async fn execute_execution_runtime_stream_inner(
)?));
}
let remote_response_observed_at_unix_ms = current_request_candidate_unix_ms();
let remote_fallback_observation = ExecutionResponseObservation {
request_started_at_unix_ms: remote_request_started_at_unix_ms,
response_headers_observed_at_unix_ms: remote_response_observed_at_unix_ms,
request_order_id: remote_request_order_id,
};
let frame_stream = response
.bytes_stream()
.map_err(|err| IoError::other(err.to_string()))
@@ -4452,6 +4504,7 @@ async fn execute_execution_runtime_stream_inner(
provider_pool_in_flight_guard.take(),
retry_scope_out.as_deref_mut(),
retry_fallback_out.as_deref_mut(),
Some(remote_fallback_observation),
)
.await;
}
@@ -5481,6 +5534,7 @@ async fn execute_stream_from_frame_stream(
in_flight_guard,
None,
None,
None,
)
.await
}
@@ -5503,6 +5557,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
in_flight_guard: Option<ProviderPoolInFlightGuard>,
mut retry_scope_out: Option<&mut AiAttemptRetryScope>,
mut retry_fallback_out: Option<&mut Option<Response<Body>>>,
fallback_response_observation: Option<ExecutionResponseObservation>,
) -> Result<Option<Response<Body>>, GatewayError> {
let request_id = plan.request_id.as_str();
let request_id_for_log = short_request_id(request_id);
@@ -5535,14 +5590,37 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
let StreamFramePayload::Headers {
status_code,
mut headers,
response_observation,
} = first_frame.payload
else {
return Err(GatewayError::Internal(
"execution runtime stream must start with headers frame".to_string(),
));
};
let mut report_context =
attach_provider_response_headers_to_report_context(report_context, &headers);
let response_observation = response_observation
.or(fallback_response_observation)
.unwrap_or(ExecutionResponseObservation {
request_started_at_unix_ms: candidate_started_unix_secs,
response_headers_observed_at_unix_ms: current_request_candidate_unix_ms(),
request_order_id: uuid::Uuid::now_v7().to_string(),
});
let mut report_context = attach_provider_response_headers_to_report_context(
report_context,
&headers,
response_observation.request_started_at_unix_ms,
response_observation.response_headers_observed_at_unix_ms,
&response_observation.request_order_id,
);
spawn_local_oauth_success_effect(
state.clone(),
&plan,
report_context.as_ref(),
LocalOAuthSuccessEffect {
status_code,
request_started_at_unix_ms: Some(response_observation.request_started_at_unix_ms),
request_order_id: Some(&response_observation.request_order_id),
},
);
if status_code == 200 {
seed_kiro_simulated_cache_enabled(state, &plan, &mut report_context).await;
if kiro_simulated_cache_enabled_from_report_context(report_context.as_ref()) {
@@ -8310,6 +8388,7 @@ mod tests {
"content-type".to_string(),
"text/event-stream".to_string(),
)]),
response_observation: None,
},
}));
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
@@ -8389,6 +8468,7 @@ mod tests {
"content-type".to_string(),
"text/event-stream".to_string(),
)]),
response_observation: None,
},
}));
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
@@ -8431,6 +8511,7 @@ mod tests {
None,
Some(&mut retry_scope),
None,
None,
)
.await
.expect("prefetch transport execution should resolve");
@@ -8480,6 +8561,7 @@ mod tests {
"content-type".to_string(),
"text/event-stream".to_string(),
)]),
response_observation: None,
},
}));
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
@@ -8522,6 +8604,7 @@ mod tests {
None,
Some(&mut retry_scope),
None,
None,
)
.await
.expect("prefetch HTTP status execution should resolve");
@@ -8680,6 +8763,7 @@ mod tests {
"content-type".to_string(),
"text/event-stream".to_string(),
)]),
response_observation: None,
},
}));
for chunk in chunks {
@@ -8725,6 +8809,7 @@ mod tests {
None,
Some(&mut retry_scope),
Some(&mut fallback_response),
None,
)
.await
.expect("native Anthropic stream execution should succeed");
@@ -9364,6 +9449,7 @@ mod tests {
"content-type".to_string(),
"text/event-stream".to_string(),
)]),
response_observation: None,
},
}));
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
@@ -9413,6 +9499,7 @@ mod tests {
None,
None,
None,
None,
),
)
.await
@@ -9850,6 +9937,7 @@ mod tests {
"content-type".to_string(),
"text/event-stream".to_string(),
)]),
response_observation: None,
},
}));
}
@@ -11532,6 +11620,7 @@ mod tests {
"content-type".to_string(),
"text/event-stream".to_string(),
)]),
response_observation: None,
},
}));
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
@@ -11660,6 +11749,7 @@ mod tests {
"content-type".to_string(),
"text/event-stream".to_string(),
)]),
response_observation: None,
},
}));
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
@@ -12384,6 +12474,7 @@ mod tests {
"content-type".to_string(),
"text/event-stream".to_string(),
)]),
response_observation: None,
},
}));
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
@@ -4,8 +4,9 @@ use std::io::Error as IoError;
use std::time::{Duration, Instant};
use aether_contracts::{
ExecutionError, ExecutionErrorKind, ExecutionPhase, ExecutionStreamTerminalSummary,
ExecutionTelemetry, StreamFrame, StreamFramePayload, StreamFrameType,
ExecutionError, ExecutionErrorKind, ExecutionPhase, ExecutionResponseObservation,
ExecutionStreamTerminalSummary, ExecutionTelemetry, StreamFrame, StreamFramePayload,
StreamFrameType,
};
use async_stream::stream;
use axum::body::Bytes;
@@ -44,6 +45,7 @@ pub(crate) fn build_direct_execution_frame_stream(
stream_precommit_committed: _,
response,
started_at,
response_observation,
stream_first_byte_timeout,
upstream_target_permit,
} = execution;
@@ -108,7 +110,11 @@ pub(crate) fn build_direct_execution_frame_stream(
}
}
match encode_headers_frame(status_code, response_headers) {
match encode_headers_frame(
status_code,
response_headers,
&response_observation,
) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
@@ -153,7 +159,11 @@ pub(crate) fn build_direct_execution_frame_stream(
upstream_bytes,
first_byte_timeout,
}) => {
match encode_headers_frame(status_code, original_headers) {
match encode_headers_frame(
status_code,
original_headers,
&response_observation,
) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
@@ -192,7 +202,11 @@ pub(crate) fn build_direct_execution_frame_stream(
return;
}
match encode_headers_frame(status_code, headers) {
match encode_headers_frame(
status_code,
headers,
&response_observation,
) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
@@ -611,12 +625,14 @@ pub(crate) fn build_direct_execution_frame_stream(
fn encode_headers_frame(
status_code: u16,
headers: BTreeMap<String, String>,
response_observation: &ExecutionResponseObservation,
) -> Result<Bytes, IoError> {
encode_stream_frame_ndjson(&StreamFrame {
frame_type: StreamFrameType::Headers,
payload: StreamFramePayload::Headers {
status_code,
headers,
response_observation: Some(response_observation.clone()),
},
})
}
@@ -1606,43 +1622,47 @@ mod tests {
.expect("listener should bind");
let addr = listener.local_addr().expect("local addr should resolve");
let server = tokio::spawn(async move {
let app = Router::new().route(
"/responses",
post(|| async {
let body = serde_json::json!({
"id": "resp_sync_bridge_123",
"object": "response",
"model": "gpt-5.4",
"status": "completed",
"output": [{
"type": "message",
"id": "msg_sync_bridge_123",
"role": "assistant",
"content": [{
"type": "output_text",
"text": "Hello from buffered JSON stream",
"annotations": []
}]
}],
"usage": {
"input_tokens": 1,
"output_tokens": 2,
"total_tokens": 3
}
});
let mut response = axum::http::Response::new(Body::from(
serde_json::to_vec(&body).expect("json should encode"),
));
response.headers_mut().insert(
header::CONTENT_TYPE,
HeaderValue::from_static("application/json"),
);
response
}),
);
axum::serve(listener, app)
let (mut socket, _) = listener.accept().await.expect("client should connect");
let mut request = [0_u8; 4096];
let _ = socket
.read(&mut request)
.await
.expect("server should start");
.expect("request should read");
let body = serde_json::to_vec(&serde_json::json!({
"id": "resp_sync_bridge_123",
"object": "response",
"model": "gpt-5.4",
"status": "completed",
"output": [{
"type": "message",
"id": "msg_sync_bridge_123",
"role": "assistant",
"content": [{
"type": "output_text",
"text": "Hello from buffered JSON stream",
"annotations": []
}]
}],
"usage": {
"input_tokens": 1,
"output_tokens": 2,
"total_tokens": 3
}
}))
.expect("json should encode");
socket
.write_all(
format!(
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\n\r\n",
body.len()
)
.as_bytes(),
)
.await
.expect("headers should write");
socket.flush().await.expect("headers should flush");
tokio::time::sleep(Duration::from_millis(75)).await;
socket.write_all(&body).await.expect("body should write");
});
let runtime = DirectSyncExecutionRuntime::new();
@@ -1678,6 +1698,12 @@ mod tests {
})
.await
.expect("stream execution should succeed");
let expected_observation = execution.response_observation.clone();
assert!(
expected_observation.response_headers_observed_at_unix_ms
>= expected_observation.request_started_at_unix_ms
);
assert!(!expected_observation.request_order_id.is_empty());
let frames = build_direct_execution_frame_stream(execution)
.map(|item| item.expect("frame should encode"))
@@ -1691,6 +1717,10 @@ mod tests {
let header_frame: Value =
serde_json::from_str(&frames[0]).expect("headers frame should parse");
let encoded_observation: aether_contracts::ExecutionResponseObservation =
serde_json::from_value(header_frame["payload"]["response_observation"].clone())
.expect("headers frame should retain the response observation");
assert_eq!(encoded_observation, expected_observation);
assert_eq!(
header_frame
.get("payload")
@@ -5,8 +5,8 @@ use std::time::{Duration, Instant};
use aether_ai_serving::{AiAttemptExecutionOutcome, AiAttemptRetryScope, UPSTREAM_IS_STREAM_KEY};
use aether_contracts::{
ExecutionError, ExecutionErrorKind, ExecutionPhase, ExecutionPlan, ExecutionResult,
ExecutionTelemetry,
ExecutionError, ExecutionErrorKind, ExecutionPhase, ExecutionPlan,
ExecutionResponseObservation, ExecutionResult, ExecutionTelemetry,
};
use aether_data_contracts::repository::candidates::RequestCandidateStatus;
use aether_scheduler_core::{
@@ -55,8 +55,9 @@ use crate::execution_runtime::submission::{
resolve_local_sync_error_status_code, submit_local_core_error_or_sync_finalize,
};
use crate::execution_runtime::transport::{
append_upstream_response_body_chunk, build_execution_response_body, build_request_body,
collect_response_headers, decode_response_body_bytes, execution_response_body_mode,
append_upstream_response_body_chunk_with_limit, build_execution_response_body,
build_request_body, collect_response_headers, decode_response_body_bytes_with_limit,
execution_plan_response_body_limit_bytes, execution_response_body_mode,
format_hyper_error_chain, format_upstream_request_error, format_wreq_upstream_request_error,
response_body_is_json, send_request, DirectHttpResponse, DirectSyncExecutionRuntime,
ExecutionRuntimeTransportError,
@@ -70,11 +71,12 @@ use crate::execution_runtime::{
};
use crate::log_ids::short_request_id;
use crate::orchestration::{
apply_local_execution_effect, build_local_error_flow_metadata, trace_upstream_response_body,
with_error_flow_report_context, with_upstream_response_report_context,
LocalAdaptiveRateLimitEffect, LocalAdaptiveSuccessEffect, LocalAttemptFailureEffect,
LocalExecutionEffect, LocalExecutionEffectContext, LocalHealthFailureEffect,
LocalHealthSuccessEffect, LocalOAuthInvalidationEffect, LocalPoolErrorEffect,
apply_local_execution_effect, build_local_error_flow_metadata,
spawn_local_oauth_success_effect, trace_upstream_response_body, with_error_flow_report_context,
with_upstream_response_report_context, LocalAdaptiveRateLimitEffect,
LocalAdaptiveSuccessEffect, LocalAttemptFailureEffect, LocalExecutionEffect,
LocalExecutionEffectContext, LocalHealthFailureEffect, LocalHealthSuccessEffect,
LocalOAuthInvalidationEffect, LocalOAuthSuccessEffect, LocalPoolErrorEffect,
};
use crate::provider_pool_demand::acquire_provider_pool_in_flight_guard;
use crate::request_candidate_runtime::{
@@ -1379,7 +1381,19 @@ async fn execute_direct_sync_runtime_candidate(
candidate_started_unix_ms,
event.status_code,
event.ttfb_ms,
)
);
spawn_local_oauth_success_effect(
state_for_response_started.clone(),
plan,
report_context,
LocalOAuthSuccessEffect {
status_code: event.status_code,
request_started_at_unix_ms: Some(
event.response_observation.request_started_at_unix_ms,
),
request_order_id: Some(&event.response_observation.request_order_id),
},
);
})
.await
.map_err(SyncExecutionFailure::from_transport);
@@ -1478,17 +1492,31 @@ async fn execute_openai_image_sync_upstream_sse_candidate(
progress_snapshot: Option<Arc<Mutex<OpenAiImageSyncProgressSnapshot>>>,
) -> Result<ExecutionResult, SyncExecutionFailure> {
let request_body = build_request_body(plan).map_err(SyncExecutionFailure::from_transport)?;
let response_body_limit_bytes = execution_plan_response_body_limit_bytes(plan);
let started_at = Instant::now();
let mut progress =
OpenAiImageSyncProgressRecorder::new(state, plan, report_context, progress_snapshot);
progress.record_connecting().await;
let request_started_at_unix_ms = current_request_candidate_unix_ms();
let request_order_id = uuid::Uuid::now_v7().to_string();
let response = send_request(plan, request_body)
.await
.map_err(SyncExecutionFailure::from_transport)?;
let ttfb_ms = started_at.elapsed().as_millis() as u64;
let response_headers_observed_at_unix_ms = current_request_candidate_unix_ms();
let status_code = response.status_code();
let headers = response.headers();
spawn_local_oauth_success_effect(
state.clone(),
plan,
report_context,
LocalOAuthSuccessEffect {
status_code,
request_started_at_unix_ms: Some(request_started_at_unix_ms),
request_order_id: Some(&request_order_id),
},
);
progress.record_response_started(status_code, ttfb_ms).await;
let mut body_bytes = Vec::new();
@@ -1503,8 +1531,12 @@ async fn execute_openai_image_sync_upstream_sse_candidate(
),
)
})?;
append_upstream_response_body_chunk(&mut body_bytes, &chunk)
.map_err(SyncExecutionFailure::from_transport)?;
append_upstream_response_body_chunk_with_limit(
&mut body_bytes,
&chunk,
response_body_limit_bytes,
)
.map_err(SyncExecutionFailure::from_transport)?;
let elapsed_ms = started_at.elapsed().as_millis() as u64;
progress
.observe_chunk(&chunk, status_code, elapsed_ms)
@@ -1521,8 +1553,12 @@ async fn execute_openai_image_sync_upstream_sse_candidate(
)),
)
})?;
append_upstream_response_body_chunk(&mut body_bytes, &chunk)
.map_err(SyncExecutionFailure::from_transport)?;
append_upstream_response_body_chunk_with_limit(
&mut body_bytes,
&chunk,
response_body_limit_bytes,
)
.map_err(SyncExecutionFailure::from_transport)?;
let elapsed_ms = started_at.elapsed().as_millis() as u64;
progress
.observe_chunk(&chunk, status_code, elapsed_ms)
@@ -1539,8 +1575,12 @@ async fn execute_openai_image_sync_upstream_sse_candidate(
),
)
})?;
append_upstream_response_body_chunk(&mut body_bytes, &chunk)
.map_err(SyncExecutionFailure::from_transport)?;
append_upstream_response_body_chunk_with_limit(
&mut body_bytes,
&chunk,
response_body_limit_bytes,
)
.map_err(SyncExecutionFailure::from_transport)?;
let elapsed_ms = started_at.elapsed().as_millis() as u64;
progress
.observe_chunk(&chunk, status_code, elapsed_ms)
@@ -1549,8 +1589,9 @@ async fn execute_openai_image_sync_upstream_sse_candidate(
}
}
let decoded_body_bytes = decode_response_body_bytes(&headers, &body_bytes)
.map_err(SyncExecutionFailure::from_transport)?;
let decoded_body_bytes =
decode_response_body_bytes_with_limit(&headers, &body_bytes, response_body_limit_bytes)
.map_err(SyncExecutionFailure::from_transport)?;
let elapsed_ms = started_at.elapsed().as_millis() as u64;
let upstream_bytes = body_bytes.len() as u64;
progress.finish(status_code, elapsed_ms).await;
@@ -1569,6 +1610,11 @@ async fn execute_openai_image_sync_upstream_sse_candidate(
candidate_id: plan.candidate_id.clone(),
status_code,
headers,
response_observation: Some(ExecutionResponseObservation {
request_started_at_unix_ms,
response_headers_observed_at_unix_ms,
request_order_id,
}),
body,
telemetry: Some(ExecutionTelemetry {
ttfb_ms: Some(ttfb_ms),
@@ -2461,6 +2507,16 @@ async fn execute_execution_runtime_sync_impl(
};
let mut candidate_first_byte_elapsed_ms =
calibrated_sync_candidate_first_byte_elapsed_ms(candidate_started_at, &result);
let initial_response_observed_at_unix_ms = current_request_candidate_unix_ms();
let mut provider_response_observation =
result
.response_observation
.clone()
.unwrap_or(ExecutionResponseObservation {
request_started_at_unix_ms: candidate_started_unix_secs,
response_headers_observed_at_unix_ms: initial_response_observed_at_unix_ms,
request_order_id: uuid::Uuid::now_v7().to_string(),
});
let mut oauth_retry_attempted = false;
let (
result_error_type,
@@ -2473,6 +2529,18 @@ async fn execute_execution_runtime_sync_impl(
local_failover_response_text,
local_failover_analysis,
) = loop {
spawn_local_oauth_success_effect(
state.clone(),
&plan,
report_context.as_ref(),
LocalOAuthSuccessEffect {
status_code: result.status_code,
request_started_at_unix_ms: Some(
provider_response_observation.request_started_at_unix_ms,
),
request_order_id: Some(&provider_response_observation.request_order_id),
},
);
let result_latency_ms = result
.telemetry
.as_ref()
@@ -2534,10 +2602,15 @@ async fn execute_execution_runtime_sync_impl(
result.status_code,
local_failover_response_text.as_deref(),
trace_id,
report_context.as_ref(),
Some(provider_response_observation.request_started_at_unix_ms),
Some(&provider_response_observation.request_order_id),
)
.await
{
oauth_retry_attempted = true;
let retry_started_at_unix_ms = current_request_candidate_unix_ms();
let retry_request_order_id = uuid::Uuid::now_v7().to_string();
match crate::execution_runtime::execute_execution_runtime_sync_plan(
state,
Some(trace_id),
@@ -2546,6 +2619,16 @@ async fn execute_execution_runtime_sync_impl(
.await
{
Ok(retry_result) => {
let retry_response_observed_at_unix_ms = current_request_candidate_unix_ms();
provider_response_observation = retry_result
.response_observation
.clone()
.unwrap_or(ExecutionResponseObservation {
request_started_at_unix_ms: retry_started_at_unix_ms,
response_headers_observed_at_unix_ms:
retry_response_observed_at_unix_ms,
request_order_id: retry_request_order_id,
});
candidate_first_byte_elapsed_ms =
calibrated_sync_candidate_first_byte_elapsed_ms(
candidate_started_at,
@@ -2594,6 +2677,13 @@ async fn execute_execution_runtime_sync_impl(
local_failover_analysis,
);
};
let mut report_context = attach_provider_response_headers_to_report_context(
report_context,
&headers,
provider_response_observation.request_started_at_unix_ms,
provider_response_observation.response_headers_observed_at_unix_ms,
&provider_response_observation.request_order_id,
);
if result.status_code >= 400 {
apply_local_execution_effect(
state,
@@ -2739,8 +2829,6 @@ async fn execute_execution_runtime_sync_impl(
}
let status_code = result.status_code;
let has_body_bytes = body_base64.is_some();
let mut report_context =
attach_provider_response_headers_to_report_context(report_context, &headers);
if (200..300).contains(&status_code) {
seed_kiro_sync_simulated_cache_enabled(state, &plan, &mut report_context).await;
if kiro_simulated_cache_enabled_from_report_context(report_context.as_ref()) {
@@ -3231,6 +3319,8 @@ async fn execute_sync_via_remote_execution_runtime(
candidate_started_unix_secs: u64,
candidate_started_at: Instant,
) -> Result<RemoteSyncFallbackOutcome, GatewayError> {
let remote_request_started_at_unix_ms = current_request_candidate_unix_ms();
let remote_request_order_id = uuid::Uuid::now_v7().to_string();
let response = match post_sync_plan_to_remote_execution_runtime(
state,
remote_execution_runtime_base_url,
@@ -3299,11 +3389,19 @@ async fn execute_sync_via_remote_execution_runtime(
));
}
response
.json()
let remote_response_observed_at_unix_ms = current_request_candidate_unix_ms();
let mut result = response
.json::<ExecutionResult>()
.await
.map(RemoteSyncFallbackOutcome::Executed)
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
result
.response_observation
.get_or_insert(ExecutionResponseObservation {
request_started_at_unix_ms: remote_request_started_at_unix_ms,
response_headers_observed_at_unix_ms: remote_response_observed_at_unix_ms,
request_order_id: remote_request_order_id,
});
Ok(RemoteSyncFallbackOutcome::Executed(result))
}
#[cfg(test)]
@@ -9,18 +9,19 @@ use std::sync::{Arc, LazyLock, Mutex as StdMutex, OnceLock, RwLock as StdRwLock}
use std::time::{Duration, Instant};
use aether_contracts::{
ExecutionPlan, ExecutionResponseBodyMode, ExecutionResult, ExecutionTelemetry, ProxySnapshot,
ResolvedTransportProfile, ResponseBody, EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER,
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, EXECUTION_REQUEST_HTTP1_ONLY_HEADER,
EXECUTION_RESPONSE_BODY_MODE_HEADER, TRANSPORT_BACKEND_BROWSER_WREQ,
TRANSPORT_BACKEND_REQWEST_RUSTLS, TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE,
TRANSPORT_HTTP_MODE_HTTP1_ONLY,
ExecutionPlan, ExecutionResponseBodyMode, ExecutionResponseObservation, ExecutionResult,
ExecutionTelemetry, ProxySnapshot, ResolvedTransportProfile, ResponseBody,
EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER, EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER,
EXECUTION_REQUEST_HTTP1_ONLY_HEADER, EXECUTION_RESPONSE_BODY_MODE_HEADER,
TRANSPORT_BACKEND_BROWSER_WREQ, TRANSPORT_BACKEND_REQWEST_RUSTLS,
TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE, TRANSPORT_HTTP_MODE_HTTP1_ONLY,
};
use aether_data::repository::proxy_nodes::ProxyNodeTrafficMutation;
use aether_http::{apply_http_client_config, HttpClientConfig};
use aether_runtime::{MetricKind, MetricSample};
use axum::body::Bytes;
use base64::Engine as _;
use brotli::Decompressor as BrotliDecoder;
use flate2::read::{DeflateDecoder, GzDecoder};
use flate2::write::GzEncoder;
use flate2::Compression;
@@ -62,6 +63,10 @@ const DEFAULT_STREAM_FIRST_BYTE_TIMEOUT_MS: u64 = 30_000;
const DEFAULT_NON_STREAM_TOTAL_TIMEOUT_MS: u64 = 300_000;
const DEFAULT_CODEX_COMPACT_TOTAL_TIMEOUT_MS: u64 = 1_200_000;
const MIN_TUNNEL_TIMEOUT_SECS: u64 = 1;
const EXECUTION_RESPONSE_BODY_LIMIT_HEADER: &str = "x-aether-execution-response-body-limit-bytes";
const DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES: usize = 8 * 1024 * 1024;
const MIN_SCOPED_RESPONSE_BODY_LIMIT_BYTES: usize = 64 * 1024;
const MAX_SCOPED_RESPONSE_BODY_LIMIT_BYTES: usize = 64 * 1024 * 1024;
const DIRECT_REQWEST_H2_CLIENT_SHARDS_ENV: &str = "AETHER_GATEWAY_DIRECT_REQWEST_H2_CLIENT_SHARDS";
const DIRECT_REQWEST_CLIENT_SHARDS_ENV: &str = "AETHER_GATEWAY_DIRECT_REQWEST_CLIENT_SHARDS";
const DIRECT_REQWEST_H2_TARGET_STREAMS_PER_CLIENT_ENV: &str =
@@ -607,6 +612,56 @@ impl std::fmt::Display for UpstreamResponseBodyPhase {
}
}
pub(crate) fn with_upstream_response_body_limit(
plan: &ExecutionPlan,
limit_bytes: usize,
) -> ExecutionPlan {
let mut bounded_plan = plan.clone();
bounded_plan
.headers
.retain(|name, _| !name.eq_ignore_ascii_case(EXECUTION_RESPONSE_BODY_LIMIT_HEADER));
bounded_plan.headers.insert(
EXECUTION_RESPONSE_BODY_LIMIT_HEADER.to_string(),
normalize_scoped_response_body_limit(limit_bytes)
.unwrap_or(DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES)
.to_string(),
);
bounded_plan
}
pub(crate) fn execution_plan_response_body_limit_bytes(plan: &ExecutionPlan) -> usize {
effective_response_body_limit_bytes(
execution_transport_header_value(&plan.headers, EXECUTION_RESPONSE_BODY_LIMIT_HEADER),
crate::headers::max_internal_buffered_body_bytes(),
)
}
fn effective_response_body_limit_bytes(
raw_scoped_limit: Option<&str>,
global_limit: usize,
) -> usize {
let Some(raw_scoped_limit) = raw_scoped_limit else {
return global_limit;
};
parse_scoped_response_body_limit(raw_scoped_limit)
.unwrap_or(DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES)
.min(global_limit)
}
fn parse_scoped_response_body_limit(value: &str) -> Option<usize> {
let raw_limit = value.trim().parse::<u64>().ok()?;
usize::try_from(raw_limit)
.ok()
.and_then(normalize_scoped_response_body_limit)
}
fn normalize_scoped_response_body_limit(limit_bytes: usize) -> Option<usize> {
(limit_bytes > 0).then_some(limit_bytes.clamp(
MIN_SCOPED_RESPONSE_BODY_LIMIT_BYTES,
MAX_SCOPED_RESPONSE_BODY_LIMIT_BYTES,
))
}
pub(crate) fn append_upstream_response_body_chunk(
body: &mut Vec<u8>,
chunk: &[u8],
@@ -618,7 +673,7 @@ pub(crate) fn append_upstream_response_body_chunk(
)
}
fn append_upstream_response_body_chunk_with_limit(
pub(crate) fn append_upstream_response_body_chunk_with_limit(
body: &mut Vec<u8>,
chunk: &[u8],
limit_bytes: usize,
@@ -691,6 +746,7 @@ pub(crate) struct DirectUpstreamStreamExecution {
pub(crate) stream_precommit_committed: bool,
pub(crate) response: DirectUpstreamResponse,
pub(crate) started_at: Instant,
pub(crate) response_observation: ExecutionResponseObservation,
pub(crate) stream_first_byte_timeout: Option<Duration>,
pub(crate) upstream_target_permit: Option<UpstreamTargetAdmissionPermit>,
}
@@ -699,6 +755,7 @@ pub(crate) struct DirectUpstreamStreamExecution {
pub(crate) struct DirectSyncResponseStarted {
pub(crate) status_code: u16,
pub(crate) ttfb_ms: u64,
pub(crate) response_observation: ExecutionResponseObservation,
}
impl DirectSyncExecutionRuntime {
@@ -722,20 +779,35 @@ impl DirectSyncExecutionRuntime {
F: FnOnce(DirectSyncResponseStarted),
{
let body_bytes = build_request_body(plan)?;
let response_body_limit_bytes = execution_plan_response_body_limit_bytes(plan);
let started_at = Instant::now();
let request_started_at_unix_ms = crate::clock::current_unix_ms();
let request_order_id = uuid::Uuid::now_v7().to_string();
with_non_stream_total_timeout(plan, async move {
let response = send_request_inner(plan, body_bytes, false).await?;
let ttfb_ms = started_at.elapsed().as_millis() as u64;
let response_headers_observed_at_unix_ms = crate::clock::current_unix_ms();
let status_code = response.status_code();
let headers = response.headers();
let response_observation = ExecutionResponseObservation {
request_started_at_unix_ms,
response_headers_observed_at_unix_ms,
request_order_id,
};
on_response_started(DirectSyncResponseStarted {
status_code,
ttfb_ms,
response_observation: response_observation.clone(),
});
let (body_bytes, stream_ttfb_ms) =
response.bytes_with_stream_timeout(plan, started_at).await?;
let decoded_body_bytes = decode_response_body_bytes(&headers, &body_bytes)?;
let (body_bytes, stream_ttfb_ms) = response
.bytes_with_stream_timeout(plan, started_at, response_body_limit_bytes)
.await?;
let decoded_body_bytes = decode_response_body_bytes_with_limit(
&headers,
&body_bytes,
response_body_limit_bytes,
)?;
let elapsed_ms = started_at.elapsed().as_millis() as u64;
let upstream_bytes = body_bytes.len() as u64;
@@ -752,6 +824,7 @@ impl DirectSyncExecutionRuntime {
candidate_id: plan.candidate_id.clone(),
status_code,
headers,
response_observation: Some(response_observation),
body,
telemetry: Some(ExecutionTelemetry {
ttfb_ms: stream_ttfb_ms.or(Some(ttfb_ms)),
@@ -776,6 +849,8 @@ impl DirectSyncExecutionRuntime {
);
let started_at = Instant::now();
let request_started_at_unix_ms = crate::clock::current_unix_ms();
let request_order_id = uuid::Uuid::now_v7().to_string();
let response = send_request(plan, body_bytes).await?;
observe_gateway_stage_ms(
"direct_send_headers",
@@ -783,6 +858,7 @@ impl DirectSyncExecutionRuntime {
);
let status_code = response.status_code();
let headers = response.headers();
let response_headers_observed_at_unix_ms = crate::clock::current_unix_ms();
let stream_summary_report_context = build_stream_summary_report_context(plan);
@@ -797,6 +873,11 @@ impl DirectSyncExecutionRuntime {
stream_precommit_committed: false,
response: response.into_direct_upstream_response(),
started_at,
response_observation: ExecutionResponseObservation {
request_started_at_unix_ms,
response_headers_observed_at_unix_ms,
request_order_id,
},
stream_first_byte_timeout: resolve_stream_first_byte_timeout(plan),
upstream_target_permit: None,
})
@@ -834,7 +915,7 @@ pub(crate) async fn execute_sync_plan_with_report_context(
}
if resolve_local_tunnel_node_id(state, plan.proxy.as_ref()).is_some() {
return execute_sync_plan_via_local_tunnel(state, plan)
return execute_sync_plan_via_local_tunnel(state, plan, report_context)
.await
.map_err(|err| GatewayError::Internal(err.to_string()));
}
@@ -857,7 +938,24 @@ pub(crate) async fn execute_sync_plan_with_report_context(
Ok(None) => {}
Err(err) => return Err(GatewayError::Internal(err.to_string())),
}
match DirectSyncExecutionRuntime::new().execute_sync(plan).await {
let state_for_response_started = state.clone();
match DirectSyncExecutionRuntime::new()
.execute_sync_with_response_started(plan, move |event| {
crate::orchestration::spawn_local_oauth_success_effect(
state_for_response_started,
plan,
report_context,
crate::orchestration::LocalOAuthSuccessEffect {
status_code: event.status_code,
request_started_at_unix_ms: Some(
event.response_observation.request_started_at_unix_ms,
),
request_order_id: Some(&event.response_observation.request_order_id),
},
);
})
.await
{
Ok(result) => {
record_manual_proxy_request_outcome(state, plan, result.status_code).await;
Ok(result)
@@ -889,6 +987,8 @@ pub(crate) async fn execute_stream_plan_via_local_tunnel(
plan.body.body_bytes_b64.is_some(),
)?;
let started_at = Instant::now();
let request_started_at_unix_ms = crate::clock::current_unix_ms();
let request_order_id = uuid::Uuid::now_v7().to_string();
let response = state
.tunnel
.open_direct_relay_stream(
@@ -900,6 +1000,7 @@ pub(crate) async fn execute_stream_plan_via_local_tunnel(
.map_err(ExecutionRuntimeTransportError::RelayError)?;
let status_code = response.status();
let headers = collect_tunnel_response_headers(response.headers());
let response_headers_observed_at_unix_ms = crate::clock::current_unix_ms();
Ok(Some(DirectUpstreamStreamExecution {
request_id: plan.request_id.clone(),
@@ -912,6 +1013,11 @@ pub(crate) async fn execute_stream_plan_via_local_tunnel(
stream_precommit_committed: false,
response: DirectUpstreamResponse::LocalTunnel(response),
started_at,
response_observation: ExecutionResponseObservation {
request_started_at_unix_ms,
response_headers_observed_at_unix_ms,
request_order_id,
},
stream_first_byte_timeout: resolve_stream_first_byte_timeout(plan),
upstream_target_permit: None,
}))
@@ -991,13 +1097,19 @@ fn manual_proxy_node_id(proxy: Option<&ProxySnapshot>) -> Option<String> {
async fn execute_sync_plan_via_local_tunnel(
state: &AppState,
plan: &ExecutionPlan,
report_context: Option<&serde_json::Value>,
) -> Result<ExecutionResult, ExecutionRuntimeTransportError> {
with_non_stream_total_timeout(plan, execute_sync_plan_via_local_tunnel_inner(state, plan)).await
with_non_stream_total_timeout(
plan,
execute_sync_plan_via_local_tunnel_inner(state, plan, report_context),
)
.await
}
async fn execute_sync_plan_via_local_tunnel_inner(
state: &AppState,
plan: &ExecutionPlan,
report_context: Option<&serde_json::Value>,
) -> Result<ExecutionResult, ExecutionRuntimeTransportError> {
let node_id = resolve_local_tunnel_node_id(state, plan.proxy.as_ref()).ok_or_else(|| {
ExecutionRuntimeTransportError::RelayError("local tunnel node unavailable".to_string())
@@ -1007,6 +1119,7 @@ async fn execute_sync_plan_via_local_tunnel_inner(
}
let body_bytes = build_request_body(plan)?;
let response_body_limit_bytes = execution_plan_response_body_limit_bytes(plan);
let transport_controls = resolve_execution_transport_controls(&plan.headers);
let headers = build_request_headers(
&plan.headers,
@@ -1030,6 +1143,8 @@ async fn execute_sync_plan_via_local_tunnel_inner(
"gateway execution runtime local tunnel request prepared"
);
let started_at = Instant::now();
let request_started_at_unix_ms = crate::clock::current_unix_ms();
let request_order_id = uuid::Uuid::now_v7().to_string();
let mut response = state
.tunnel
.open_direct_relay_stream(
@@ -1040,12 +1155,30 @@ async fn execute_sync_plan_via_local_tunnel_inner(
.await
.map_err(ExecutionRuntimeTransportError::RelayError)?;
let ttfb_ms = started_at.elapsed().as_millis() as u64;
let response_headers_observed_at_unix_ms = crate::clock::current_unix_ms();
let status_code = response.status();
let headers = collect_tunnel_response_headers(response.headers());
let response_observation = ExecutionResponseObservation {
request_started_at_unix_ms,
response_headers_observed_at_unix_ms,
request_order_id,
};
crate::orchestration::spawn_local_oauth_success_effect(
state.clone(),
plan,
report_context,
crate::orchestration::LocalOAuthSuccessEffect {
status_code,
request_started_at_unix_ms: Some(response_observation.request_started_at_unix_ms),
request_order_id: Some(&response_observation.request_order_id),
},
);
let proxy_timing = execution_header_for_log(&headers, "x-proxy-timing").unwrap_or("-");
let (body_bytes, stream_ttfb_ms) =
collect_local_tunnel_response_body(response, plan, started_at).await?;
let decoded_body_bytes = decode_response_body_bytes(&headers, &body_bytes)?;
collect_local_tunnel_response_body(response, plan, started_at, response_body_limit_bytes)
.await?;
let decoded_body_bytes =
decode_response_body_bytes_with_limit(&headers, &body_bytes, response_body_limit_bytes)?;
let elapsed_ms = started_at.elapsed().as_millis() as u64;
let upstream_bytes = body_bytes.len() as u64;
if status_code >= 400 {
@@ -1095,6 +1228,7 @@ async fn execute_sync_plan_via_local_tunnel_inner(
candidate_id: plan.candidate_id.clone(),
status_code,
headers,
response_observation: Some(response_observation),
body,
telemetry: Some(ExecutionTelemetry {
ttfb_ms: stream_ttfb_ms.or(Some(ttfb_ms)),
@@ -1109,6 +1243,7 @@ async fn collect_local_tunnel_response_body(
mut response: tunnel::DirectRelayResponse,
plan: &ExecutionPlan,
started_at: Instant,
response_body_limit_bytes: usize,
) -> Result<(Vec<u8>, Option<u64>), ExecutionRuntimeTransportError> {
let mut body_bytes = Vec::new();
let mut first_byte_ms = None;
@@ -1131,7 +1266,11 @@ async fn collect_local_tunnel_response_body(
if plan.stream && first_byte_ms.is_none() && !chunk.is_empty() {
first_byte_ms = Some(started_at.elapsed().as_millis() as u64);
}
append_upstream_response_body_chunk(&mut body_bytes, &chunk)?;
append_upstream_response_body_chunk_with_limit(
&mut body_bytes,
&chunk,
response_body_limit_bytes,
)?;
}
Ok((body_bytes, first_byte_ms))
@@ -1292,20 +1431,28 @@ impl DirectHttpResponse {
}
pub(crate) async fn bytes(self) -> Result<Bytes, ExecutionRuntimeTransportError> {
self.bytes_with_limit(crate::headers::max_internal_buffered_body_bytes())
.await
}
async fn bytes_with_limit(
self,
response_body_limit_bytes: usize,
) -> Result<Bytes, ExecutionRuntimeTransportError> {
let started_at = Instant::now();
match self {
DirectHttpResponse::Reqwest(response) => {
collect_reqwest_stream_body(response, started_at, None)
collect_reqwest_stream_body(response, started_at, None, response_body_limit_bytes)
.await
.map(|(body, _)| body)
}
DirectHttpResponse::HyperH2c(response) => {
collect_hyper_stream_body(response, started_at, None)
collect_hyper_stream_body(response, started_at, None, response_body_limit_bytes)
.await
.map(|(body, _)| body)
}
DirectHttpResponse::BrowserWreq(response) => {
collect_wreq_stream_body(response, started_at, None)
collect_wreq_stream_body(response, started_at, None, response_body_limit_bytes)
.await
.map(|(body, _)| body)
}
@@ -1316,21 +1463,43 @@ impl DirectHttpResponse {
self,
plan: &ExecutionPlan,
started_at: Instant,
response_body_limit_bytes: usize,
) -> Result<(Bytes, Option<u64>), ExecutionRuntimeTransportError> {
if !plan.stream {
return self.bytes().await.map(|bytes| (bytes, None));
return self
.bytes_with_limit(response_body_limit_bytes)
.await
.map(|bytes| (bytes, None));
}
let first_byte_timeout = resolve_stream_first_byte_timeout(plan);
match self {
DirectHttpResponse::Reqwest(response) => {
collect_reqwest_stream_body(response, started_at, first_byte_timeout).await
collect_reqwest_stream_body(
response,
started_at,
first_byte_timeout,
response_body_limit_bytes,
)
.await
}
DirectHttpResponse::HyperH2c(response) => {
collect_hyper_stream_body(response, started_at, first_byte_timeout).await
collect_hyper_stream_body(
response,
started_at,
first_byte_timeout,
response_body_limit_bytes,
)
.await
}
DirectHttpResponse::BrowserWreq(response) => {
collect_wreq_stream_body(response, started_at, first_byte_timeout).await
collect_wreq_stream_body(
response,
started_at,
first_byte_timeout,
response_body_limit_bytes,
)
.await
}
}
}
@@ -1376,6 +1545,7 @@ async fn collect_reqwest_stream_body(
response: reqwest::Response,
started_at: Instant,
first_byte_timeout: Option<Duration>,
response_body_limit_bytes: usize,
) -> Result<(Bytes, Option<u64>), ExecutionRuntimeTransportError> {
let mut stream = response.bytes_stream();
let mut body_bytes = Vec::new();
@@ -1396,7 +1566,11 @@ async fn collect_reqwest_stream_body(
if first_byte_ms.is_none() && !chunk.is_empty() {
first_byte_ms = Some(started_at.elapsed().as_millis() as u64);
}
append_upstream_response_body_chunk(&mut body_bytes, &chunk)?;
append_upstream_response_body_chunk_with_limit(
&mut body_bytes,
&chunk,
response_body_limit_bytes,
)?;
}
Ok((Bytes::from(body_bytes), first_byte_ms))
@@ -1406,6 +1580,7 @@ async fn collect_hyper_stream_body(
response: hyper::Response<HyperIncomingBody>,
started_at: Instant,
first_byte_timeout: Option<Duration>,
response_body_limit_bytes: usize,
) -> Result<(Bytes, Option<u64>), ExecutionRuntimeTransportError> {
let mut stream = response.into_body().into_data_stream();
let mut body_bytes = Vec::new();
@@ -1426,7 +1601,11 @@ async fn collect_hyper_stream_body(
if first_byte_ms.is_none() && !chunk.is_empty() {
first_byte_ms = Some(started_at.elapsed().as_millis() as u64);
}
append_upstream_response_body_chunk(&mut body_bytes, &chunk)?;
append_upstream_response_body_chunk_with_limit(
&mut body_bytes,
&chunk,
response_body_limit_bytes,
)?;
}
Ok((Bytes::from(body_bytes), first_byte_ms))
@@ -1436,6 +1615,7 @@ async fn collect_wreq_stream_body(
response: wreq::Response,
started_at: Instant,
first_byte_timeout: Option<Duration>,
response_body_limit_bytes: usize,
) -> Result<(Bytes, Option<u64>), ExecutionRuntimeTransportError> {
let mut stream = response.bytes_stream();
let mut body_bytes = Vec::new();
@@ -1456,7 +1636,11 @@ async fn collect_wreq_stream_body(
if first_byte_ms.is_none() && !chunk.is_empty() {
first_byte_ms = Some(started_at.elapsed().as_millis() as u64);
}
append_upstream_response_body_chunk(&mut body_bytes, &chunk)?;
append_upstream_response_body_chunk_with_limit(
&mut body_bytes,
&chunk,
response_body_limit_bytes,
)?;
}
Ok((Bytes::from(body_bytes), first_byte_ms))
@@ -2308,10 +2492,31 @@ async fn send_via_tunnel_relay(
error_kind = %kind,
"gateway execution runtime tunnel relay returned relay error"
);
let message = response
.text()
.await
.unwrap_or_else(|_| format!("hub relay error: {kind}"));
let response_headers = collect_response_headers(response.headers());
let response_body_limit_bytes = execution_plan_response_body_limit_bytes(plan);
let (wire_body, _) =
collect_reqwest_stream_body(response, Instant::now(), None, response_body_limit_bytes)
.await
.map_err(|error| {
ExecutionRuntimeTransportError::RelayError(format!(
"hub relay error: {kind}: bounded error body read failed: {error}"
))
})?;
let decoded_body = decode_response_body_bytes_with_limit(
&response_headers,
&wire_body,
response_body_limit_bytes,
)
.map_err(|error| {
ExecutionRuntimeTransportError::RelayError(format!(
"hub relay error: {kind}: bounded error body decode failed: {error}"
))
})?;
let message = if decoded_body.is_empty() {
format!("hub relay error: {kind}")
} else {
String::from_utf8_lossy(decoded_body.as_ref()).into_owned()
};
return Err(ExecutionRuntimeTransportError::RelayError(message));
}
@@ -3833,6 +4038,7 @@ pub(crate) fn build_request_headers(
|| normalized_key == EXECUTION_REQUEST_HTTP1_ONLY_HEADER
|| normalized_key == EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER
|| normalized_key == EXECUTION_RESPONSE_BODY_MODE_HEADER
|| normalized_key == EXECUTION_RESPONSE_BODY_LIMIT_HEADER
{
continue;
}
@@ -3981,7 +4187,7 @@ pub(crate) fn decode_response_body_bytes<'a>(
)
}
fn decode_response_body_bytes_with_limit<'a>(
pub(crate) fn decode_response_body_bytes_with_limit<'a>(
headers: &BTreeMap<String, String>,
body_bytes: &'a [u8],
limit_bytes: usize,
@@ -4003,6 +4209,11 @@ fn decode_response_body_bytes_with_limit<'a>(
read_upstream_response_decoder_with_limit("deflate", &mut decoder, limit_bytes)
.map(Cow::Owned)
}
Some("br") => {
let mut decoder = BrotliDecoder::new(body_bytes, 4_096);
read_upstream_response_decoder_with_limit("br", &mut decoder, limit_bytes)
.map(Cow::Owned)
}
_ => Ok(Cow::Borrowed(body_bytes)),
}
}
@@ -4122,12 +4333,16 @@ mod tests {
use super::{
append_upstream_response_body_chunk_with_limit, build_browser_wreq_client, build_client,
build_direct_tunnel_request_meta, build_execution_response_body, build_request_headers,
decode_response_body_bytes_with_limit, execute_sync_plan, execution_response_body_mode,
decode_response_body_bytes_with_limit, effective_response_body_limit_bytes,
execute_sync_plan, execution_plan_response_body_limit_bytes, execution_response_body_mode,
record_manual_proxy_request_failure, record_manual_proxy_request_outcome,
record_manual_proxy_request_success, record_manual_proxy_stream_error,
resolve_execution_transport_controls, resolve_non_stream_total_timeout,
resolve_stream_first_byte_timeout, response_body_is_json, DirectSyncExecutionRuntime,
resolve_stream_first_byte_timeout, response_body_is_json,
with_upstream_response_body_limit, DirectSyncExecutionRuntime,
ExecutionRuntimeTransportError, ExecutionTransportControls, UpstreamResponseBodyPhase,
DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES, EXECUTION_RESPONSE_BODY_LIMIT_HEADER,
MAX_SCOPED_RESPONSE_BODY_LIMIT_BYTES, MIN_SCOPED_RESPONSE_BODY_LIMIT_BYTES,
};
use crate::constants::{
EXECUTION_RUNTIME_LOOP_GUARD_HEADER, EXECUTION_RUNTIME_LOOP_GUARD_VIA_TOKEN,
@@ -4182,6 +4397,162 @@ mod tests {
assert!(!materialized.contains_key("x-aether-future-control"));
}
#[test]
fn scoped_response_body_limit_injection_preserves_transport_profile_and_extra() {
let mut plan = tunnel_timeout_plan(false);
let original_profile = ResolvedTransportProfile {
profile_id: "existing-profile".into(),
backend: TRANSPORT_BACKEND_BROWSER_WREQ.into(),
http_mode: TRANSPORT_HTTP_MODE_HTTP1_ONLY.into(),
pool_scope: "provider".into(),
header_fingerprint: Some(json!({"user_agent": "existing"})),
extra: Some(json!({"existing": {"nested": true}})),
};
plan.transport_profile = Some(original_profile.clone());
let bounded_plan =
with_upstream_response_body_limit(&plan, DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES);
assert_eq!(plan.transport_profile, Some(original_profile.clone()));
assert_eq!(bounded_plan.transport_profile, Some(original_profile));
assert_eq!(
bounded_plan
.headers
.get(EXECUTION_RESPONSE_BODY_LIMIT_HEADER)
.and_then(|value| value.parse::<usize>().ok()),
Some(DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES)
);
assert_eq!(
execution_plan_response_body_limit_bytes(&bounded_plan),
DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES
);
let unprofiled_plan = tunnel_timeout_plan(false);
let bounded_unprofiled_plan = with_upstream_response_body_limit(
&unprofiled_plan,
DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES,
);
assert!(unprofiled_plan.transport_profile.is_none());
assert!(bounded_unprofiled_plan.transport_profile.is_none());
assert_eq!(
execution_plan_response_body_limit_bytes(&bounded_unprofiled_plan),
DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES
);
let mut shadowed_plan = tunnel_timeout_plan(false);
shadowed_plan.headers.insert(
EXECUTION_RESPONSE_BODY_LIMIT_HEADER.to_ascii_uppercase(),
"65536".to_string(),
);
let bounded_shadowed_plan = with_upstream_response_body_limit(
&shadowed_plan,
DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES,
);
assert_eq!(
bounded_shadowed_plan
.headers
.keys()
.filter(|name| name.eq_ignore_ascii_case(EXECUTION_RESPONSE_BODY_LIMIT_HEADER))
.count(),
1
);
}
#[test]
fn scoped_response_body_limit_parsing_rejects_invalid_values_and_clamps_bounds() {
let scoped_plan = |raw_limit: &str| {
let mut plan = tunnel_timeout_plan(false);
plan.headers.insert(
EXECUTION_RESPONSE_BODY_LIMIT_HEADER.to_string(),
raw_limit.to_string(),
);
plan
};
for invalid in ["0", "-1", "1.5", "", "invalid"] {
assert_eq!(
execution_plan_response_body_limit_bytes(&scoped_plan(invalid)),
DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES
);
}
assert_eq!(
execution_plan_response_body_limit_bytes(&scoped_plan("1")),
MIN_SCOPED_RESPONSE_BODY_LIMIT_BYTES
);
assert_eq!(
execution_plan_response_body_limit_bytes(&scoped_plan(
&(MAX_SCOPED_RESPONSE_BODY_LIMIT_BYTES as u64 + 1).to_string()
)),
MAX_SCOPED_RESPONSE_BODY_LIMIT_BYTES
);
assert_eq!(
execution_plan_response_body_limit_bytes(&scoped_plan("1048576")),
1_048_576
);
assert_eq!(
effective_response_body_limit_bytes(
Some(&(DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES * 2).to_string()),
1024 * 1024,
),
1024 * 1024,
"a scoped limit must never raise the operator's global cap"
);
}
#[test]
fn scoped_response_body_wire_limit_rejects_overflow() {
let bounded_plan = with_upstream_response_body_limit(
&tunnel_timeout_plan(false),
MIN_SCOPED_RESPONSE_BODY_LIMIT_BYTES,
);
let limit_bytes = execution_plan_response_body_limit_bytes(&bounded_plan);
let mut body = vec![b'x'; limit_bytes];
let error =
append_upstream_response_body_chunk_with_limit(&mut body, b"overflow", limit_bytes)
.expect_err("wire body above the plan-scoped limit should fail");
assert!(matches!(
error,
ExecutionRuntimeTransportError::UpstreamResponseTooLarge {
phase: UpstreamResponseBodyPhase::Wire,
limit_bytes: MIN_SCOPED_RESPONSE_BODY_LIMIT_BYTES,
}
));
}
#[test]
fn scoped_response_body_limit_rejects_gzip_bomb_after_wire_check() {
let bounded_plan = with_upstream_response_body_limit(
&tunnel_timeout_plan(false),
MIN_SCOPED_RESPONSE_BODY_LIMIT_BYTES,
);
let limit_bytes = execution_plan_response_body_limit_bytes(&bounded_plan);
let payload = vec![b'x'; limit_bytes + 1];
let mut encoder = flate2::write::GzEncoder::new(Vec::new(), flate2::Compression::default());
encoder
.write_all(&payload)
.expect("gzip payload should encode");
let encoded = encoder.finish().expect("gzip payload should finish");
assert!(encoded.len() < limit_bytes);
let mut wire_body = Vec::new();
append_upstream_response_body_chunk_with_limit(&mut wire_body, &encoded, limit_bytes)
.expect("compressed wire body should fit within the plan-scoped limit");
let headers = BTreeMap::from([("content-encoding".to_string(), "gzip".to_string())]);
let error = decode_response_body_bytes_with_limit(&headers, &wire_body, limit_bytes)
.expect_err("decoded body above the plan-scoped limit should fail");
assert!(matches!(
error,
ExecutionRuntimeTransportError::UpstreamResponseTooLarge {
phase: UpstreamResponseBodyPhase::Decoded,
limit_bytes: MIN_SCOPED_RESPONSE_BODY_LIMIT_BYTES,
}
));
}
#[test]
fn upstream_response_wire_limit_allows_exact_body_and_rejects_next_byte() {
let mut body = Vec::new();
@@ -5605,6 +5976,8 @@ mod tests {
)
.await
.expect("headers should write");
socket.flush().await.expect("headers should flush");
tokio::time::sleep(std::time::Duration::from_millis(40)).await;
socket
.write_all(b"b\r\ndata: one\n\n\r\n")
.await
@@ -5634,12 +6007,34 @@ mod tests {
let body = result
.body
.clone()
.and_then(|body| body.body_bytes_b64)
.and_then(|body| base64::engine::general_purpose::STANDARD.decode(body).ok())
.expect("stream body should be captured as bytes");
let body = String::from_utf8(body).expect("stream body should be utf8");
assert!(body.contains("data: one"));
assert!(body.contains("data: two"));
let observation = result
.response_observation
.expect("stream sync execution should preserve header observation");
let telemetry = result
.telemetry
.expect("stream sync execution should include telemetry");
let ttfb_ms = telemetry
.ttfb_ms
.expect("stream sync execution should measure the first body byte");
assert!(
observation.response_headers_observed_at_unix_ms
>= observation.request_started_at_unix_ms
);
assert!(
observation
.response_headers_observed_at_unix_ms
.saturating_sub(observation.request_started_at_unix_ms)
< ttfb_ms,
"header observation must not be derived from body-byte ttfb"
);
assert!(!observation.request_order_id.is_empty());
}
#[tokio::test]
@@ -266,6 +266,7 @@ pub(crate) async fn maybe_execute_windsurf_sync(
candidate_id: prepared.candidate_id,
status_code: 200,
headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]),
response_observation: None,
body: Some(ResponseBody {
json_body: Some(body_json),
body_bytes_b64: None,
@@ -527,6 +528,7 @@ fn build_windsurf_stream_frame_stream(
("cache-control".to_string(), "no-cache".to_string()),
("content-type".to_string(), "text/event-stream".to_string()),
]),
response_observation: None,
},
});
@@ -20,9 +20,11 @@ use crate::ai_serving::LocalExecutionAttemptSource;
use crate::clock::current_unix_ms;
use crate::control::GatewayControlDecision;
use crate::execution_runtime::{
build_transport_error_stop_response, execute_execution_runtime_stream_with_retry_scope,
acquire_upstream_execution_gate, build_transport_error_stop_response,
execute_execution_runtime_stream_with_retry_scope,
execute_execution_runtime_sync_with_retry_scope,
mark_stream_candidate_watchdog_terminal_started, StreamCandidateWatchdogProgress,
UpstreamExecutionGateProvider, UPSTREAM_EXECUTION_GATE_NAME,
};
use crate::executor::{
build_local_execution_exhaustion, mark_deferred_upstream_response, LocalExecutionRequestOutcome,
@@ -43,7 +45,6 @@ use crate::stage_metrics::observe_gateway_stage_ms;
use crate::{AppState, GatewayError};
const DEFAULT_STREAM_FIRST_BYTE_WATCHDOG_TIMEOUT_MS: u64 = 30_000;
const UPSTREAM_EXECUTION_GATE_NAME: &str = "gateway_upstream_execution";
const UPSTREAM_TARGET_GATE_NAME: &str = "gateway_upstream_target";
const UPSTREAM_EXECUTION_GATE_HOLD_STREAM_RESPONSE_ENV: &str =
"AETHER_GATEWAY_UPSTREAM_EXECUTION_GATE_HOLD_STREAM_RESPONSE";
@@ -1612,47 +1613,6 @@ fn hold_response_upstream_execution_permit(
Response::from_parts(parts, Body::from_stream(stream))
}
trait UpstreamExecutionGateProvider {
fn upstream_execution_gate(&self) -> Option<&aether_runtime::ConcurrencyGate>;
fn upstream_execution_gate_queue_budget(&self) -> Duration;
}
impl UpstreamExecutionGateProvider for AppState {
fn upstream_execution_gate(&self) -> Option<&aether_runtime::ConcurrencyGate> {
self.upstream_execution_gate.as_deref()
}
fn upstream_execution_gate_queue_budget(&self) -> Duration {
self.frontdoor_runtime_guards.internal_gate_queue_budget
}
}
async fn acquire_upstream_execution_gate(
state: &(impl UpstreamExecutionGateProvider + ?Sized),
trace_id: &str,
) -> Result<Option<ConcurrencyPermit>, GatewayError> {
let Some(gate) = state.upstream_execution_gate() else {
return Ok(None);
};
let budget = state.upstream_execution_gate_queue_budget();
let gate_wait_started_at = std::time::Instant::now();
match timeout(budget, gate.acquire()).await {
Ok(Ok(permit)) => {
observe_gateway_stage_ms(
"upstream_execution_gate_wait",
gate_wait_started_at.elapsed().as_millis() as u64,
);
Ok(Some(permit))
}
Ok(Err(err)) => Err(GatewayError::Internal(err.to_string())),
Err(_) => Err(GatewayError::AdmissionTimeout {
trace_id: trace_id.to_string(),
gate: UPSTREAM_EXECUTION_GATE_NAME,
queue_budget_ms: budget.as_millis() as u64,
}),
}
}
pub(crate) async fn mark_unused_local_candidate_items<T, FPlan, FContext>(
state: &AppState,
remaining: Vec<T>,
@@ -1687,6 +1687,7 @@ mod tests {
CONTENT_TYPE.as_str().to_string(),
"application/json".to_string(),
)]),
response_observation: None,
body: Some(aether_contracts::ResponseBody {
json_body: Some(body_json),
body_bytes_b64: None,
@@ -55,6 +55,33 @@ pub(super) async fn maybe_handle(
if idempotency_key.is_empty() {
return Ok(Some(bad_request_response("idempotency_key 不能为空")));
}
if idempotency_key.len() > 256 {
return Ok(Some(bad_request_response(
"idempotency_key 不能超过 256 个字节",
)));
}
let expected_credential_generation = match payload.expected_credential_generation {
serde_json::Value::Null => None,
serde_json::Value::String(value) => {
let value = value.trim().to_string();
if value.is_empty() {
return Ok(Some(bad_request_response(
"expected_credential_generation 不能为空字符串",
)));
}
if value.len() > 256 {
return Ok(Some(bad_request_response(
"expected_credential_generation 不能超过 256 个字节",
)));
}
Some(value)
}
_ => {
return Ok(Some(bad_request_response(
"expected_credential_generation 必须是字符串或 null",
)));
}
};
let Some(key) = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
@@ -93,9 +120,15 @@ pub(super) async fn maybe_handle(
)));
};
let (status, payload) =
consume_codex_reset_credit_locally(state, &provider, &endpoint, key, &idempotency_key)
.await?;
let (status, payload) = consume_codex_reset_credit_locally(
state,
&provider,
&endpoint,
key,
&idempotency_key,
expected_credential_generation.as_deref(),
)
.await?;
Ok(Some((status, Json(payload)).into_response()))
}
@@ -1,7 +1,10 @@
use crate::handlers::admin::admin_provider_pool_config;
use crate::handlers::admin::provider::shared::paths::admin_update_key_id;
use crate::handlers::admin::provider::shared::payloads::AdminProviderKeyUpdatePatch;
use crate::handlers::admin::provider::write::keys::admin_provider_key_update_requires_immediate_model_fetch;
use crate::handlers::admin::provider::write::keys::{
admin_provider_key_update_requires_immediate_model_fetch,
build_provider_catalog_key_admin_cas_update,
};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::maintenance::ensure_provider_key_pool_scores_for_keys;
use crate::provider_key_auth::provider_key_effective_api_formats;
@@ -82,7 +85,25 @@ pub(super) async fn maybe_handle(
Ok(record) => record,
Err(detail) => return Ok(Some(bad_request_response(detail))),
};
let Some(mut updated) = state.update_provider_catalog_key(&updated_record).await? else {
let admin_update = build_provider_catalog_key_admin_cas_update(
&existing_key,
updated_record.clone(),
&provider.provider_type,
);
if !state
.compare_and_update_provider_catalog_key_admin_state(&admin_update)
.await?
{
return Ok(Some(conflict_response(
"Key 凭据或配置已被其他请求更新,请刷新后重试",
)));
}
let Some(mut updated) = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
.await?
.into_iter()
.next()
else {
return Ok(None);
};
if updated_record.learned_rpm_limit != existing_key.learned_rpm_limit {
@@ -183,3 +204,11 @@ fn not_found_response(detail: impl Into<String>) -> Response<Body> {
)
.into_response()
}
fn conflict_response(detail: impl Into<String>) -> Response<Body> {
(
http::StatusCode::CONFLICT,
Json(json!({ "detail": detail.into() })),
)
.into_response()
}
@@ -22,16 +22,60 @@ use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::shared::sync_provider_key_oauth_status_snapshot;
use crate::provider_key_auth::provider_key_is_oauth_managed;
use crate::GatewayError;
use aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyOAuthRuntimeStateCasUpdate;
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyOAuthRuntimeStateCasUpdate,
ProviderCatalogUpstreamMetadataNamespaceExpectation,
};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
use serde_json::{json, Value};
use std::time::{SystemTime, UNIX_EPOCH};
const CODEX_OAUTH_COMPLETE_NAMESPACE_CAS_MAX_RETRIES: usize = 3;
const CODEX_CREDENTIAL_GENERATION_KEY: &str = "credential_generation";
#[derive(Debug, PartialEq)]
enum CodexOAuthCompleteCasMissAction {
AlreadyCompleted,
RetryNamespace(Option<Value>),
Conflict,
}
fn codex_oauth_complete_cas_miss_action(
latest_encrypted_auth_config: Option<&str>,
latest_upstream_metadata: Option<&Value>,
latest_status_snapshot: Option<&Value>,
expected_encrypted_auth_config: Option<&str>,
persisted_encrypted_auth_config: &str,
expected_codex_metadata_value: Option<&Value>,
replacement_codex_metadata_value: &Value,
) -> CodexOAuthCompleteCasMissAction {
let latest_codex_metadata_value = latest_upstream_metadata
.and_then(Value::as_object)
.and_then(|metadata| metadata.get("codex"))
.cloned();
let quota_is_cleared = latest_status_snapshot
.and_then(Value::as_object)
.and_then(|snapshot| snapshot.get("quota"))
== Some(&Value::Null);
if latest_encrypted_auth_config == Some(persisted_encrypted_auth_config)
&& latest_codex_metadata_value.as_ref() == Some(replacement_codex_metadata_value)
&& quota_is_cleared
{
return CodexOAuthCompleteCasMissAction::AlreadyCompleted;
}
if latest_encrypted_auth_config != expected_encrypted_auth_config
|| latest_codex_metadata_value.as_ref() == expected_codex_metadata_value
{
return CodexOAuthCompleteCasMissAction::Conflict;
}
CodexOAuthCompleteCasMissAction::RetryNamespace(latest_codex_metadata_value)
}
pub(super) async fn handle_admin_provider_oauth_complete_key(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
@@ -277,29 +321,98 @@ pub(super) async fn handle_admin_provider_oauth_complete_key(
.and_then(|snapshot| snapshot.get("oauth"))
.cloned()
.unwrap_or(serde_json::Value::Null);
let mut expected_codex_metadata_value = key
.upstream_metadata
.as_ref()
.and_then(serde_json::Value::as_object)
.and_then(|metadata| metadata.get("codex"))
.cloned();
let mut status_snapshot_patch =
serde_json::Map::from_iter([("oauth".to_string(), oauth_status)]);
if provider_type == "codex" {
status_snapshot_patch.insert("quota".to_string(), serde_json::Value::Null);
}
let persisted_encrypted_auth_config = recovered_key
.encrypted_auth_config
.clone()
.expect("recovered auth config should be present");
let updated_result = state
.app()
.compare_and_update_provider_catalog_key_oauth_runtime_state(
&ProviderCatalogKeyOAuthRuntimeStateCasUpdate {
key_id: key_id.clone(),
expected_encrypted_auth_config: state_data.expected_encrypted_auth_config,
expected_credential: None,
encrypted_auth_config: persisted_encrypted_auth_config.clone(),
encrypted_api_key_update: Some(encrypted_api_key),
expires_at_unix_secs_update: Some(expires_at),
oauth_invalid_at_unix_secs: None,
oauth_invalid_reason: None,
reset_error_count: true,
upstream_metadata_patch: None,
status_snapshot_patch: json!({ "oauth": oauth_status }),
updated_at_unix_secs: Some(now_unix_secs),
},
)
.await;
let replacement_codex_metadata_value = json!({
CODEX_CREDENTIAL_GENERATION_KEY: uuid::Uuid::now_v7().to_string()
});
let expected_encrypted_auth_config = state_data.expected_encrypted_auth_config.clone();
let updated_result: Result<bool, GatewayError> = async {
let max_namespace_retries = if provider_type == "codex" {
CODEX_OAUTH_COMPLETE_NAMESPACE_CAS_MAX_RETRIES
} else {
0
};
for retry in 0..=max_namespace_retries {
let updated = state
.app()
.compare_and_update_provider_catalog_key_oauth_runtime_state(
&ProviderCatalogKeyOAuthRuntimeStateCasUpdate {
key_id: key_id.clone(),
expected_encrypted_auth_config: expected_encrypted_auth_config.clone(),
expected_credential: None,
expected_upstream_metadata_namespace: (provider_type == "codex").then(
|| ProviderCatalogUpstreamMetadataNamespaceExpectation {
namespace: "codex".to_string(),
expected_value: expected_codex_metadata_value.clone(),
},
),
encrypted_auth_config: persisted_encrypted_auth_config.clone(),
encrypted_api_key_update: Some(encrypted_api_key.clone()),
expires_at_unix_secs_update: Some(expires_at),
oauth_invalid_at_unix_secs: None,
oauth_invalid_reason: None,
reset_error_count: true,
upstream_metadata_patch: (provider_type == "codex")
.then(|| json!({"codex": replacement_codex_metadata_value.clone()})),
upstream_metadata_namespace_to_remove: None,
status_snapshot_patch: serde_json::Value::Object(
status_snapshot_patch.clone(),
),
updated_at_unix_secs: Some(now_unix_secs),
},
)
.await?;
if updated {
return Ok(true);
}
if provider_type != "codex" {
return Ok(false);
}
let Some(latest_key) = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
.await?
.into_iter()
.next()
else {
return Ok(false);
};
match codex_oauth_complete_cas_miss_action(
latest_key.encrypted_auth_config.as_deref(),
latest_key.upstream_metadata.as_ref(),
latest_key.status_snapshot.as_ref(),
expected_encrypted_auth_config.as_deref(),
&persisted_encrypted_auth_config,
expected_codex_metadata_value.as_ref(),
&replacement_codex_metadata_value,
) {
CodexOAuthCompleteCasMissAction::AlreadyCompleted => return Ok(true),
CodexOAuthCompleteCasMissAction::Conflict => return Ok(false),
CodexOAuthCompleteCasMissAction::RetryNamespace(latest_codex_metadata_value) => {
if retry == max_namespace_retries {
return Ok(false);
}
expected_codex_metadata_value = latest_codex_metadata_value;
}
}
}
Ok(false)
}
.await;
let _ = state
.app()
.invalidate_local_oauth_refresh_entry(&key_id)
@@ -397,3 +510,78 @@ pub(super) async fn handle_admin_provider_oauth_complete_key(
}))
.into_response())
}
#[cfg(test)]
mod tests {
use super::{codex_oauth_complete_cas_miss_action, CodexOAuthCompleteCasMissAction};
use serde_json::json;
#[test]
fn codex_oauth_complete_retries_only_when_namespace_changed() {
let expected_codex = json!({"request_id": "old"});
let replacement_codex = json!({"credential_generation": "generation-new"});
let latest_metadata = json!({
"codex": {"request_id": "new"},
"unrelated": {"preserved": true}
});
assert_eq!(
codex_oauth_complete_cas_miss_action(
Some("old-auth"),
Some(&latest_metadata),
None,
Some("old-auth"),
"new-auth",
Some(&expected_codex),
&replacement_codex,
),
CodexOAuthCompleteCasMissAction::RetryNamespace(Some(json!({
"request_id": "new"
})))
);
assert_eq!(
codex_oauth_complete_cas_miss_action(
Some("old-auth"),
Some(&json!({"codex": expected_codex.clone()})),
None,
Some("old-auth"),
"new-auth",
Some(&expected_codex),
&replacement_codex,
),
CodexOAuthCompleteCasMissAction::Conflict
);
}
#[test]
fn codex_oauth_complete_accepts_an_ambiguous_success_but_rejects_auth_rotation() {
let replacement_codex = json!({"credential_generation": "generation-new"});
assert_eq!(
codex_oauth_complete_cas_miss_action(
Some("new-auth"),
Some(&json!({
"codex": replacement_codex.clone(),
"unrelated": {"preserved": true}
})),
Some(&json!({"quota": null})),
Some("old-auth"),
"new-auth",
Some(&json!({"request_id": "old"})),
&replacement_codex,
),
CodexOAuthCompleteCasMissAction::AlreadyCompleted
);
assert_eq!(
codex_oauth_complete_cas_miss_action(
Some("other-auth"),
Some(&json!({"codex": {"request_id": "new"}})),
Some(&json!({"quota": null})),
Some("old-auth"),
"new-auth",
Some(&json!({"request_id": "old"})),
&replacement_codex,
),
CodexOAuthCompleteCasMissAction::Conflict
);
}
}
@@ -12,6 +12,7 @@ use crate::ai_serving::{
build_provider_key_pool_score_upsert, provider_key_pool_score_id, provider_key_pool_score_scope,
};
use crate::handlers::admin::admin_provider_pool_config;
use crate::handlers::admin::provider::write::keys::build_provider_catalog_key_admin_cas_update;
use crate::handlers::admin::request::AdminAppState;
use crate::provider_key_auth::provider_active_api_formats;
use crate::GatewayError;
@@ -286,6 +287,61 @@ fn grok_oauth_catalog_key_fingerprint(
grok_browser_transport_fingerprint_from_auth_config(auth_config)
}
pub(crate) fn rotate_codex_credential_generation(
key: &mut StoredProviderCatalogKey,
provider_type: &str,
) {
if !provider_type.trim().eq_ignore_ascii_case("codex") {
return;
}
let mut upstream_metadata = key
.upstream_metadata
.as_ref()
.and_then(Value::as_object)
.cloned()
.unwrap_or_default();
upstream_metadata.insert(
"codex".to_string(),
json!({
aether_admin::provider::quota::CODEX_CREDENTIAL_GENERATION_KEY:
Uuid::now_v7().to_string(),
}),
);
key.upstream_metadata = Some(Value::Object(upstream_metadata));
if let Some(mut status_snapshot) = key
.status_snapshot
.as_ref()
.and_then(Value::as_object)
.cloned()
{
status_snapshot.insert("quota".to_string(), Value::Null);
key.status_snapshot = Some(Value::Object(status_snapshot));
}
}
pub(crate) fn ensure_codex_credential_generation_rotated(
key: &mut StoredProviderCatalogKey,
provider_type: &str,
previous_generation: Option<&str>,
) {
if !provider_type.trim().eq_ignore_ascii_case("codex") {
return;
}
let current_generation = key
.upstream_metadata
.as_ref()
.and_then(Value::as_object)
.and_then(|metadata| metadata.get("codex"))
.and_then(|codex| aether_admin::provider::quota::codex_credential_generation(Some(codex)));
let already_rotated = current_generation.is_some() && current_generation != previous_generation;
if !already_rotated {
rotate_codex_credential_generation(key, provider_type);
}
}
pub(crate) async fn create_provider_oauth_catalog_key(
state: &AdminAppState<'_>,
provider_id: &str,
@@ -344,6 +400,7 @@ pub(crate) async fn create_provider_oauth_catalog_key(
record.circuit_breaker_by_format = Some(json!({}));
record.created_at_unix_ms = Some(now_unix_secs);
record.updated_at_unix_secs = Some(now_unix_secs);
rotate_codex_credential_generation(&mut record, provider_type);
let created = state.create_provider_catalog_key(&record).await?;
if let Some(key) = created.as_ref() {
let _ = state
@@ -395,17 +452,23 @@ pub(crate) async fn update_existing_provider_oauth_catalog_key(
updated.proxy = Some(proxy);
}
updated.updated_at_unix_secs = Some(now_unix_secs);
if state.update_provider_catalog_key(&updated).await?.is_none() {
return Ok(None);
}
rotate_codex_credential_generation(&mut updated, provider_type);
let admin_update =
build_provider_catalog_key_admin_cas_update(existing_key, updated.clone(), provider_type);
if !state
.clear_provider_catalog_key_oauth_invalid_marker(&updated.id)
.compare_and_update_provider_catalog_key_admin_state(&admin_update)
.await?
{
return Ok(None);
}
let persisted = state
.reset_provider_catalog_key_recovery_state(&updated.id)
.reset_provider_catalog_key_recovery_state_fenced(
&updated.id,
updated
.encrypted_auth_config
.as_deref()
.expect("OAuth update always supplies encrypted auth_config"),
)
.await?;
if let Some(key) = persisted.as_ref() {
let _ = state
@@ -502,10 +565,12 @@ fn provider_oauth_catalog_key_api_formats(
#[cfg(test)]
mod tests {
use super::{
grok_oauth_catalog_key_fingerprint, provider_oauth_token_payload_expires_at_unix_secs,
ensure_codex_credential_generation_rotated, grok_oauth_catalog_key_fingerprint,
provider_oauth_token_payload_expires_at_unix_secs, rotate_codex_credential_generation,
};
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
use serde_json::json;
use serde_json::{json, Value};
fn sample_unsigned_jwt(payload: serde_json::Value) -> String {
let header = URL_SAFE_NO_PAD.encode(r#"{"alg":"none","typ":"JWT"}"#);
@@ -609,4 +674,101 @@ mod tests {
assert!(grok_oauth_catalog_key_fingerprint("openai", auth_config).is_none());
}
#[test]
fn codex_credential_rotation_replaces_quota_namespace_and_preserves_unrelated_state() {
let mut key = StoredProviderCatalogKey::new(
"key".to_string(),
"provider".to_string(),
"Codex".to_string(),
"oauth".to_string(),
None,
true,
)
.expect("key should build");
key.upstream_metadata = Some(json!({
"codex": {
"credential_generation": "old-generation",
"primary_used_percent": 75.0,
},
"unrelated": {"preserved": true},
}));
key.status_snapshot = Some(json!({
"oauth": {"status": "valid"},
"quota": {"used_ratio": 0.75},
}));
rotate_codex_credential_generation(&mut key, "codex");
let codex = key
.upstream_metadata
.as_ref()
.and_then(|metadata| metadata.get("codex"))
.and_then(Value::as_object)
.expect("codex namespace should exist");
assert_eq!(codex.len(), 1);
assert_ne!(
codex
.get(aether_admin::provider::quota::CODEX_CREDENTIAL_GENERATION_KEY)
.and_then(Value::as_str),
Some("old-generation")
);
assert_eq!(
key.upstream_metadata
.as_ref()
.and_then(|metadata| metadata.pointer("/unrelated/preserved")),
Some(&json!(true))
);
assert_eq!(
key.status_snapshot
.as_ref()
.and_then(|snapshot| snapshot.get("quota")),
Some(&Value::Null)
);
assert_eq!(
key.status_snapshot
.as_ref()
.and_then(|snapshot| snapshot.pointer("/oauth/status")),
Some(&json!("valid"))
);
}
#[test]
fn codex_credential_rotation_ensure_does_not_rotate_twice_in_one_write() {
let mut key = StoredProviderCatalogKey::new(
"key".to_string(),
"provider".to_string(),
"Codex".to_string(),
"oauth".to_string(),
None,
true,
)
.expect("key should build");
key.upstream_metadata = Some(json!({
"codex": {"credential_generation": "generation-before-write"}
}));
rotate_codex_credential_generation(&mut key, "codex");
let builder_generation = key
.upstream_metadata
.as_ref()
.and_then(|metadata| metadata.pointer("/codex/credential_generation"))
.and_then(Value::as_str)
.expect("builder should rotate the generation")
.to_string();
ensure_codex_credential_generation_rotated(
&mut key,
"codex",
Some("generation-before-write"),
);
assert_eq!(
key.upstream_metadata
.as_ref()
.and_then(|metadata| metadata.pointer("/codex/credential_generation"))
.and_then(Value::as_str),
Some(builder_generation.as_str())
);
}
}
@@ -470,6 +470,7 @@ mod tests {
candidate_id: None,
status_code: 403,
headers: BTreeMap::new(),
response_observation: None,
body: Some(ResponseBody {
json_body: None,
body_bytes_b64: Some(base64::engine::general_purpose::STANDARD.encode(body)),
@@ -18,24 +18,229 @@ use self::plan::{
execute_codex_reset_credit_plan,
};
use super::shared::{
build_quota_snapshot_payload, extract_execution_error_message,
oauth_refresh_auto_removed_result, persist_fenced_provider_quota_refresh_state,
persist_provider_quota_refresh_state, provider_auto_remove_banned_keys,
build_quota_snapshot_payload, complete_codex_account_reset, extract_execution_error_message,
oauth_refresh_auto_removed_result, persist_codex_provider_quota_refresh_state,
persist_fenced_provider_quota_refresh_state, provider_auto_remove_banned_keys,
provider_auto_remove_quota_exhausted_keys, quota_key_auto_removed,
quota_refresh_success_invalid_state, should_auto_remove_oauth_invalid_key,
ProviderQuotaExecutionOutcome,
quota_refresh_success_invalid_state, reserve_codex_account_reset,
should_auto_remove_oauth_invalid_key, CodexAccountResetCompleteResult,
CodexAccountResetReserveResult, CodexAccountResetTerminal, ProviderQuotaExecutionOutcome,
};
use crate::handlers::admin::request::AdminAppState;
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
use crate::provider_key_auth::provider_key_is_oauth_managed;
use crate::state::ProviderTransportCredentialFence;
use crate::GatewayError;
use aether_contracts::ProxySnapshot;
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthCredentialFence,
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
ProviderCatalogKeyOAuthCredentialCasDelete,
ProviderCatalogUpstreamMetadataNamespaceExpectation, StoredProviderCatalogEndpoint,
StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use axum::http::StatusCode;
use serde_json::{json, Map, Value};
use std::time::{SystemTime, UNIX_EPOCH};
const CODEX_OAUTH_CREDENTIAL_STABILIZATION_ATTEMPTS: usize = 3;
const CODEX_RESET_QUOTA_RECONCILIATION_DELAYS_MS: [u64; 4] = [1_000, 2_000, 4_000, 8_000];
enum CodexOAuthRequestPreparation {
Ready {
transport: AdminGatewayProviderTransportSnapshot,
auth: (String, String),
credential_fence: ProviderTransportCredentialFence,
},
MissingAuth,
Conflict,
}
async fn prepare_codex_oauth_request(
state: &AdminAppState<'_>,
initial_transport: &AdminGatewayProviderTransportSnapshot,
) -> Result<CodexOAuthRequestPreparation, GatewayError> {
for _ in 0..CODEX_OAUTH_CREDENTIAL_STABILIZATION_ATTEMPTS {
let Some(transport) = state
.read_provider_transport_snapshot_uncached(
&initial_transport.provider.id,
&initial_transport.endpoint.id,
&initial_transport.key.id,
)
.await?
else {
return Ok(CodexOAuthRequestPreparation::Conflict);
};
if !crate::state::provider_transport_context_allows_credential_rotation(
initial_transport,
&transport,
) {
return Ok(CodexOAuthRequestPreparation::Conflict);
}
let Some(before_fence) = state
.app()
.capture_provider_transport_credential_fence(&transport)
.await?
else {
continue;
};
let resolved_auth = state.resolve_local_oauth_header_auth(&transport).await?;
let Some(current_transport) = state
.read_provider_transport_snapshot_uncached(
&initial_transport.provider.id,
&initial_transport.endpoint.id,
&initial_transport.key.id,
)
.await?
else {
return Ok(CodexOAuthRequestPreparation::Conflict);
};
if !crate::state::provider_transport_context_allows_credential_rotation(
initial_transport,
&current_transport,
) {
return Ok(CodexOAuthRequestPreparation::Conflict);
}
let Some(after_fence) = state
.app()
.capture_provider_transport_credential_fence(&current_transport)
.await?
else {
continue;
};
if before_fence != after_fence {
continue;
}
return Ok(match resolved_auth {
Some(auth) => CodexOAuthRequestPreparation::Ready {
transport: current_transport,
auth,
credential_fence: after_fence,
},
None => CodexOAuthRequestPreparation::MissingAuth,
});
}
Ok(CodexOAuthRequestPreparation::Conflict)
}
fn codex_reset_refresh_succeeded(payload: Option<&Value>, key_id: &str) -> bool {
payload
.and_then(|payload| payload.get("results"))
.and_then(Value::as_array)
.into_iter()
.flatten()
.filter_map(Value::as_object)
.find(|item| item.get("key_id").and_then(Value::as_str) == Some(key_id))
.and_then(|item| item.get("status"))
.and_then(Value::as_str)
.is_some_and(|status| status.eq_ignore_ascii_case("success"))
}
async fn codex_reset_fence_is_still_pending(
state: &AdminAppState<'_>,
key_id: &str,
expected_credential: &ProviderTransportCredentialFence,
reset_fence: &super::shared::CodexAccountResetFence,
) -> Result<bool, GatewayError> {
let Some(key) = state
.read_provider_catalog_keys_by_ids(&[key_id.to_string()])
.await?
.into_iter()
.next()
else {
return Ok(false);
};
if key.encrypted_auth_config.as_deref()
!= Some(expected_credential.encrypted_auth_config.as_str())
|| key.encrypted_api_key != expected_credential.credential.encrypted_api_key
|| key.auth_type != expected_credential.credential.auth_type
|| key.provider_id != expected_credential.credential.provider_id
{
return Ok(false);
}
let provider_type_matches = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&key.provider_id))
.await?
.into_iter()
.next()
.is_some_and(|provider| {
provider.provider_type == expected_credential.credential.provider_type
});
if !provider_type_matches {
return Ok(false);
}
let codex = key
.upstream_metadata
.as_ref()
.and_then(Value::as_object)
.and_then(|metadata| metadata.get("codex"))
.and_then(Value::as_object);
Ok(codex.is_some_and(|codex| {
codex
.get(aether_admin::provider::quota::CODEX_QUOTA_ACCOUNT_RESET_FENCE_ID_KEY)
.and_then(Value::as_str)
== Some(reset_fence.id.as_str())
&& codex
.get(aether_admin::provider::quota::CODEX_QUOTA_ACCOUNT_RESET_GENERATION_KEY)
.and_then(aether_admin::provider::quota::coerce_json_u64)
== Some(reset_fence.generation)
&& codex
.get(
aether_admin::provider::quota::CODEX_QUOTA_ACCOUNT_RESET_PENDING_GENERATION_KEY,
)
.and_then(aether_admin::provider::quota::coerce_json_u64)
== Some(reset_fence.generation)
&& codex
.get(aether_admin::provider::quota::CODEX_QUOTA_ACCOUNT_RESET_PENDING_KEY)
.and_then(Value::as_bool)
== Some(true)
}))
}
async fn refresh_codex_quota_after_reset_until_settled(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
endpoint: &StoredProviderCatalogEndpoint,
key: &StoredProviderCatalogKey,
reset_fence: &super::shared::CodexAccountResetFence,
expected_credential: &ProviderTransportCredentialFence,
) -> Result<Option<Value>, GatewayError> {
let mut latest_payload = None;
for attempt in 0..=CODEX_RESET_QUOTA_RECONCILIATION_DELAYS_MS.len() {
if attempt > 0 {
tokio::time::sleep(std::time::Duration::from_millis(
CODEX_RESET_QUOTA_RECONCILIATION_DELAYS_MS[attempt - 1],
))
.await;
}
if !codex_reset_fence_is_still_pending(state, &key.id, expected_credential, reset_fence)
.await?
{
break;
}
let payload = refresh_codex_provider_quota_locally_with_reset_fence(
state,
provider,
endpoint,
vec![key.clone()],
None,
Some(reset_fence.id.as_str()),
Some(reset_fence.generation),
Some(expected_credential),
)
.await?;
let refresh_succeeded = codex_reset_refresh_succeeded(payload.as_ref(), &key.id);
latest_payload = payload;
if !refresh_succeeded
|| !codex_reset_fence_is_still_pending(state, &key.id, expected_credential, reset_fence)
.await?
{
break;
}
}
Ok(latest_payload)
}
fn merge_codex_quota_metadata(
header_metadata: Option<&serde_json::Value>,
@@ -53,6 +258,26 @@ fn merge_codex_quota_metadata(
serde_json::Value::Object(merged)
}
fn codex_quota_window_coverage(
body_json: Option<&Value>,
) -> aether_admin::provider::quota::CodexQuotaWindowCoverage {
let body = body_json.and_then(Value::as_object);
let has_account_snapshot = body
.and_then(|body| body.get("rate_limit"))
.and_then(Value::as_object)
.is_some();
let has_spark_snapshot = body
.and_then(|body| body.get("additional_rate_limits"))
.and_then(Value::as_array)
.is_some();
match (has_account_snapshot, has_spark_snapshot) {
(true, true) => aether_admin::provider::quota::CodexQuotaWindowCoverage::FullSnapshot,
(true, false) => aether_admin::provider::quota::CodexQuotaWindowCoverage::AccountSnapshot,
_ => aether_admin::provider::quota::CodexQuotaWindowCoverage::Patch,
}
}
fn truncate_codex_reset_credit_detail_error(message: impl Into<String>) -> String {
let message = message.into();
let mut sanitized = message.replace('\n', " ");
@@ -198,6 +423,10 @@ fn codex_consume_success_status(outcome: &str) -> &'static str {
}
}
fn codex_reset_credit_outcome_allows_usage_drop(outcome: &str) -> bool {
matches!(outcome, "reset" | "already_redeemed")
}
fn codex_extract_refresh_result_fields(
refresh_payload: Option<&Value>,
key_id: &str,
@@ -247,12 +476,74 @@ fn codex_extract_refresh_result_fields(
)
}
async fn finish_codex_reset_replay(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
endpoint: &StoredProviderCatalogEndpoint,
key: &StoredProviderCatalogKey,
credential: &ProviderTransportCredentialFence,
terminal: CodexAccountResetTerminal,
) -> Result<(StatusCode, Value), GatewayError> {
let mut refresh_status = "skipped".to_string();
let mut refresh_error = None;
let mut metadata = None;
let mut quota_snapshot = None;
if codex_reset_credit_outcome_allows_usage_drop(&terminal.outcome) {
let fence = super::shared::CodexAccountResetFence {
unix_ms: crate::clock::current_unix_ms(),
id: format!("reset:{}", terminal.idempotency_key),
generation: terminal.generation,
};
if codex_reset_fence_is_still_pending(state, &key.id, credential, &fence).await? {
match refresh_codex_quota_after_reset_until_settled(
state, provider, endpoint, key, &fence, credential,
)
.await
{
Ok(payload) => {
(refresh_status, refresh_error, metadata, quota_snapshot) =
codex_extract_refresh_result_fields(payload.as_ref(), &key.id);
}
Err(err) => {
refresh_status = "failed".to_string();
refresh_error =
Some(truncate_codex_reset_credit_detail_error(err.into_message()));
}
}
}
}
let mut payload = Map::new();
payload.insert("key_id".to_string(), json!(key.id));
payload.insert(
"status".to_string(),
json!(codex_consume_success_status(&terminal.outcome)),
);
payload.insert("outcome".to_string(), json!(terminal.outcome));
payload.insert(
"idempotency_key".to_string(),
json!(terminal.idempotency_key),
);
payload.insert("replay".to_string(), json!(true));
payload.insert("refresh_status".to_string(), json!(refresh_status));
if let Some(refresh_error) = refresh_error {
payload.insert("refresh_error".to_string(), json!(refresh_error));
}
if let Some(metadata) = metadata {
payload.insert("metadata".to_string(), metadata);
}
if let Some(quota_snapshot) = quota_snapshot {
payload.insert("quota_snapshot".to_string(), quota_snapshot);
}
Ok((StatusCode::OK, Value::Object(payload)))
}
pub(crate) async fn consume_codex_reset_credit_locally(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
endpoint: &StoredProviderCatalogEndpoint,
key: StoredProviderCatalogKey,
idempotency_key: &str,
expected_credential_generation: Option<&str>,
) -> Result<(StatusCode, Value), GatewayError> {
let transport = match state
.read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id)
@@ -273,22 +564,48 @@ pub(crate) async fn consume_codex_reset_credit_locally(
};
let is_oauth_managed = provider_key_is_oauth_managed(&key, provider.provider_type.as_str());
let resolved_oauth_auth = if is_oauth_managed {
state.resolve_local_oauth_header_auth(&transport).await?
} else {
None
};
if is_oauth_managed && resolved_oauth_auth.is_none() {
if !is_oauth_managed {
return Ok((
StatusCode::BAD_REQUEST,
json!({
"key_id": key.id,
"status": "error",
"outcome": "error",
"message": "缺少 Codex OAuth 认证信息,请先重新授权/刷新 Token",
"message": "Codex reset credit 仅支持 OAuth 托管账号",
}),
));
}
let (transport, resolved_oauth_auth, reset_credential_fence) =
match prepare_codex_oauth_request(state, &transport).await? {
CodexOAuthRequestPreparation::Ready {
transport,
auth,
credential_fence,
} => (transport, Some(auth), credential_fence),
CodexOAuthRequestPreparation::MissingAuth => {
return Ok((
StatusCode::BAD_REQUEST,
json!({
"key_id": key.id,
"status": "error",
"outcome": "error",
"message": "缺少 Codex OAuth 认证信息,请先重新授权/刷新 Token",
}),
));
}
CodexOAuthRequestPreparation::Conflict => {
return Ok((
StatusCode::CONFLICT,
json!({
"key_id": key.id,
"status": "error",
"outcome": "error",
"idempotency_key": idempotency_key,
"message": "Codex credential changed before reset credit could be consumed",
}),
));
}
};
let request_spec = match build_codex_reset_credit_consume_request_spec(
&transport,
@@ -309,6 +626,79 @@ pub(crate) async fn consume_codex_reset_credit_locally(
}
};
let reservation = match reserve_codex_account_reset(
state,
&key.id,
reset_credential_fence.encrypted_auth_config.as_str(),
&reset_credential_fence.credential,
expected_credential_generation,
idempotency_key,
)
.await?
{
Some(CodexAccountResetReserveResult::Reserved(reservation)) => reservation,
Some(CodexAccountResetReserveResult::Replay(terminal)) => {
return finish_codex_reset_replay(
state,
provider,
endpoint,
&key,
&reset_credential_fence,
terminal,
)
.await;
}
Some(CodexAccountResetReserveResult::LegacyReplay) => {
return Ok((
StatusCode::OK,
json!({
"key_id": key.id,
"status": "success",
"outcome": "historical_replay",
"idempotency_key": idempotency_key,
"refresh_status": "skipped",
}),
));
}
Some(CodexAccountResetReserveResult::Busy(active)) => {
return Ok((
StatusCode::CONFLICT,
json!({
"key_id": key.id,
"status": "error",
"outcome": "busy",
"idempotency_key": idempotency_key,
"active_idempotency_key": active.idempotency_key,
"message": "Another Codex reset credit operation is unresolved",
}),
));
}
Some(CodexAccountResetReserveResult::CredentialGenerationMismatch) => {
return Ok((
StatusCode::CONFLICT,
json!({
"key_id": key.id,
"status": "error",
"outcome": "credential_changed",
"idempotency_key": idempotency_key,
"message": "Codex credential changed since this reset request was prepared",
}),
));
}
None => {
return Ok((
StatusCode::CONFLICT,
json!({
"key_id": key.id,
"status": "error",
"outcome": "error",
"idempotency_key": idempotency_key,
"message": "Codex reset reservation could not be persisted",
}),
));
}
};
let result =
match execute_codex_reset_credit_plan(state, &transport, request_spec, None).await? {
ProviderQuotaExecutionOutcome::Response(result) => result,
@@ -331,11 +721,11 @@ pub(crate) async fn consume_codex_reset_credit_locally(
.and_then(|body| body.json_body.as_ref());
let outcome = normalize_codex_reset_credit_consume_outcome(body_json)
.unwrap_or_else(|| "unknown".to_string());
let known_non_error_outcome = matches!(
let known_terminal_outcome = matches!(
outcome.as_str(),
"reset" | "already_redeemed" | "nothing_to_reset" | "no_credit"
);
if result.status_code >= 400 && !known_non_error_outcome {
if !known_terminal_outcome {
let detail = extract_execution_error_message(&result)
.unwrap_or_else(|| format!("HTTP {}", result.status_code));
return Ok((
@@ -345,19 +735,88 @@ pub(crate) async fn consume_codex_reset_credit_locally(
"status": "error",
"outcome": "error",
"idempotency_key": idempotency_key,
"message": format!("reset credit consume 返回状态码 {}: {detail}", result.status_code),
"message": format!("reset credit consume outcome is ambiguous: {detail}"),
"status_code": result.status_code,
}),
));
}
let (refresh_status, refresh_error, metadata, quota_snapshot) =
match refresh_codex_provider_quota_locally(
let fence_unix_ms = result
.response_observation
.as_ref()
.map(|observation| observation.response_headers_observed_at_unix_ms)
.unwrap_or_else(crate::clock::current_unix_ms);
let Some(completed) = complete_codex_account_reset(
state,
&key.id,
reset_credential_fence.encrypted_auth_config.as_str(),
&reset_credential_fence.credential,
&reservation,
&outcome,
fence_unix_ms,
)
.await?
else {
return Ok((
StatusCode::CONFLICT,
json!({
"key_id": key.id,
"status": "error",
"outcome": "error",
"idempotency_key": idempotency_key,
"message": "Codex reset completion could not be persisted",
}),
));
};
let (effective_outcome, reset_fence) = match completed {
CodexAccountResetCompleteResult::Activated(fence) => (outcome.clone(), Some(fence)),
CodexAccountResetCompleteResult::Noop(terminal)
| CodexAccountResetCompleteResult::Replay(terminal) => {
let fence =
codex_reset_credit_outcome_allows_usage_drop(&terminal.outcome).then(|| {
super::shared::CodexAccountResetFence {
unix_ms: fence_unix_ms,
id: format!("reset:{}", terminal.idempotency_key),
generation: terminal.generation,
}
});
(terminal.outcome, fence)
}
};
let (refresh_status, refresh_error, metadata, quota_snapshot) = match reset_fence.as_ref() {
Some(reset_fence) => {
match refresh_codex_quota_after_reset_until_settled(
state,
provider,
endpoint,
&key,
reset_fence,
&reset_credential_fence,
)
.await
{
Ok(refresh_payload) => {
codex_extract_refresh_result_fields(refresh_payload.as_ref(), &key.id)
}
Err(err) => (
"failed".to_string(),
Some(truncate_codex_reset_credit_detail_error(err.into_message())),
None,
None,
),
}
}
None => match refresh_codex_provider_quota_locally_with_reset_fence(
state,
provider,
endpoint,
vec![key.clone()],
None,
None,
None,
Some(&reset_credential_fence),
)
.await
{
@@ -370,15 +829,16 @@ pub(crate) async fn consume_codex_reset_credit_locally(
None,
None,
),
};
},
};
let mut payload = Map::new();
payload.insert("key_id".to_string(), json!(key.id));
payload.insert(
"status".to_string(),
json!(codex_consume_success_status(&outcome)),
json!(codex_consume_success_status(&effective_outcome)),
);
payload.insert("outcome".to_string(), json!(outcome));
payload.insert("outcome".to_string(), json!(effective_outcome));
payload.insert("idempotency_key".to_string(), json!(idempotency_key));
payload.insert("refresh_status".to_string(), json!(refresh_status));
if let Some(refresh_error) = refresh_error {
@@ -400,6 +860,29 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
endpoint: &StoredProviderCatalogEndpoint,
keys: Vec<StoredProviderCatalogKey>,
proxy_override: Option<ProxySnapshot>,
) -> Result<Option<serde_json::Value>, GatewayError> {
refresh_codex_provider_quota_locally_with_reset_fence(
state,
provider,
endpoint,
keys,
proxy_override,
None,
None,
None,
)
.await
}
async fn refresh_codex_provider_quota_locally_with_reset_fence(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
endpoint: &StoredProviderCatalogEndpoint,
keys: Vec<StoredProviderCatalogKey>,
proxy_override: Option<ProxySnapshot>,
account_reset_fence_id: Option<&str>,
authoritative_reset_generation: Option<u64>,
expected_reset_credential: Option<&crate::state::ProviderTransportCredentialFence>,
) -> Result<Option<serde_json::Value>, GatewayError> {
let mut results = Vec::new();
let mut success_count = 0usize;
@@ -412,7 +895,7 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
for key in keys {
let had_oauth_refresh_issue =
codex_oauth_refresh_issue_reason(key.oauth_invalid_reason.as_deref());
let transport = match state
let initial_transport = match state
.read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id)
.await?
{
@@ -429,48 +912,70 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
}
};
let is_oauth_managed = provider_key_is_oauth_managed(&key, provider.provider_type.as_str());
let quota_auth_config_fence = if is_oauth_managed {
match state
.app()
.capture_provider_transport_auth_config_fence(&transport)
.await?
{
Some(ciphertext) => Some(ciphertext),
None => {
let (transport, resolved_oauth_auth, quota_credential_fence) = if is_oauth_managed {
match prepare_codex_oauth_request(state, &initial_transport).await? {
CodexOAuthRequestPreparation::Ready {
transport,
auth,
credential_fence,
} => (transport, Some(auth), Some(credential_fence)),
CodexOAuthRequestPreparation::MissingAuth => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "OAuth credential changed before quota refresh",
"message": "缺少 Codex OAuth 认证信息,请先重新授权/刷新 Token",
}));
continue;
}
CodexOAuthRequestPreparation::Conflict => {
if quota_key_auto_removed(state, &key.id).await? {
auto_removed_count += 1;
results.push(oauth_refresh_auto_removed_result(&key));
} else {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "OAuth credential changed before quota refresh",
}));
}
continue;
}
}
} else {
None
(initial_transport, None, None)
};
let resolved_oauth_auth = if is_oauth_managed {
state.resolve_local_oauth_header_auth(&transport).await?
} else {
None
};
if is_oauth_managed && quota_key_auto_removed(state, &key.id).await? {
auto_removed_count += 1;
results.push(oauth_refresh_auto_removed_result(&key));
continue;
}
if is_oauth_managed && resolved_oauth_auth.is_none() {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "缺少 Codex OAuth 认证信息,请先重新授权/刷新 Token",
}));
continue;
if let Some(expected_reset_credential) = expected_reset_credential {
if quota_credential_fence.as_ref() != Some(expected_reset_credential) {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "Codex credential changed after reset credit was consumed",
}));
continue;
}
}
let transport_codex_metadata = transport
.key
.upstream_metadata
.as_ref()
.and_then(Value::as_object)
.and_then(|metadata| metadata.get("codex"));
let observed_reset_generation = authoritative_reset_generation.or_else(|| {
Some(
aether_admin::provider::quota::codex_quota_account_reset_generation(
transport_codex_metadata,
),
)
});
let observed_credential_generation =
aether_admin::provider::quota::codex_credential_generation(transport_codex_metadata)
.map(ToOwned::to_owned);
let request_spec =
match build_codex_quota_request_spec(&transport, resolved_oauth_auth.clone()) {
@@ -487,6 +992,8 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
}
};
let quota_request_fallback_started_at_unix_ms = crate::clock::current_unix_ms();
let quota_request_fallback_order_id = uuid::Uuid::now_v7().to_string();
let result = match execute_codex_quota_plan(
state,
&transport,
@@ -508,16 +1015,25 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
continue;
}
};
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let quota_response_fallback_observed_at_unix_ms = crate::clock::current_unix_ms();
let quota_response_observation = result.response_observation.as_ref();
let quota_request_started_at_unix_ms = quota_response_observation
.map(|observation| observation.request_started_at_unix_ms)
.unwrap_or(quota_request_fallback_started_at_unix_ms);
let quota_response_observed_at_unix_ms = quota_response_observation
.map(|observation| observation.response_headers_observed_at_unix_ms)
.unwrap_or(quota_response_fallback_observed_at_unix_ms);
let quota_request_order_id = quota_response_observation
.map(|observation| observation.request_order_id.as_str())
.unwrap_or(quota_request_fallback_order_id.as_str());
let now_unix_secs = quota_response_observed_at_unix_ms / 1_000;
let header_metadata = parse_codex_usage_headers(&result.headers, now_unix_secs);
let mut metadata_update = header_metadata
.as_ref()
.map(|metadata| json!({ "codex": metadata }));
let mut quota_window_coverage =
aether_admin::provider::quota::CodexQuotaWindowCoverage::Patch;
let (mut oauth_invalid_at_unix_secs, mut oauth_invalid_reason) = (None, None);
let mut status = "error".to_string();
let mut message = None::<String>;
@@ -544,6 +1060,7 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
now_unix_secs,
)
.await?;
quota_window_coverage = codex_quota_window_coverage(Some(body_json));
metadata_update = Some(json!({
"codex": codex_metadata
}));
@@ -586,6 +1103,8 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
}
402 => {
if codex_looks_like_workspace_deactivated(err_msg.as_deref()) {
quota_window_coverage =
aether_admin::provider::quota::CodexQuotaWindowCoverage::Patch;
let mut codex_meta = metadata_update
.as_ref()
.and_then(|value| value.get("codex"))
@@ -627,6 +1146,8 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
oauth_invalid_reason = reason;
status = "workspace_deactivated".to_string();
} else {
quota_window_coverage =
aether_admin::provider::quota::CodexQuotaWindowCoverage::Patch;
let plan_type = transport
.key
.decrypted_auth_config
@@ -667,24 +1188,45 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
}
}
let persisted = if let Some(expected_auth_config) = quota_auth_config_fence.as_deref() {
let persisted = if let Some(expected_credential) = quota_credential_fence.as_ref() {
persist_fenced_provider_quota_refresh_state(
state,
&key.id,
expected_auth_config,
expected_credential.encrypted_auth_config.as_str(),
metadata_update.as_ref(),
oauth_invalid_at_unix_secs,
oauth_invalid_reason.clone(),
aether_admin::provider::quota::CodexQuotaMergeContext {
observed_at_unix_secs: now_unix_secs,
request_started_at_unix_ms: Some(quota_request_started_at_unix_ms),
request_order_id: Some(quota_request_order_id),
observed_reset_generation,
authoritative_reset_generation,
observed_credential_generation: observed_credential_generation.as_deref(),
account_reset_fence_id,
coverage: quota_window_coverage,
},
Some(&expected_credential.credential),
)
.await?
} else {
persist_provider_quota_refresh_state(
persist_codex_provider_quota_refresh_state(
state,
&key.id,
metadata_update.as_ref(),
oauth_invalid_at_unix_secs,
oauth_invalid_reason.clone(),
None,
aether_admin::provider::quota::CodexQuotaMergeContext {
observed_at_unix_secs: now_unix_secs,
request_started_at_unix_ms: Some(quota_request_started_at_unix_ms),
request_order_id: Some(quota_request_order_id),
observed_reset_generation,
authoritative_reset_generation,
observed_credential_generation: observed_credential_generation.as_deref(),
account_reset_fence_id,
coverage: quota_window_coverage,
},
)
.await?
};
@@ -698,26 +1240,57 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
}));
continue;
}
let credential_cas_delete = quota_auth_config_fence.as_ref().map(|auth_config| {
let persisted_key = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key.id))
.await?
.into_iter()
.next();
let persisted_codex_metadata = persisted_key
.as_ref()
.and_then(|key| key.upstream_metadata.as_ref())
.and_then(serde_json::Value::as_object)
.and_then(|metadata| metadata.get("codex"))
.cloned();
if let Some(codex_metadata) = persisted_codex_metadata.as_ref() {
metadata_update = Some(json!({"codex": codex_metadata}));
}
let persisted_codex_object = persisted_codex_metadata
.as_ref()
.and_then(serde_json::Value::as_object);
let request_owns_persisted_oauth_state = quota_credential_fence.is_none()
|| (persisted_codex_object
.and_then(|codex| codex.get("oauth_state_request_started_at_unix_ms"))
.and_then(aether_admin::provider::quota::coerce_json_u64)
== Some(quota_request_started_at_unix_ms)
&& persisted_codex_object
.and_then(|codex| codex.get("oauth_state_request_id"))
.and_then(serde_json::Value::as_str)
== Some(quota_request_order_id));
let credential_cas_delete = quota_credential_fence.as_ref().map(|credential_fence| {
ProviderCatalogKeyOAuthCredentialCasDelete {
key_id: key.id.clone(),
expected_encrypted_auth_config: Some(auth_config.clone()),
expected_credential: ProviderCatalogKeyOAuthCredentialFence {
encrypted_api_key: key.encrypted_api_key.clone(),
auth_type: key.auth_type.clone(),
provider_id: key.provider_id.clone(),
provider_type: provider.provider_type.clone(),
},
expected_encrypted_auth_config: Some(
credential_fence.encrypted_auth_config.clone(),
),
expected_credential: credential_fence.credential.clone(),
expected_upstream_metadata_namespace: Some(
ProviderCatalogUpstreamMetadataNamespaceExpectation {
namespace: "codex".to_string(),
expected_value: persisted_codex_metadata.clone(),
},
),
}
});
let should_auto_remove_hard_banned =
provider_auto_remove_banned_keys(provider.config.as_ref())
&& should_auto_remove_oauth_invalid_key(
&key,
oauth_invalid_reason.as_deref(),
matches!(status_code, Some(401 | 403)),
now_unix_secs,
);
let should_auto_remove_hard_banned = request_owns_persisted_oauth_state
&& provider_auto_remove_banned_keys(provider.config.as_ref())
&& should_auto_remove_oauth_invalid_key(
persisted_key.as_ref().unwrap_or(&key),
persisted_key
.as_ref()
.and_then(|key| key.oauth_invalid_reason.as_deref()),
matches!(status_code, Some(401 | 403)),
now_unix_secs,
);
let auto_removed_hard_banned = if should_auto_remove_hard_banned {
match credential_cas_delete.as_ref() {
Some(delete) => {
@@ -735,6 +1308,7 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
auto_removed_hard_banned_count += 1;
}
let auto_removed_quota_exhausted = if !auto_removed_hard_banned
&& request_owns_persisted_oauth_state
&& status == "quota_exhausted"
&& provider_auto_remove_quota_exhausted_keys(provider.config.as_ref())
{
@@ -798,7 +1372,9 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
}
if let Some(quota_snapshot) = build_quota_snapshot_payload(
"codex",
key.status_snapshot.as_ref(),
persisted_key
.as_ref()
.and_then(|key| key.status_snapshot.as_ref()),
metadata_update.as_ref(),
) {
payload.insert("quota_snapshot".to_string(), quota_snapshot);
@@ -866,6 +1442,19 @@ mod tests {
);
}
#[test]
fn codex_reset_credit_only_allows_usage_drop_after_confirmed_redemption() {
assert!(codex_reset_credit_outcome_allows_usage_drop("reset"));
assert!(codex_reset_credit_outcome_allows_usage_drop(
"already_redeemed"
));
assert!(!codex_reset_credit_outcome_allows_usage_drop(
"nothing_to_reset"
));
assert!(!codex_reset_credit_outcome_allows_usage_drop("no_credit"));
assert!(!codex_reset_credit_outcome_allows_usage_drop("unknown"));
}
#[test]
fn codex_reset_credit_detail_failure_records_attempt_time() {
let mut metadata = Map::new();
@@ -880,4 +1469,29 @@ mod tests {
Some(&json!(1_777_000_000u64))
);
}
#[test]
fn codex_quota_coverage_only_replaces_observed_window_families() {
assert_eq!(
codex_quota_window_coverage(Some(&json!({"credits":{"balance":5}}))),
aether_admin::provider::quota::CodexQuotaWindowCoverage::Patch
);
assert_eq!(
codex_quota_window_coverage(None),
aether_admin::provider::quota::CodexQuotaWindowCoverage::Patch
);
assert_eq!(
codex_quota_window_coverage(Some(&json!({
"rate_limit":{"primary_window":{}}
}))),
aether_admin::provider::quota::CodexQuotaWindowCoverage::AccountSnapshot
);
assert_eq!(
codex_quota_window_coverage(Some(&json!({
"rate_limit":{"primary_window":{}},
"additional_rate_limits":[]
}))),
aether_admin::provider::quota::CodexQuotaWindowCoverage::FullSnapshot
);
}
}
@@ -796,6 +796,7 @@ mod tests {
candidate_id: None,
status_code: 403,
headers: BTreeMap::new(),
response_observation: None,
body: Some(ResponseBody {
json_body: None,
body_bytes_b64: Some(base64::engine::general_purpose::STANDARD.encode(body)),
File diff suppressed because it is too large Load Diff
@@ -399,7 +399,8 @@ async fn provider_query_read_cached_models(
let cache_key = format!("upstream_models:{provider_id}:{key_id}");
let raw = state.runtime_state().kv_get(&cache_key).await.ok()??;
let parsed = serde_json::from_str::<Vec<Value>>(&raw).ok()?;
Some(aggregate_models_for_cache(&parsed))
let models = aggregate_models_for_cache(&parsed);
(!models.is_empty()).then_some(models)
}
async fn provider_query_read_provider_cached_models(
@@ -409,7 +410,8 @@ async fn provider_query_read_provider_cached_models(
let cache_key = format!("{ANTIGRAVITY_PROVIDER_CACHE_KEY_PREFIX}{provider_id}");
let raw = state.runtime_state().kv_get(&cache_key).await.ok()??;
let parsed = serde_json::from_str::<Vec<Value>>(&raw).ok()?;
Some(aggregate_models_for_cache(&parsed))
let models = aggregate_models_for_cache(&parsed);
(!models.is_empty()).then_some(models)
}
async fn provider_query_write_provider_cached_models(
@@ -417,7 +419,11 @@ async fn provider_query_write_provider_cached_models(
provider_id: &str,
models: &[Value],
) {
let Ok(serialized) = serde_json::to_string(&aggregate_models_for_cache(models)) else {
let models = aggregate_models_for_cache(models);
if models.is_empty() {
return;
}
let Ok(serialized) = serde_json::to_string(&models) else {
return;
};
let cache_key = format!("{ANTIGRAVITY_PROVIDER_CACHE_KEY_PREFIX}{provider_id}");
@@ -577,7 +583,7 @@ async fn provider_query_fetch_models_for_key(
};
all_errors.extend(outcome.errors);
let unique_models = aggregate_models_for_cache(&outcome.cached_models);
let unique_models = outcome.legacy_models;
if outcome.has_success && !unique_models.is_empty() {
<AppState as ModelFetchRuntimeState>::write_upstream_models_cache(
state.app(),
@@ -635,7 +635,6 @@ fn provider_query_build_test_request_body_for_api_format_with_search_session(
"model": model,
"input": message,
"max_output_tokens": 30,
"temperature": 0.7,
"stream": true,
}),
"openai:search" => json!({
@@ -654,7 +653,6 @@ fn provider_query_build_test_request_body_for_api_format_with_search_session(
"content": message
}],
"max_tokens": 30,
"temperature": 0.7,
"stream": true,
}),
_ => json!({
@@ -664,7 +662,6 @@ fn provider_query_build_test_request_body_for_api_format_with_search_session(
"content": message
}],
"max_tokens": 30,
"temperature": 0.7,
"stream": true,
}),
}
@@ -811,7 +808,6 @@ fn provider_query_build_test_request_body_with_model_policy(
"content": provider_query_extract_message(payload)
.unwrap_or_else(|| DEFAULT_PROVIDER_QUERY_TEST_MESSAGE.to_string())
}],
"temperature": 0.7,
"stream": true,
})
}
@@ -1304,6 +1300,10 @@ fn provider_query_pool_catalog_key_context(
provider_type,
quota_snapshot,
),
quota_hard_blocked: admin_provider_pool_pure::admin_pool_key_quota_hard_blocked(
key,
provider_type,
),
health_score,
latency_avg_ms,
catalog_lru_score: Some(key.last_used_at_unix_secs.unwrap_or(0) as f64),
@@ -178,6 +178,7 @@ fn provider_query_execution_json_body_decodes_stream_encoded_json_response() {
"content-type".to_string(),
"application/json".to_string(),
)]),
response_observation: None,
body: Some(aether_contracts::ResponseBody {
json_body: None,
body_bytes_b64: Some(encoded_body),
@@ -242,6 +243,26 @@ fn provider_query_default_test_request_body_does_not_set_max_tokens() {
);
}
#[test]
fn provider_query_default_test_request_bodies_do_not_set_temperature() {
let payload = json!({});
let default_body = provider_query_build_test_request_body(&payload, "fallback-model");
assert!(default_body.get("temperature").is_none());
for api_format in ["openai:chat", "openai:responses", "claude:messages"] {
let body = provider_query_build_test_request_body_for_api_format(
&payload,
"fallback-model",
"/api/admin/provider-query/test-model",
api_format,
);
assert!(
body.get("temperature").is_none(),
"admin model test must not set temperature for {api_format}"
);
}
}
#[test]
fn provider_query_failover_request_body_overrides_custom_model() {
let payload = json!({
@@ -416,6 +437,7 @@ fn provider_query_standard_test_aggregates_responses_stream_body() {
candidate_id: Some("candidate-0".to_string()),
status_code: 200,
headers: BTreeMap::new(),
response_observation: None,
body: Some(aether_contracts::ResponseBody {
json_body: None,
body_bytes_b64: Some(
@@ -448,6 +470,7 @@ fn provider_query_standard_test_aggregates_responses_image_generation_call() {
candidate_id: Some("candidate-0".to_string()),
status_code: 200,
headers: BTreeMap::new(),
response_observation: None,
body: Some(aether_contracts::ResponseBody {
json_body: None,
body_bytes_b64: Some(
@@ -581,6 +604,7 @@ fn provider_query_search_success_requires_non_empty_output() {
candidate_id: Some("candidate-0".to_string()),
status_code: 200,
headers: BTreeMap::new(),
response_observation: None,
body: Some(aether_contracts::ResponseBody {
json_body: Some(body),
body_bytes_b64: None,
@@ -731,6 +755,7 @@ fn provider_query_standard_test_rejects_gemini_success_without_visible_output()
candidate_id: Some("candidate-0".to_string()),
status_code: 200,
headers: BTreeMap::new(),
response_observation: None,
body: Some(aether_contracts::ResponseBody {
json_body: Some(json!({
"candidates": [{
@@ -122,6 +122,7 @@ pub(crate) struct AdminProviderQuotaRefreshRequest {
#[derive(Debug, Deserialize)]
pub(crate) struct AdminCodexResetCreditConsumeRequest {
pub(crate) idempotency_key: String,
pub(crate) expected_credential_generation: serde_json::Value,
}
#[derive(Debug, Deserialize)]
@@ -151,6 +152,10 @@ pub(crate) struct AdminProviderCreateRequest {
#[serde(default)]
pub(crate) keep_priority_on_conversion: Option<bool>,
#[serde(default)]
pub(crate) codex_fingerprint_convergence_enabled: Option<bool>,
#[serde(default)]
pub(crate) responses_websocket_enabled: Option<bool>,
#[serde(default)]
pub(crate) is_active: Option<bool>,
#[serde(default)]
pub(crate) concurrent_limit: Option<i32>,
@@ -210,6 +215,10 @@ pub(crate) struct AdminProviderUpdateRequest {
#[serde(default)]
pub(crate) keep_priority_on_conversion: Option<bool>,
#[serde(default)]
pub(crate) codex_fingerprint_convergence_enabled: Option<bool>,
#[serde(default)]
pub(crate) responses_websocket_enabled: Option<bool>,
#[serde(default)]
pub(crate) is_active: Option<bool>,
#[serde(default)]
pub(crate) concurrent_limit: Option<i32>,
@@ -335,3 +344,37 @@ pub(crate) struct AdminImportProviderModelsRequest {
)]
pub(crate) price_per_request: Option<f64>,
}
#[cfg(test)]
mod tests {
use super::AdminCodexResetCreditConsumeRequest;
#[test]
fn codex_reset_credit_consume_requires_an_explicit_credential_generation() {
assert!(
serde_json::from_value::<AdminCodexResetCreditConsumeRequest>(
serde_json::json!({"idempotency_key":"reset-old-client"}),
)
.is_err()
);
let legacy_account =
serde_json::from_value::<AdminCodexResetCreditConsumeRequest>(serde_json::json!({
"idempotency_key":"reset-legacy-account",
"expected_credential_generation":null,
}))
.expect("explicit null should fence an account without a generation");
assert!(legacy_account.expected_credential_generation.is_null());
let generated_account =
serde_json::from_value::<AdminCodexResetCreditConsumeRequest>(serde_json::json!({
"idempotency_key":"reset-generated-account",
"expected_credential_generation":"credential-v2",
}))
.expect("string generation should deserialize");
assert_eq!(
generated_account.expected_credential_generation,
serde_json::json!("credential-v2")
);
}
}
@@ -4,7 +4,7 @@ use crate::handlers::admin::provider::shared::support::{
};
use crate::handlers::admin::shared::unix_secs_to_rfc3339;
use crate::handlers::public::{request_candidate_event_unix_ms, request_candidate_status_label};
use crate::orchestration::codex_cyber_flag_passthrough_enabled;
use crate::orchestration::{codex_cyber_flag_passthrough_enabled, responses_websocket_adapter};
use crate::provider_key_auth::provider_key_effective_api_formats;
use aether_data_contracts::repository::candidates::{
RequestCandidateStatus, StoredRequestCandidate,
@@ -218,6 +218,11 @@ pub(crate) fn build_admin_provider_summary_value(
"ops_architecture_id": ops_architecture_id,
"kiro_simulated_cache_enabled": kiro_simulated_cache_enabled,
"codex_cyber_flag_passthrough_enabled": codex_cyber_flag_passthrough_enabled(&provider.provider_type, provider.config.as_ref()),
"codex_fingerprint_convergence_enabled": crate::provider_transport::codex_fingerprint_convergence_enabled(
&provider.provider_type,
provider.config.as_ref(),
),
"responses_websocket_enabled": responses_websocket_adapter(&provider.provider_type, provider.config.as_ref()).is_some(),
"ops_quota_alert_enabled": ops_quota_alert_enabled,
"created_at": endpoint_timestamp_or_now(provider.created_at_unix_ms, now_unix_secs),
"updated_at": endpoint_timestamp_or_now(provider.updated_at_unix_secs, now_unix_secs),
@@ -1,3 +1,4 @@
use crate::handlers::admin::provider::oauth::provisioning::rotate_codex_credential_generation;
use crate::handlers::admin::provider::shared::payloads::AdminProviderKeyCreateRequest;
use crate::handlers::admin::provider::write::normalize::{
normalize_allow_auth_channel_mismatch_formats, normalize_api_format_json_object_keys,
@@ -216,6 +217,7 @@ pub(crate) async fn build_admin_create_provider_key_record(
)?;
key.created_at_unix_ms = Some(now_unix_secs);
key.updated_at_unix_secs = Some(now_unix_secs);
rotate_codex_credential_generation(&mut key, &provider.provider_type);
Ok(key)
}
@@ -6,6 +6,7 @@ pub(crate) use self::update::build_admin_update_provider_key_record;
pub(crate) use self::update::{
admin_provider_key_update_requires_immediate_model_fetch,
build_admin_update_provider_key_record_with_existing_keys,
build_provider_catalog_key_admin_cas_update,
};
mod batch;
@@ -1,3 +1,4 @@
use crate::handlers::admin::provider::oauth::provisioning::rotate_codex_credential_generation;
use crate::handlers::admin::provider::shared::payloads::AdminProviderKeyUpdatePatch;
use crate::handlers::admin::provider::write::normalize::{
normalize_allow_auth_channel_mismatch_formats, normalize_api_format_json_object_keys,
@@ -13,6 +14,7 @@ use crate::handlers::admin::shared::{
use crate::handlers::shared::normalize_optional_api_key_concurrent_limit;
use crate::provider_key_auth::provider_key_is_oauth_managed;
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyOAuthCredentialFence,
StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use aether_provider_transport::provider_types::provider_type_is_fixed;
@@ -368,6 +370,12 @@ pub(crate) fn build_admin_update_provider_key_record_with_existing_keys(
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs());
let credential_identity_changed = !updated.auth_type.eq_ignore_ascii_case(&existing.auth_type)
|| updated.encrypted_api_key != existing.encrypted_api_key
|| updated.encrypted_auth_config != existing.encrypted_auth_config;
if credential_identity_changed {
rotate_codex_credential_generation(&mut updated, &provider.provider_type);
}
Ok(updated)
}
@@ -382,6 +390,49 @@ pub(crate) fn admin_provider_key_update_requires_immediate_model_fetch(
&& (!existing.auto_fetch_models || filters_changed || locked_models_changed)
}
pub(crate) fn build_provider_catalog_key_admin_cas_update(
existing: &StoredProviderCatalogKey,
updated: StoredProviderCatalogKey,
provider_type: &str,
) -> ProviderCatalogKeyAdminCasUpdate {
let previous_generation = existing
.upstream_metadata
.as_ref()
.and_then(|metadata| metadata.pointer("/codex/credential_generation"))
.and_then(serde_json::Value::as_str);
let next_generation = updated
.upstream_metadata
.as_ref()
.and_then(|metadata| metadata.pointer("/codex/credential_generation"))
.and_then(serde_json::Value::as_str);
let credential_changed = existing.auth_type != updated.auth_type
|| existing.encrypted_api_key != updated.encrypted_api_key
|| existing.encrypted_auth_config != updated.encrypted_auth_config;
let codex_rotation = provider_type
.trim()
.eq_ignore_ascii_case("codex")
.then(|| next_generation.filter(|next| Some(*next) != previous_generation))
.flatten()
.map(|generation| {
json!({
aether_admin::provider::quota::CODEX_CREDENTIAL_GENERATION_KEY: generation,
})
});
ProviderCatalogKeyAdminCasUpdate {
expected_encrypted_auth_config: existing.encrypted_auth_config.clone(),
expected_credential: ProviderCatalogKeyOAuthCredentialFence {
encrypted_api_key: existing.encrypted_api_key.clone(),
auth_type: existing.auth_type.clone(),
provider_id: existing.provider_id.clone(),
provider_type: provider_type.to_string(),
},
key: updated,
codex_rotation,
reset_oauth_runtime: credential_changed,
}
}
fn raw_secret_auth_type(value: &str) -> bool {
matches!(
value.trim().to_ascii_lowercase().as_str(),
@@ -208,6 +208,53 @@ pub(crate) fn normalize_chat_pii_redaction_config(
}
}
pub(crate) fn set_responses_websocket_enabled(
config: &mut serde_json::Map<String, serde_json::Value>,
enabled: bool,
) -> Result<(), String> {
let mut responses = match config.remove("responses_websocket") {
None => serde_json::Map::new(),
Some(serde_json::Value::Object(config)) => config,
Some(_) => return Err("config.responses_websocket 必须是 JSON 对象".to_string()),
};
responses.insert("enabled".to_string(), serde_json::Value::Bool(enabled));
config.insert(
"responses_websocket".to_string(),
serde_json::Value::Object(responses),
);
Ok(())
}
pub(crate) fn remove_responses_websocket_enabled(
config: &mut serde_json::Map<String, serde_json::Value>,
) {
let Some(serde_json::Value::Object(responses)) = config.get_mut("responses_websocket") else {
return;
};
responses.remove("enabled");
if responses.is_empty() {
config.remove("responses_websocket");
}
}
pub(crate) fn validate_responses_websocket_config(
config: &serde_json::Map<String, serde_json::Value>,
) -> Result<(), String> {
if let Some(value) = config.get("responses_websocket") {
let responses = value
.as_object()
.ok_or_else(|| "config.responses_websocket 必须是 JSON 对象".to_string())?;
let enabled = responses
.get("enabled")
.ok_or_else(|| "config.responses_websocket.enabled 为必填布尔值".to_string())?;
if !enabled.is_boolean() {
return Err("config.responses_websocket.enabled 必须是布尔值".to_string());
}
}
Ok(())
}
pub(crate) fn validate_vertex_api_formats(
provider_type: &str,
auth_type: &str,
@@ -258,7 +305,9 @@ mod tests {
normalize_api_format_list, normalize_auth_type, normalize_auth_type_by_format,
normalize_chat_pii_redaction_config, normalize_pool_advanced_config,
normalize_provider_type_input, normalize_rate_multipliers,
reconcile_allow_auth_channel_mismatch_formats, validate_vertex_api_formats,
reconcile_allow_auth_channel_mismatch_formats, remove_responses_websocket_enabled,
set_responses_websocket_enabled, validate_responses_websocket_config,
validate_vertex_api_formats,
};
use serde_json::json;
@@ -317,6 +366,21 @@ mod tests {
);
}
#[test]
fn responses_websocket_setting_is_available_to_explicitly_enabled_providers() {
let mut config = serde_json::Map::new();
set_responses_websocket_enabled(&mut config, true)
.expect("Responses setting should be accepted");
assert_eq!(
config.get("responses_websocket"),
Some(&json!({"enabled": true}))
);
validate_responses_websocket_config(&config).expect("Responses setting should validate");
remove_responses_websocket_enabled(&mut config);
assert!(config.get("responses_websocket").is_none());
}
#[test]
fn normalize_auth_type_supports_bearer() {
assert_eq!(
@@ -7,6 +7,8 @@ use crate::handlers::admin::provider::shared::support::{
use crate::handlers::admin::provider::write::normalize::normalize_chat_pii_redaction_config;
use crate::handlers::admin::provider::write::normalize::normalize_pool_advanced_config;
use crate::handlers::admin::provider::write::normalize::normalize_provider_type_input;
use crate::handlers::admin::provider::write::normalize::set_responses_websocket_enabled;
use crate::handlers::admin::provider::write::normalize::validate_responses_websocket_config;
use crate::handlers::admin::request::AdminAppState;
use crate::handlers::admin::shared::normalize_json_object;
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogProvider;
@@ -140,6 +142,28 @@ pub(crate) async fn build_admin_create_provider_record(
if let Some(value) = normalize_pool_advanced_config(payload.pool_advanced)? {
config_map.insert("pool_advanced".to_string(), value);
}
if let Some(enabled) = payload.codex_fingerprint_convergence_enabled {
if provider_type != "codex" && enabled {
return Err(
"codex_fingerprint_convergence_enabled 仅适用于 provider_type=codex".to_string(),
);
}
if provider_type == "codex" {
let codex_config = config_map
.entry(crate::provider_transport::CODEX_FINGERPRINT_CONFIG_NAMESPACE.to_string())
.or_insert_with(|| json!({}));
let Some(codex_config) = codex_config.as_object_mut() else {
return Err("config.codex 必须是 JSON 对象".to_string());
};
codex_config.insert(
crate::provider_transport::CODEX_FINGERPRINT_ENABLED_CONFIG_KEY.to_string(),
json!(enabled),
);
}
}
if provider_type != "codex" {
remove_codex_fingerprint_config(&mut config_map);
}
if let Some(value) = normalize_json_object(payload.failover_rules, "failover_rules")? {
config_map.insert("failover_rules".to_string(), value);
}
@@ -157,6 +181,10 @@ pub(crate) async fn build_admin_create_provider_record(
config_map.insert("chat_pii_redaction".to_string(), value);
}
}
if let Some(enabled) = payload.responses_websocket_enabled {
set_responses_websocket_enabled(&mut config_map, enabled)?;
}
validate_responses_websocket_config(&config_map)?;
let config = (!config_map.is_empty()).then_some(serde_json::Value::Object(config_map));
crate::provider_transport::validate_anthropic_compatibility_profile_config(config.as_ref())
.map_err(|_| "无效的 Anthropic compatibility profile".to_string())?;
@@ -204,3 +232,46 @@ pub(crate) async fn build_admin_create_provider_record(
Ok((record, shift_existing_priorities_from))
}
fn remove_codex_fingerprint_config(config_map: &mut serde_json::Map<String, serde_json::Value>) {
let namespace = crate::provider_transport::CODEX_FINGERPRINT_CONFIG_NAMESPACE;
let key = crate::provider_transport::CODEX_FINGERPRINT_ENABLED_CONFIG_KEY;
let mut remove_namespace = false;
if let Some(codex_config) = config_map
.get_mut(namespace)
.and_then(|value| value.as_object_mut())
{
codex_config.remove(key);
remove_namespace = codex_config.is_empty();
}
if remove_namespace {
config_map.remove(namespace);
}
}
#[cfg(test)]
mod tests {
use serde_json::json;
#[test]
fn removing_fingerprint_setting_preserves_other_codex_config() {
let mut config = json!({
"codex": {
"fingerprint_convergence_enabled": true,
"pass_through_cyber_flag_interrupt": true
},
"other": {"kept": true}
})
.as_object()
.expect("config object")
.clone();
super::remove_codex_fingerprint_config(&mut config);
assert_eq!(
config["codex"],
json!({"pass_through_cyber_flag_interrupt": true})
);
assert_eq!(config["other"], json!({"kept": true}));
}
}
@@ -7,6 +7,8 @@ use crate::handlers::admin::provider::shared::support::{
use crate::handlers::admin::provider::write::normalize::normalize_chat_pii_redaction_config;
use crate::handlers::admin::provider::write::normalize::normalize_pool_advanced_config;
use crate::handlers::admin::provider::write::normalize::normalize_provider_type_input;
use crate::handlers::admin::provider::write::normalize::set_responses_websocket_enabled;
use crate::handlers::admin::provider::write::normalize::validate_responses_websocket_config;
use crate::handlers::admin::request::AdminAppState;
use crate::handlers::admin::shared::normalize_json_object;
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogProvider;
@@ -244,6 +246,32 @@ pub(crate) async fn build_admin_update_provider_record(
}
}
if fields.contains("codex_fingerprint_convergence_enabled") {
let Some(enabled) = payload.codex_fingerprint_convergence_enabled else {
return Err("codex_fingerprint_convergence_enabled 必须是布尔值".to_string());
};
if target_provider_type != "codex" && enabled {
return Err(
"codex_fingerprint_convergence_enabled 仅适用于 provider_type=codex".to_string(),
);
}
if target_provider_type == "codex" {
let codex_config = config_map
.entry(crate::provider_transport::CODEX_FINGERPRINT_CONFIG_NAMESPACE.to_string())
.or_insert_with(|| json!({}));
let Some(codex_config) = codex_config.as_object_mut() else {
return Err("config.codex 必须是 JSON 对象".to_string());
};
codex_config.insert(
crate::provider_transport::CODEX_FINGERPRINT_ENABLED_CONFIG_KEY.to_string(),
json!(enabled),
);
}
}
if target_provider_type != "codex" {
remove_codex_fingerprint_config(&mut config_map);
}
for (field_name, payload_value) in [
(
PROVIDER_MAX_TRANSFER_COUNT_CONFIG_KEY,
@@ -311,6 +339,14 @@ pub(crate) async fn build_admin_update_provider_record(
}
}
if fields.contains("responses_websocket_enabled") {
let enabled = payload
.responses_websocket_enabled
.ok_or_else(|| "responses_websocket_enabled 必须是布尔值".to_string())?;
set_responses_websocket_enabled(&mut config_map, enabled)?;
}
validate_responses_websocket_config(&config_map)?;
updated.config = (!config_map.is_empty()).then_some(serde_json::Value::Object(config_map));
crate::provider_transport::validate_anthropic_compatibility_profile_config(
updated.config.as_ref(),
@@ -322,3 +358,46 @@ pub(crate) async fn build_admin_update_provider_record(
.map(|duration| duration.as_secs());
Ok(updated)
}
fn remove_codex_fingerprint_config(config_map: &mut serde_json::Map<String, serde_json::Value>) {
let namespace = crate::provider_transport::CODEX_FINGERPRINT_CONFIG_NAMESPACE;
let key = crate::provider_transport::CODEX_FINGERPRINT_ENABLED_CONFIG_KEY;
let mut remove_namespace = false;
if let Some(codex_config) = config_map
.get_mut(namespace)
.and_then(|value| value.as_object_mut())
{
codex_config.remove(key);
remove_namespace = codex_config.is_empty();
}
if remove_namespace {
config_map.remove(namespace);
}
}
#[cfg(test)]
mod tests {
use serde_json::json;
#[test]
fn removing_fingerprint_setting_preserves_other_codex_config() {
let mut config = json!({
"codex": {
"fingerprint_convergence_enabled": true,
"pass_through_cyber_flag_interrupt": true
},
"other": {"kept": true}
})
.as_object()
.expect("config object")
.clone();
super::remove_codex_fingerprint_config(&mut config);
assert_eq!(
config["codex"],
json!({"pass_through_cyber_flag_interrupt": true})
);
assert_eq!(config["other"], json!({"kept": true}));
}
}
@@ -186,6 +186,15 @@ impl<'a> AdminAppState<'a> {
self.app.update_provider_catalog_key(key).await
}
pub(crate) async fn compare_and_update_provider_catalog_key_admin_state(
&self,
update: &aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyAdminCasUpdate,
) -> Result<bool, GatewayError> {
self.app
.compare_and_update_provider_catalog_key_admin_state(update)
.await
}
pub(crate) async fn compare_and_update_provider_catalog_key_adaptive_state(
&self,
update: &aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyAdaptiveStateUpdate,
@@ -513,6 +513,7 @@ impl<'a> AdminAppState<'a> {
provider_id: key.provider_id.clone(),
provider_type: provider.provider_type.clone(),
},
expected_upstream_metadata_namespace: None,
},
)
.await
@@ -4,10 +4,12 @@ use crate::api::ai::admin_endpoint_signature_parts;
use crate::handlers::admin::admin_provider_pool_config;
use crate::handlers::admin::model::ADMIN_EXTERNAL_MODELS_PROXY_NODE_CONFIG_KEY;
use crate::handlers::admin::provider::endpoints_admin::payloads::AdminProviderEndpointUpdatePatch;
use crate::handlers::admin::provider::oauth::provisioning::ensure_codex_credential_generation_rotated;
use crate::handlers::admin::provider::shared::payloads::{
AdminProviderCreateRequest, AdminProviderKeyCreateRequest, AdminProviderKeyUpdatePatch,
AdminProviderUpdatePatch,
};
use crate::handlers::admin::provider::write::keys::build_provider_catalog_key_admin_cas_update;
use crate::handlers::admin::shared::{
normalize_json_array, normalize_json_object, normalize_string_list,
};
@@ -377,10 +379,14 @@ fn normalize_import_key_raw_payload(
fn apply_imported_oauth_key_credentials(
state: &AdminAppState<'_>,
provider_type: &str,
previous_codex_credential_generation: Option<&str>,
raw_key: &Map<String, Value>,
normalized_auth_config: Option<&Value>,
record: &mut aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey,
) -> Result<bool, String> {
let previous_encrypted_api_key = record.encrypted_api_key.clone();
let previous_encrypted_auth_config = record.encrypted_auth_config.clone();
let mut credentials_supplied = false;
let mut api_key_supplied = false;
if let Some(api_key_value) = raw_key.get("api_key") {
@@ -424,10 +430,19 @@ fn apply_imported_oauth_key_credentials(
api_key_supplied,
);
let credential_material_changed = record.encrypted_api_key != previous_encrypted_api_key
|| record.encrypted_auth_config != previous_encrypted_auth_config;
if credentials_supplied {
record.oauth_invalid_at_unix_secs = None;
record.oauth_invalid_reason = None;
}
if credential_material_changed {
ensure_codex_credential_generation_rotated(
record,
provider_type,
previous_codex_credential_generation,
);
}
Ok(credentials_supplied)
}
@@ -1861,6 +1876,15 @@ impl<'a> AdminAppState<'a> {
if let Some(existing_index) = existing_key_index {
let existing_key = existing_keys[existing_index].clone();
let previous_codex_credential_generation = existing_key
.upstream_metadata
.as_ref()
.and_then(Value::as_object)
.and_then(|metadata| metadata.get("codex"))
.and_then(|codex| {
aether_admin::provider::quota::codex_credential_generation(Some(codex))
})
.map(ToOwned::to_owned);
match merge_mode {
AdminImportMergeMode::Skip => {
stats.keys.skipped += 1;
@@ -1890,6 +1914,8 @@ impl<'a> AdminAppState<'a> {
let oauth_credentials_supplied = if auth_type == "oauth" {
invalid!(apply_imported_oauth_key_credentials(
self,
&provider.provider_type,
previous_codex_credential_generation.as_deref(),
&raw_key,
normalized_auth_config.as_ref(),
&mut updated,
@@ -1903,8 +1929,31 @@ impl<'a> AdminAppState<'a> {
imported_key.fingerprint.clone(),
"fingerprint",
));
let Some(mut persisted) =
self.update_provider_catalog_key(&updated).await?
let admin_update = build_provider_catalog_key_admin_cas_update(
&existing_key,
updated.clone(),
&provider.provider_type,
);
if !self
.compare_and_update_provider_catalog_key_admin_state(&admin_update)
.await?
{
return Ok(Err((
http::StatusCode::CONFLICT,
json!({
"detail": format!(
"Provider '{provider_name}' 的 Key 已被其他请求更新,请重试"
)
}),
)));
}
let Some(mut persisted) = self
.read_provider_catalog_keys_by_ids(std::slice::from_ref(
&updated.id,
))
.await?
.into_iter()
.next()
else {
return Ok(Err(invalid_request(format!(
"更新 Provider '{provider_name}' 的 Key 失败"
@@ -1926,16 +1975,15 @@ impl<'a> AdminAppState<'a> {
persisted = reloaded;
}
if oauth_credentials_supplied {
if !self
.clear_provider_catalog_key_oauth_invalid_marker(&updated.id)
.await?
{
return Ok(Err(invalid_request(format!(
"更新 Provider '{provider_name}' 的 Key 失败"
))));
}
let Some(reloaded) = self
.reset_provider_catalog_key_recovery_state(&updated.id)
.reset_provider_catalog_key_recovery_state_fenced(
&updated.id,
updated.encrypted_auth_config.as_deref().ok_or_else(|| {
GatewayError::Internal(format!(
"OAuth Provider '{provider_name}' imported without auth_config"
))
})?,
)
.await?
else {
return Ok(Err(invalid_request(format!(
@@ -1975,6 +2023,8 @@ impl<'a> AdminAppState<'a> {
let oauth_credentials_supplied = if auth_type == "oauth" {
invalid!(apply_imported_oauth_key_credentials(
self,
&provider.provider_type,
None,
&raw_key,
normalized_auth_config.as_ref(),
&mut record,
@@ -1,5 +1,6 @@
mod body_buffer;
mod local;
mod websocket;
use self::body_buffer::{
buffer_and_normalize_request_body, build_request_body_buffer_error_response,
@@ -8,6 +9,7 @@ use self::body_buffer::{
use self::local::{
maybe_build_local_admin_proxy_response, maybe_build_local_internal_proxy_response,
};
pub(crate) use self::websocket::responses::responses_websocket;
use super::internal::resolve_local_proxy_execution_path;
pub(crate) use super::public::matches_model_mapping_for_models;
use crate::ai_serving::api::{
@@ -0,0 +1,527 @@
//! Authenticated public WebSocket upgrade admission shared by AI adapters.
use std::future::Future;
use std::net::{IpAddr, SocketAddr};
use axum::body::Body;
use axum::extract::ws::{WebSocket, WebSocketUpgrade};
use axum::http::header::{
AUTHORIZATION, CONNECTION, COOKIE, HOST, PROXY_AUTHORIZATION, TE, TRAILER, TRANSFER_ENCODING,
UPGRADE,
};
use axum::http::uri::PathAndQuery;
use axum::http::{HeaderMap, HeaderName, Method, Response, StatusCode, Uri};
use tracing::{info, warn};
use crate::api::response::{
build_local_auth_rejection_response, build_local_http_error_response,
build_local_overloaded_response,
};
use crate::control::{
trusted_auth_local_rejection, GatewayControlDecision, GatewayCredentialCarrier,
GatewayLocalAuthRejection,
};
use crate::handlers::proxy::websocket::session::{WebSocketSessionLimits, WEBSOCKET_LOG_TRANSPORT};
use crate::handlers::shared::ip_rules_allow;
use crate::headers::{effective_client_ip, extract_or_generate_trace_id};
use crate::router::RequestAdmissionError;
use crate::{AppState, GatewayError};
/// Request facts that survive the HTTP Upgrade and are needed by a protocol
/// adapter for planning, rate limiting, and connection-scoped audit logs.
pub(crate) struct WebSocketRequestContext {
pub(crate) trace_id: String,
pub(crate) headers: HeaderMap,
pub(crate) uri: Uri,
pub(crate) remote_addr: SocketAddr,
/// Effective client IP resolved once from the authenticated Upgrade. Every
/// turn re-checks live API-key/admin IP policy against this immutable fact.
pub(crate) client_ip: IpAddr,
pub(crate) decision: GatewayControlDecision,
/// Held for the lifetime of the upgraded socket. The Responses session
/// polls its health and closes the client when a distributed lease is
/// revoked or expires.
pub(crate) websocket_connection_permit: Option<aether_runtime::AdmissionPermit>,
}
/// Adapter-specific wording and event identifiers for generic upgrade checks.
#[derive(Clone, Copy)]
pub(crate) struct WebSocketIngressSpec {
pub(crate) route_unavailable_message: &'static str,
}
/// Performs the HTTP-only part of an AI WebSocket request.
///
/// The ordinary request permit covers only the HTTP Upgrade window. A
/// dedicated WebSocket connection permit is held for the socket lifetime so
/// idle clients cannot consume capacity reserved for normal HTTP requests.
pub(crate) async fn upgrade_authenticated_ai_websocket<F, Fut>(
state: AppState,
remote_addr: SocketAddr,
ws: WebSocketUpgrade,
headers: HeaderMap,
uri: Uri,
limits: WebSocketSessionLimits,
spec: WebSocketIngressSpec,
run_session: F,
) -> Result<Response<Body>, GatewayError>
where
F: FnOnce(WebSocket, AppState, WebSocketRequestContext) -> Fut + Send + 'static,
Fut: Future<Output = ()> + Send + 'static,
{
let trace_id = extract_or_generate_trace_id(&headers);
let client_ip = effective_client_ip(&headers, &remote_addr);
if state.admin_security_ip_blacklisted(client_ip).await? {
return build_local_http_error_response(
&trace_id,
None,
StatusCode::FORBIDDEN,
"当前 IP 已被禁止访问",
);
}
let request_context = crate::control::resolve_public_request_context(
&state,
&Method::GET,
&uri,
&headers,
&trace_id,
)
.await?;
let Some(mut decision) = request_context.control_decision else {
return build_local_http_error_response(
&trace_id,
None,
StatusCode::NOT_FOUND,
spec.route_unavailable_message,
);
};
if let Some(rejection) = trusted_auth_local_rejection(Some(&decision), &headers) {
return build_local_auth_rejection_response(&trace_id, Some(&decision), &rejection);
}
// Browsers attach cookies to WebSocket handshakes automatically and the
// WebSocket API does not let callers add an Authorization header. A
// cookie-only public upgrade would therefore be vulnerable to cross-site
// WebSocket hijacking unless every deployment maintained an Origin
// allowlist. Explicit API-key/bearer credentials (or trusted internal
// auth resolved by the control plane) remain supported.
if !websocket_credential_carrier_is_allowed(decision.gateway_credential_carrier) {
warn!(
event_name = "ai_websocket_cookie_only_auth_rejected",
log_type = "security",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %trace_id,
client_ip = %client_ip,
"gateway rejected cookie-only public WebSocket authentication"
);
return build_local_auth_rejection_response(
&trace_id,
Some(&decision),
&GatewayLocalAuthRejection::InvalidApiKey,
);
}
let Some(auth_context) = decision.auth_context.as_ref() else {
return build_local_auth_rejection_response(
&trace_id,
Some(&decision),
&GatewayLocalAuthRejection::InvalidApiKey,
);
};
if !auth_context.access_allowed
|| auth_context.user_id.trim().is_empty()
|| auth_context.api_key_id.trim().is_empty()
{
return build_local_auth_rejection_response(
&trace_id,
Some(&decision),
&GatewayLocalAuthRejection::InvalidApiKey,
);
}
if !ip_rules_allow(auth_context.ip_rules.as_deref(), client_ip) {
return build_local_auth_rejection_response(
&trace_id,
Some(&decision),
&GatewayLocalAuthRejection::IpNotAllowed {
remote_ip: client_ip.to_string(),
},
);
}
let request_permit = match state.try_acquire_request_permit().await {
Ok(permit) => permit,
Err(error) => {
return websocket_admission_error_response(
&trace_id,
&decision,
Some(uri.path()),
error,
)
}
};
let websocket_connection_permit = match state.try_acquire_websocket_connection_permit().await {
Ok(permit) => permit,
Err(error) => {
return websocket_admission_error_response(
&trace_id,
&decision,
Some(uri.path()),
error,
)
}
};
// Authentication has consumed the downstream credentials. From this
// point on the URI and headers become planner input, so retain neither an
// API key from the query string nor client authentication/handshake
// headers. Provider authentication is added independently by the
// planner and is therefore unaffected by this boundary.
let uri = websocket_planning_uri(&uri);
decision.public_query_string = uri.query().map(ToOwned::to_owned);
let headers = websocket_planning_headers(headers);
let context = WebSocketRequestContext {
trace_id,
headers,
uri,
remote_addr,
client_ip,
decision,
websocket_connection_permit,
};
Ok(ws
.max_frame_size(limits.max_frame_size)
.max_message_size(limits.max_message_size)
.on_upgrade(move |socket| async move {
drop(request_permit);
run_session(socket, state, context).await;
}))
}
fn websocket_credential_carrier_is_allowed(carrier: Option<GatewayCredentialCarrier>) -> bool {
carrier != Some(GatewayCredentialCarrier::CookieHeader)
}
fn websocket_planning_uri(uri: &Uri) -> Uri {
let Some(query) = uri.query() else {
return uri.clone();
};
let mut retained = Vec::new();
let mut removed_sensitive_value = false;
for (name, value) in url::form_urlencoded::parse(query.as_bytes()) {
if websocket_query_parameter_is_sensitive(name.as_ref()) {
removed_sensitive_value = true;
} else {
retained.push((name.into_owned(), value.into_owned()));
}
}
if !removed_sensitive_value {
return uri.clone();
}
let mut serializer = url::form_urlencoded::Serializer::new(String::new());
serializer.extend_pairs(retained.iter().map(|(name, value)| (name, value)));
let retained_query = serializer.finish();
let path_and_query = if retained_query.is_empty() {
uri.path().to_string()
} else {
format!("{}?{retained_query}", uri.path())
};
let path_and_query = path_and_query
.parse::<PathAndQuery>()
.expect("a valid URI path plus form-encoded query must remain valid");
let mut parts = uri.clone().into_parts();
parts.path_and_query = Some(path_and_query);
Uri::from_parts(parts).expect("replacing only path-and-query must preserve a valid URI")
}
fn websocket_query_parameter_is_sensitive(name: &str) -> bool {
matches!(
name.to_ascii_lowercase().as_str(),
"key" | "api_key" | "api-key" | "access_token" | "authorization" | "token"
)
}
fn websocket_planning_headers(mut headers: HeaderMap) -> HeaderMap {
// RFC 9110 permits Connection to name additional hop-by-hop fields. Read
// those names before removing Connection itself.
let connection_scoped_names = headers
.get_all(CONNECTION)
.iter()
.filter_map(|value| value.to_str().ok())
.flat_map(|value| value.split(','))
.filter_map(|name| HeaderName::from_bytes(name.trim().as_bytes()).ok())
.collect::<Vec<_>>();
for name in connection_scoped_names {
headers.remove(name);
}
for name in [
AUTHORIZATION,
CONNECTION,
COOKIE,
HOST,
PROXY_AUTHORIZATION,
TE,
TRAILER,
TRANSFER_ENCODING,
UPGRADE,
] {
headers.remove(name);
}
for name in [
"api-key",
"keep-alive",
"proxy-connection",
"x-api-key",
"x-goog-api-key",
crate::constants::GATEWAY_HEADER,
crate::constants::TRUSTED_AUTH_USER_ID_HEADER,
crate::constants::TRUSTED_AUTH_API_KEY_ID_HEADER,
crate::constants::TRUSTED_AUTH_BALANCE_HEADER,
crate::constants::TRUSTED_AUTH_ACCESS_ALLOWED_HEADER,
crate::constants::TRUSTED_ADMIN_USER_ID_HEADER,
crate::constants::TRUSTED_ADMIN_USER_ROLE_HEADER,
crate::constants::TRUSTED_ADMIN_SESSION_ID_HEADER,
crate::constants::TRUSTED_ADMIN_MANAGEMENT_TOKEN_ID_HEADER,
] {
headers.remove(name);
}
let websocket_managed_names = headers
.keys()
.filter(|name| name.as_str().starts_with("sec-websocket-"))
.cloned()
.collect::<Vec<_>>();
for name in websocket_managed_names {
headers.remove(name);
}
headers
}
fn websocket_admission_error_response(
trace_id: &str,
decision: &GatewayControlDecision,
request_path: Option<&str>,
error: RequestAdmissionError,
) -> Result<Response<Body>, GatewayError> {
match error {
RequestAdmissionError::Local(aether_runtime::ConcurrencyError::Saturated {
gate,
limit,
})
| RequestAdmissionError::Distributed(
aether_runtime_state::RuntimeSemaphoreError::Saturated { gate, limit },
)
| RequestAdmissionError::Distributed(
aether_runtime_state::RuntimeSemaphoreError::Unavailable { gate, limit, .. },
) => build_local_overloaded_response(trace_id, Some(decision), request_path, gate, limit),
RequestAdmissionError::Local(aether_runtime::ConcurrencyError::Closed { gate }) => Err(
GatewayError::Internal(format!("gateway concurrency gate {gate} is closed")),
),
RequestAdmissionError::Distributed(
aether_runtime_state::RuntimeSemaphoreError::InvalidConfiguration(message),
) => Err(GatewayError::Internal(message)),
}
}
/// Connection-level access log fields which are independent of a protocol's
/// per-turn usage lifecycle.
#[derive(Clone, Copy)]
pub(crate) struct WebSocketConnectionLogSpec {
pub(crate) opened_event_name: &'static str,
pub(crate) closed_event_name: &'static str,
pub(crate) opened_message: &'static str,
pub(crate) closed_message: &'static str,
pub(crate) execution_path: &'static str,
pub(crate) provider_type: &'static str,
}
pub(crate) struct WebSocketConnectionLog {
spec: WebSocketConnectionLogSpec,
trace_id: String,
remote_addr: SocketAddr,
path: String,
route_class: String,
user_id: String,
api_key_id: String,
started_at: std::time::Instant,
}
impl WebSocketConnectionLog {
pub(crate) fn new(context: &WebSocketRequestContext, spec: WebSocketConnectionLogSpec) -> Self {
let auth_context = context.decision.auth_context.as_ref();
Self {
spec,
trace_id: context.trace_id.clone(),
remote_addr: context.remote_addr,
path: context.uri.path().to_string(),
route_class: context
.decision
.route_class
.as_deref()
.unwrap_or("ai_public")
.to_string(),
user_id: auth_context
.map(|auth_context| auth_context.user_id.clone())
.unwrap_or_else(|| "-".to_string()),
api_key_id: auth_context
.map(|auth_context| auth_context.api_key_id.clone())
.unwrap_or_else(|| "-".to_string()),
started_at: std::time::Instant::now(),
}
}
pub(crate) fn log_opened(&self) {
info!(
event_name = self.spec.opened_event_name,
log_type = "access",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
status = "upgraded",
status_code = 101u16,
trace_id = %self.trace_id,
remote_addr = %self.remote_addr,
method = "GET",
path = %self.path,
user_id = %self.user_id,
api_key_id = %self.api_key_id,
route_class = %self.route_class,
execution_path = self.spec.execution_path,
provider_type = self.spec.provider_type,
message = self.spec.opened_message,
);
}
}
impl Drop for WebSocketConnectionLog {
fn drop(&mut self) {
info!(
event_name = self.spec.closed_event_name,
log_type = "access",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
status = "closed",
status_code = 101u16,
trace_id = %self.trace_id,
remote_addr = %self.remote_addr,
method = "GET",
path = %self.path,
user_id = %self.user_id,
api_key_id = %self.api_key_id,
route_class = %self.route_class,
execution_path = self.spec.execution_path,
provider_type = self.spec.provider_type,
elapsed_ms = self.started_at.elapsed().as_millis() as u64,
message = self.spec.closed_message,
);
}
}
#[cfg(test)]
mod tests {
use axum::http::header::{
AUTHORIZATION, CONNECTION, COOKIE, HOST, ORIGIN, SEC_WEBSOCKET_KEY, UPGRADE, USER_AGENT,
};
use axum::http::{HeaderMap, HeaderValue, Uri};
use super::{
websocket_credential_carrier_is_allowed, websocket_planning_headers, websocket_planning_uri,
};
use crate::control::GatewayCredentialCarrier;
#[test]
fn planning_uri_removes_query_credentials_without_losing_safe_parameters() {
let uri: Uri = "/v1/responses?key=downstream-secret&client_hint=a%20b&token=also-secret"
.parse()
.expect("request URI should parse");
let sanitized = websocket_planning_uri(&uri);
assert_eq!(sanitized.path(), "/v1/responses");
assert_eq!(sanitized.query(), Some("client_hint=a+b"));
assert!(!sanitized.to_string().contains("downstream-secret"));
assert!(!sanitized.to_string().contains("also-secret"));
}
#[test]
fn planning_uri_leaves_an_uncredentialed_query_byte_for_byte_unchanged() {
let uri: Uri = "/v1/responses?client_hint=a%20b&empty="
.parse()
.expect("request URI should parse");
assert_eq!(websocket_planning_uri(&uri), uri);
}
#[test]
fn planning_headers_drop_client_auth_cookie_and_websocket_transport_state() {
let mut headers = HeaderMap::new();
headers.insert(
AUTHORIZATION,
HeaderValue::from_static("Bearer client-secret"),
);
headers.insert(COOKIE, HeaderValue::from_static("session=client-secret"));
headers.insert("x-api-key", HeaderValue::from_static("client-secret"));
headers.insert(HOST, HeaderValue::from_static("gateway.example"));
headers.insert(
CONNECTION,
HeaderValue::from_static("keep-alive, Upgrade, x-connection-secret"),
);
headers.insert(UPGRADE, HeaderValue::from_static("websocket"));
headers.insert(SEC_WEBSOCKET_KEY, HeaderValue::from_static("handshake-key"));
headers.insert(
"sec-websocket-future-field",
HeaderValue::from_static("future-handshake-value"),
);
headers.insert(
"x-connection-secret",
HeaderValue::from_static("connection-secret"),
);
headers.insert(ORIGIN, HeaderValue::from_static("https://client.example"));
headers.insert(USER_AGENT, HeaderValue::from_static("codex-cli/test"));
headers.insert("x-client-hint", HeaderValue::from_static("safe"));
let sanitized = websocket_planning_headers(headers);
for name in [
AUTHORIZATION.as_str(),
COOKIE.as_str(),
"x-api-key",
HOST.as_str(),
CONNECTION.as_str(),
UPGRADE.as_str(),
SEC_WEBSOCKET_KEY.as_str(),
"sec-websocket-future-field",
"x-connection-secret",
] {
assert!(sanitized.get(name).is_none(), "{name} must not survive");
}
assert_eq!(
sanitized.get(ORIGIN),
Some(&HeaderValue::from_static("https://client.example"))
);
assert_eq!(
sanitized.get(USER_AGENT),
Some(&HeaderValue::from_static("codex-cli/test"))
);
assert_eq!(
sanitized.get("x-client-hint"),
Some(&HeaderValue::from_static("safe"))
);
}
#[test]
fn websocket_auth_requires_an_explicit_credential_instead_of_cookie_only() {
assert!(!websocket_credential_carrier_is_allowed(Some(
GatewayCredentialCarrier::CookieHeader
)));
for carrier in [
None,
Some(GatewayCredentialCarrier::AuthorizationBearer),
Some(GatewayCredentialCarrier::XApiKey),
Some(GatewayCredentialCarrier::ApiKey),
Some(GatewayCredentialCarrier::XGoogApiKey),
Some(GatewayCredentialCarrier::QueryKey),
] {
assert!(websocket_credential_carrier_is_allowed(carrier));
}
}
}
@@ -0,0 +1,12 @@
//! Shared infrastructure for public AI WebSocket bridges.
//!
//! Protocol adapters live below [`responses`]. This layer deliberately owns
//! only transport concerns that are common to future adapters: authenticated
//! upgrade admission, connection limits, upstream handshakes, and frame
//! conversion. It does not interpret provider events or make routing
//! decisions.
pub(crate) mod ingress;
pub(crate) mod responses;
pub(crate) mod session;
pub(crate) mod transport;
@@ -0,0 +1,234 @@
//! Provider-specific hooks for the standard Responses WebSocket session.
use async_trait::async_trait;
use serde_json::Value;
use super::adapters::CODEX_RESPONSES_WEBSOCKET_ADAPTER;
use crate::ai_serving::AiExecutionDecision;
use crate::handlers::proxy::websocket::transport::UpstreamWebSocketErrorCodes;
use crate::orchestration::ResponsesWebSocketAdapter;
use crate::AppState;
#[derive(Debug, Clone, Copy)]
pub(super) struct ResponsesWebSocketDrainDirective {
pub(super) error_code: &'static str,
/// The terminal upstream event may be replayed only when the session has
/// not exposed any standard Responses event to the client.
pub(super) retry_current_turn: bool,
/// When present, the exhausted provider key remains excluded from later
/// turns on this client socket until the upstream's reported reset time.
pub(super) retry_exclusion_until_unix_secs: Option<u64>,
}
/// Provider-specific observation produced while relaying an upstream frame.
/// The session can make the retry/drain decision synchronously, while the
/// optional persistence sink runs outside the frame-forwarding path.
#[derive(Debug, Clone)]
pub(super) struct ResponsesWebSocketAdapterObservation {
pub(super) drain: Option<ResponsesWebSocketDrainDirective>,
pub(super) quota_metadata: Option<Value>,
}
/// Provider identity used by the shared session's temporary exclusion table.
/// The session does not need to know how a provider derives its account id.
#[derive(Debug, Clone, Default)]
pub(super) struct ResponsesWebSocketExclusionIdentity {
pub(super) account_id: Option<String>,
}
/// Whether receiving an upstream event still leaves the active client turn
/// safe to replay on a freshly bound upstream. The shared session keeps the
/// conservative default; provider adapters may explicitly whitelist their
/// documented, pre-response advisory events.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) enum ResponsesWebSocketRebindSafety {
Safe,
Unsafe { reason: &'static str },
}
/// How an upstream text frame crosses the public Responses WebSocket boundary.
///
/// The normal path is deliberately byte-opaque: callers forward the parsed
/// frame's original text without rebuilding it from a gateway-owned schema.
/// Codex is the only adapter that may peel its documented private batch
/// envelope. Even then, the retained events are borrowed whole so unknown
/// `response.*` event types and unknown fields survive unchanged.
#[derive(Debug, Clone, PartialEq)]
pub(super) enum ResponsesWebSocketRelayDirective<'a> {
/// Forward the provider frame's original text exactly as received.
ForwardOriginal,
/// The provider frame was a private batch envelope. Forward each retained
/// event in document order by serializing the complete borrowed value.
ForwardEvents(Vec<&'a Value>),
/// The entire frame was an explicitly recognized provider-private
/// envelope and therefore has no public event to relay.
SuppressProviderPrivate,
}
/// Boundary between the standard Responses protocol engine and provider
/// behavior. Adapters receive already-planned provider requests; they never
/// own public WebSocket parsing, turn accounting, or model scheduling.
#[async_trait]
pub(super) trait ResponsesWebSocketProtocolAdapter: Send + Sync {
fn kind(&self) -> ResponsesWebSocketAdapter;
fn upstream_errors(&self) -> UpstreamWebSocketErrorCodes;
/// Adds provider-specific metadata to an otherwise standard Responses
/// stream report. The event payload is never rewritten for the client.
fn decorate_turn_report_context(&self, report_context: &mut Option<Value>, event: &Value);
/// Whether this adapter needs the shared session to parse each upstream
/// text event before normal turn accounting runs.
fn observes_upstream_events(&self) -> bool;
/// Classifies whether a received upstream event can be followed by a
/// transparent quota-driven rebind. An adapter must return `Safe` only
/// for events that neither create public Responses state nor make a replay
/// observably ambiguous to the client.
fn rebind_safety_for_upstream_event(&self, event: &Value) -> ResponsesWebSocketRebindSafety;
/// Selects the public relay shape without projecting a provider event
/// through an Aether-owned field or event-type allowlist.
fn relay_directive_for_upstream_event<'a>(
&self,
_event: &'a Value,
) -> ResponsesWebSocketRelayDirective<'a> {
ResponsesWebSocketRelayDirective::ForwardOriginal
}
/// Lets an adapter classify provider-only events. Returning a directive
/// asks the shared session to drain after the active standard response.
fn observe_upstream_event(&self, event: &Value)
-> Option<ResponsesWebSocketAdapterObservation>;
fn exhaustion_exclusion_identity(
&self,
_decision: &AiExecutionDecision,
) -> Option<ResponsesWebSocketExclusionIdentity> {
None
}
/// Persists an adapter observation outside the frame-forwarding path.
async fn persist_upstream_observation(
&self,
state: &AppState,
trace_id: &str,
report_context: Option<&Value>,
observation: ResponsesWebSocketAdapterObservation,
);
}
pub(super) fn resolve_responses_websocket_adapter(
kind: ResponsesWebSocketAdapter,
) -> &'static dyn ResponsesWebSocketProtocolAdapter {
match kind {
ResponsesWebSocketAdapter::Standard => &STANDARD_RESPONSES_WEBSOCKET_ADAPTER,
ResponsesWebSocketAdapter::Codex => &CODEX_RESPONSES_WEBSOCKET_ADAPTER,
}
}
struct StandardResponsesWebSocketAdapter;
const STANDARD_UPSTREAM_WEBSOCKET_ERRORS: UpstreamWebSocketErrorCodes =
UpstreamWebSocketErrorCodes {
upstream_url_missing: "responses_upstream_url_missing",
upstream_url_invalid: "responses_upstream_url_invalid",
headers_invalid: "responses_websocket_headers_invalid",
client_build_failed: "responses_websocket_client_build_failed",
proxy_invalid: "responses_websocket_proxy_invalid",
tunnel_proxy_unsupported: "responses_websocket_tunnel_proxy_unsupported",
handshake_failed: "responses_websocket_handshake_failed",
upgrade_rejected: "responses_websocket_upgrade_rejected",
upgrade_failed: "responses_websocket_upgrade_failed",
};
static STANDARD_RESPONSES_WEBSOCKET_ADAPTER: StandardResponsesWebSocketAdapter =
StandardResponsesWebSocketAdapter;
#[async_trait]
impl ResponsesWebSocketProtocolAdapter for StandardResponsesWebSocketAdapter {
fn kind(&self) -> ResponsesWebSocketAdapter {
ResponsesWebSocketAdapter::Standard
}
fn upstream_errors(&self) -> UpstreamWebSocketErrorCodes {
STANDARD_UPSTREAM_WEBSOCKET_ERRORS
}
fn decorate_turn_report_context(&self, _report_context: &mut Option<Value>, _event: &Value) {}
fn observes_upstream_events(&self) -> bool {
false
}
fn rebind_safety_for_upstream_event(&self, event: &Value) -> ResponsesWebSocketRebindSafety {
let reason = if is_standard_responses_event(event) {
"standard_response_event"
} else {
"unrecognized_upstream_event"
};
ResponsesWebSocketRebindSafety::Unsafe { reason }
}
fn observe_upstream_event(
&self,
_event: &Value,
) -> Option<ResponsesWebSocketAdapterObservation> {
None
}
async fn persist_upstream_observation(
&self,
_state: &AppState,
_trace_id: &str,
_report_context: Option<&Value>,
_observation: ResponsesWebSocketAdapterObservation,
) {
}
}
pub(super) fn is_standard_responses_event(event: &Value) -> bool {
event
.get("type")
.and_then(Value::as_str)
.is_some_and(|event_type| event_type.starts_with("response."))
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::{
resolve_responses_websocket_adapter, ResponsesWebSocketProtocolAdapter,
ResponsesWebSocketRelayDirective,
};
use crate::orchestration::ResponsesWebSocketAdapter;
#[test]
fn standard_adapter_has_no_codex_extensions() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
assert_eq!(adapter.kind(), ResponsesWebSocketAdapter::Standard);
assert!(!adapter.observes_upstream_events());
assert_eq!(
adapter.upstream_errors().handshake_failed,
"responses_websocket_handshake_failed"
);
}
#[test]
fn standard_adapter_always_forwards_future_events_opaquely() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
let event = json!({
"type": "response.future_capability.delta",
"delta": {"future_shape": [1, {"nested": true}]},
"unknown_top_level": {"must": "survive"},
});
assert_eq!(
adapter.relay_directive_for_upstream_event(&event),
ResponsesWebSocketRelayDirective::ForwardOriginal
);
}
}
@@ -0,0 +1,472 @@
//! Codex-specific extensions for the standard Responses WebSocket session.
use async_trait::async_trait;
use serde_json::{Map, Value};
use super::super::adapter::{
is_standard_responses_event, ResponsesWebSocketAdapterObservation,
ResponsesWebSocketDrainDirective, ResponsesWebSocketExclusionIdentity,
ResponsesWebSocketProtocolAdapter, ResponsesWebSocketRebindSafety,
ResponsesWebSocketRelayDirective,
};
use crate::ai_serving::AiExecutionDecision;
use crate::clock::current_unix_secs;
use crate::handlers::proxy::websocket::transport::UpstreamWebSocketErrorCodes;
use crate::orchestration::{
codex_account_id_from_headers, codex_quota_exhaustion_reset_at,
sync_codex_websocket_quota_metadata, ResponsesWebSocketAdapter,
};
use crate::AppState;
const CODEX_WEBSOCKET_LOG_TARGET: &str = "aether_gateway::handlers::proxy::codex_ws";
const CODEX_WEBSOCKET_RATE_LIMITS_REPORT_CONTEXT_FIELD: &str = "codex_websocket_rate_limits";
const CODEX_UPSTREAM_WEBSOCKET_ERRORS: UpstreamWebSocketErrorCodes = UpstreamWebSocketErrorCodes {
upstream_url_missing: "codex_upstream_url_missing",
upstream_url_invalid: "codex_upstream_url_invalid",
headers_invalid: "codex_websocket_headers_invalid",
client_build_failed: "codex_websocket_client_build_failed",
proxy_invalid: "codex_websocket_proxy_invalid",
tunnel_proxy_unsupported: "codex_websocket_tunnel_proxy_unsupported",
handshake_failed: "codex_websocket_handshake_failed",
upgrade_rejected: "codex_websocket_upgrade_rejected",
upgrade_failed: "codex_websocket_upgrade_failed",
};
pub(crate) static CODEX_RESPONSES_WEBSOCKET_ADAPTER: CodexResponsesWebSocketAdapter =
CodexResponsesWebSocketAdapter;
pub(crate) struct CodexResponsesWebSocketAdapter;
#[async_trait]
impl ResponsesWebSocketProtocolAdapter for CodexResponsesWebSocketAdapter {
fn kind(&self) -> ResponsesWebSocketAdapter {
ResponsesWebSocketAdapter::Codex
}
fn upstream_errors(&self) -> UpstreamWebSocketErrorCodes {
CODEX_UPSTREAM_WEBSOCKET_ERRORS
}
fn decorate_turn_report_context(&self, report_context: &mut Option<Value>, event: &Value) {
let Some(rate_limits) = parse_codex_rate_limits(event) else {
return;
};
let context = report_context.get_or_insert_with(|| Value::Object(Map::new()));
let Some(context) = context.as_object_mut() else {
return;
};
context.insert(
CODEX_WEBSOCKET_RATE_LIMITS_REPORT_CONTEXT_FIELD.to_string(),
rate_limits,
);
}
fn observes_upstream_events(&self) -> bool {
true
}
fn rebind_safety_for_upstream_event(&self, event: &Value) -> ResponsesWebSocketRebindSafety {
let mut saw_event = false;
if event.get("type").and_then(Value::as_str).is_some() {
saw_event = true;
let safety = codex_direct_rebind_safety(event);
if matches!(safety, ResponsesWebSocketRebindSafety::Unsafe { .. }) {
return safety;
}
}
match event.get("chunks") {
Some(Value::Array(chunks)) => {
for chunk in chunks {
saw_event = true;
let safety = codex_direct_rebind_safety(chunk);
if matches!(safety, ResponsesWebSocketRebindSafety::Unsafe { .. }) {
return safety;
}
}
}
Some(_) => {
return ResponsesWebSocketRebindSafety::Unsafe {
reason: "unrecognized_upstream_event",
};
}
None => {}
}
if saw_event {
ResponsesWebSocketRebindSafety::Safe
} else {
ResponsesWebSocketRebindSafety::Unsafe {
reason: "unrecognized_upstream_event",
}
}
}
fn relay_directive_for_upstream_event<'a>(
&self,
event: &'a Value,
) -> ResponsesWebSocketRelayDirective<'a> {
codex_relay_directive(event)
}
fn observe_upstream_event(
&self,
event: &Value,
) -> Option<ResponsesWebSocketAdapterObservation> {
let rate_limits = parse_codex_rate_limits(event)?;
let exhausted =
aether_admin::provider::quota::codex_rate_limit_metadata_exhausted(&rate_limits);
let retry_exclusion_until_unix_secs =
codex_quota_exhaustion_reset_at(&rate_limits, current_unix_secs());
Some(ResponsesWebSocketAdapterObservation {
drain: exhausted.then_some(ResponsesWebSocketDrainDirective {
error_code: "codex_account_quota_exhausted",
retry_current_turn: true,
retry_exclusion_until_unix_secs,
}),
quota_metadata: Some(rate_limits),
})
}
fn exhaustion_exclusion_identity(
&self,
decision: &AiExecutionDecision,
) -> Option<ResponsesWebSocketExclusionIdentity> {
Some(ResponsesWebSocketExclusionIdentity {
account_id: codex_account_id_from_headers(&decision.provider_request_headers)
.map(str::to_string),
})
}
async fn persist_upstream_observation(
&self,
state: &AppState,
trace_id: &str,
report_context: Option<&Value>,
observation: ResponsesWebSocketAdapterObservation,
) {
let Some(rate_limits) = observation.quota_metadata else {
return;
};
if let Err(error) =
sync_codex_websocket_quota_metadata(state, report_context, rate_limits).await
{
tracing::warn!(
target: CODEX_WEBSOCKET_LOG_TARGET,
event_name = "codex_websocket_quota_sync_failed",
log_type = "ops",
transport = "websocket",
websocket = true,
trace_id = %trace_id,
error = ?error,
"gateway failed to persist Codex WebSocket quota metadata"
);
}
}
}
fn codex_direct_rebind_safety(event: &Value) -> ResponsesWebSocketRebindSafety {
let event_type = event
.get("type")
.and_then(Value::as_str)
.unwrap_or_default();
if matches!(event_type, "codex.rate_limits" | "codex.response.metadata") {
// Codex emits these as pre-response advisory metadata. They do
// not create a public `response.*` object, so a replacement
// upstream can safely emit its own current snapshot.
return ResponsesWebSocketRebindSafety::Safe;
}
if event_type == "error"
&& event.pointer("/error/type").and_then(Value::as_str) == Some("usage_limit_reached")
&& parse_codex_rate_limits(event).is_some()
{
// This terminal quota event has not been relayed yet. It can trigger
// one transparent attempt on another key as long as no earlier public
// response event made the logical turn unsafe. If replanning fails,
// the connection layer forwards this exact upstream error instead of
// manufacturing a gateway continuation error.
return ResponsesWebSocketRebindSafety::Safe;
}
let reason = if is_standard_responses_event(event) {
"standard_response_event"
} else {
"unrecognized_upstream_event"
};
ResponsesWebSocketRebindSafety::Unsafe { reason }
}
fn codex_relay_directive(event: &Value) -> ResponsesWebSocketRelayDirective<'_> {
match event.get("chunks") {
Some(Value::Array(chunks)) if is_explicit_codex_batch_envelope(event) => {
let public_events = chunks
.iter()
.filter(|chunk| !is_codex_private_leaf_event(chunk))
.collect::<Vec<_>>();
if public_events.is_empty() {
ResponsesWebSocketRelayDirective::SuppressProviderPrivate
} else {
ResponsesWebSocketRelayDirective::ForwardEvents(public_events)
}
}
// A malformed or future shape is not proven private. Preserve it
// opaquely rather than guessing at a provider schema.
Some(_) => ResponsesWebSocketRelayDirective::ForwardOriginal,
None if is_codex_private_leaf_event(event) => {
ResponsesWebSocketRelayDirective::SuppressProviderPrivate
}
None => ResponsesWebSocketRelayDirective::ForwardOriginal,
}
}
/// Recognizes only Codex's private batch container. A type-less object must
/// contain exactly `chunks`; unknown siblings could be future public protocol
/// data and therefore force opaque forwarding. A named Codex private root may
/// carry provider metadata alongside its chunks and is safe to peel.
fn is_explicit_codex_batch_envelope(event: &Value) -> bool {
if is_codex_private_event_type(event) {
return true;
}
event.as_object().is_some_and(|object| {
object.len() == 1
&& object.contains_key("chunks")
&& event.get("type").and_then(Value::as_str).is_none()
})
}
fn is_codex_private_leaf_event(event: &Value) -> bool {
is_codex_private_event_type(event) && event.get("chunks").is_none()
}
fn is_codex_private_event_type(event: &Value) -> bool {
matches!(
event.get("type").and_then(Value::as_str),
Some("codex.rate_limits" | "codex.response.metadata")
)
}
fn parse_codex_rate_limits(event: &Value) -> Option<Value> {
aether_admin::provider::quota::parse_codex_websocket_rate_limits_response(
event,
current_unix_secs(),
)
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::{
CodexResponsesWebSocketAdapter, ResponsesWebSocketProtocolAdapter,
ResponsesWebSocketRebindSafety, ResponsesWebSocketRelayDirective,
};
#[test]
fn codex_rate_limit_chunk_is_kept_for_the_terminal_report() {
let adapter = CodexResponsesWebSocketAdapter;
assert!(adapter.observes_upstream_events());
let mut context = Some(json!({"key_id": "codex-key"}));
adapter.decorate_turn_report_context(
&mut context,
&json!({
"chunks": [{
"type": "codex.rate_limits",
"plan_type": "free",
"rate_limits": {
"allowed": true,
"limit_reached": false,
"primary": {
"used_percent": 91,
"window_minutes": 43200,
"reset_after_seconds": 2590791
}
}
}]
}),
);
assert_eq!(
context.as_ref().and_then(
|context| context.pointer("/codex_websocket_rate_limits/primary_used_percent")
),
Some(&json!(91.0))
);
}
#[test]
fn usage_limit_error_is_kept_for_the_terminal_report() {
let adapter = CodexResponsesWebSocketAdapter;
let mut context = Some(json!({"key_id": "codex-key"}));
adapter.decorate_turn_report_context(
&mut context,
&json!({
"type": "error",
"error": {
"type": "usage_limit_reached",
"plan_type": "free",
"resets_at": 1_787_274_385u64,
},
"status_code": 429,
"headers": {
"X-Codex-Primary-Used-Percent": "100",
"X-Codex-Primary-Reset-At": "1787274385",
},
}),
);
assert_eq!(
context
.as_ref()
.and_then(|context| context.pointer("/codex_websocket_rate_limits/allowed")),
Some(&json!(false))
);
assert_eq!(
context.as_ref().and_then(|context| {
context.pointer("/codex_websocket_rate_limits/primary_used_percent")
}),
Some(&json!(100.0))
);
}
#[test]
fn only_known_codex_pre_response_signals_are_safe_to_rebind() {
let adapter = CodexResponsesWebSocketAdapter;
assert_eq!(
adapter.rebind_safety_for_upstream_event(&json!({
"type": "codex.rate_limits",
"rate_limits": {"allowed": true}
})),
ResponsesWebSocketRebindSafety::Safe
);
assert_eq!(
adapter.rebind_safety_for_upstream_event(&json!({
"type": "codex.response.metadata"
})),
ResponsesWebSocketRebindSafety::Safe
);
assert_eq!(
adapter.rebind_safety_for_upstream_event(&json!({
"chunks": [
{"type": "codex.rate_limits", "rate_limits": {"allowed": true}},
{"type": "codex.response.metadata"}
]
})),
ResponsesWebSocketRebindSafety::Safe
);
assert_eq!(
adapter.rebind_safety_for_upstream_event(&json!({
"type": "response.created"
})),
ResponsesWebSocketRebindSafety::Unsafe {
reason: "standard_response_event"
}
);
assert_eq!(
adapter.rebind_safety_for_upstream_event(&json!({
"type": "codex.unknown"
})),
ResponsesWebSocketRebindSafety::Unsafe {
reason: "unrecognized_upstream_event"
}
);
assert_eq!(
adapter.rebind_safety_for_upstream_event(&json!({
"type": "error",
"error": {
"type": "usage_limit_reached",
"plan_type": "plus",
"resets_in_seconds": 3_600
},
"status_code": 429
})),
ResponsesWebSocketRebindSafety::Safe
);
assert_eq!(
adapter.rebind_safety_for_upstream_event(&json!({
"type": "error",
"error": {"type": "usage_limit_reached"}
})),
ResponsesWebSocketRebindSafety::Unsafe {
reason: "unrecognized_upstream_event"
}
);
assert_eq!(
adapter.rebind_safety_for_upstream_event(&json!({
"type": "response.future_capability.delta",
"chunks": [{"type": "codex.rate_limits"}]
})),
ResponsesWebSocketRebindSafety::Unsafe {
reason: "standard_response_event"
}
);
}
#[test]
fn codex_suppresses_only_explicit_private_events_and_envelopes() {
let adapter = CodexResponsesWebSocketAdapter;
for event in [
json!({"type": "codex.rate_limits", "rate_limits": {"allowed": true}}),
json!({"type": "codex.response.metadata", "account_hint": "private"}),
json!({"chunks": [
{"type": "codex.rate_limits"},
{"type": "codex.response.metadata"}
]}),
] {
assert_eq!(
adapter.relay_directive_for_upstream_event(&event),
ResponsesWebSocketRelayDirective::SuppressProviderPrivate
);
}
for event in [
json!({"type": "error", "error": {"type": "usage_limit_reached"}}),
json!({"type": "codex.future_private_maybe", "future": true}),
json!({"chunks": [], "future_envelope_field": {"must": "survive"}}),
json!({"type": "response.future.done", "future_capability": true}),
] {
assert_eq!(
adapter.relay_directive_for_upstream_event(&event),
ResponsesWebSocketRelayDirective::ForwardOriginal
);
}
}
#[test]
fn mixed_codex_batch_forwards_whole_non_private_events_in_order() {
let adapter = CodexResponsesWebSocketAdapter;
let event = json!({
"chunks": [
{
"type": "response.created",
"response": {"id": "resp_future"},
"future_created_field": {"opaque": true}
},
{"type": "codex.rate_limits", "account_hint": "private"},
{
"type": "response.future_capability.delta",
"future_capability": {"nested": [1, 2, 3]},
"sequence_number": 2
},
{"provider_future_event": {"unknown": "must be forwarded"}},
{
"type": "error",
"error": {"type": "future_error", "future_detail": 7}
}
]
});
let ResponsesWebSocketRelayDirective::ForwardEvents(events) =
adapter.relay_directive_for_upstream_event(&event)
else {
panic!("a mixed private envelope must retain all non-private events");
};
assert_eq!(events.len(), 4);
assert_eq!(events[0]["future_created_field"], json!({"opaque": true}));
assert_eq!(events[1]["future_capability"], json!({"nested": [1, 2, 3]}));
assert_eq!(
events[2]["provider_future_event"],
json!({"unknown": "must be forwarded"})
);
assert_eq!(events[3]["error"]["future_detail"], json!(7));
}
}
@@ -0,0 +1,5 @@
//! Provider-specific Responses WebSocket adapters.
mod codex;
pub(super) use codex::CODEX_RESPONSES_WEBSOCKET_ADAPTER;
@@ -0,0 +1,80 @@
//! Per-turn resource admission for the Responses WebSocket bridge.
//!
//! A WebSocket connection may live for a long time, but each `response.create`
//! is still one active upstream execution. Keep the resource leases attached
//! to the turn instead of the socket so idle connections do not consume
//! upstream capacity.
use std::time::Instant;
use aether_contracts::ExecutionPlan;
use crate::execution_runtime::acquire_upstream_execution_gate;
use crate::provider_pool_demand::{
acquire_provider_pool_in_flight_guard, ProviderPoolInFlightGuard,
};
use crate::upstream_admission::UpstreamTargetAdmissionPermit;
use crate::{AppState, GatewayError};
pub(super) struct ResponsesWebSocketTurnAdmission {
upstream_execution: Option<aether_runtime::ConcurrencyPermit>,
upstream_target: Option<UpstreamTargetAdmissionPermit>,
provider_pool: Option<ProviderPoolInFlightGuard>,
acquired_at: Instant,
}
impl ResponsesWebSocketTurnAdmission {
pub(super) async fn acquire(
state: &AppState,
plan: &ExecutionPlan,
trace_id: &str,
) -> Result<Self, GatewayError> {
let upstream_execution = acquire_upstream_execution_gate(state, trace_id).await?;
let upstream_target = match state
.upstream_target_admission
.acquire(plan, trace_id)
.await
{
Ok(permit) => permit,
Err(error) => {
drop(upstream_execution);
return Err(error);
}
};
let provider_pool = acquire_provider_pool_in_flight_guard(
state.runtime_state.clone(),
&plan.provider_id,
&plan.request_id,
plan.candidate_id.as_deref(),
&plan.key_id,
)
.await;
Ok(Self {
upstream_execution,
upstream_target,
provider_pool,
acquired_at: Instant::now(),
})
}
/// Release the distributed provider token before the turn's persistence
/// work. The remaining permits are local RAII guards and are dropped with
/// this value.
pub(super) async fn release(mut self) {
if let Some(provider_pool) = self.provider_pool.take() {
provider_pool.release().await;
}
drop(self.upstream_target.take());
drop(self.upstream_execution.take());
}
}
impl Drop for ResponsesWebSocketTurnAdmission {
fn drop(&mut self) {
crate::stage_metrics::observe_gateway_stage_ms(
"websocket_turn_admission_held",
self.acquired_at.elapsed().as_millis() as u64,
);
}
}
@@ -0,0 +1,560 @@
//! Identity of the physical upstream connection backing a Responses session.
//!
//! A Responses continuation carries state that lives on one provider socket.
//! Comparing only the selected key is therefore not sufficient: transport
//! settings, stable account headers, credentials, and the protocol adapter can
//! all change the connection that would receive the next event. Ordinary Codex
//! OAuth access-token refreshes retain the credential generation and therefore
//! do not unnecessarily replace an already-upgraded socket.
use std::collections::{BTreeMap, BTreeSet};
use std::fmt;
use aether_contracts::{ProxySnapshot, ResolvedTransportProfile};
use sha2::{Digest, Sha256};
use super::adapter::ResponsesWebSocketProtocolAdapter;
use crate::ai_serving::AiExecutionDecision;
use crate::handlers::proxy::websocket::transport::{
websocket_handshake_headers, websocket_upstream_url,
};
use crate::orchestration::ResponsesWebSocketAdapter;
/// Stable, comparable identity for the actual WebSocket connection target.
///
/// The identity deliberately owns the normalized handshake values rather than
/// retaining a reference to the planner decision. A later re-plan can then
/// be compared without accidentally ignoring a field that changes the
/// physical connection.
#[derive(Clone, PartialEq)]
pub(super) struct UpstreamBindingIdentity {
adapter_kind: ResponsesWebSocketAdapter,
provider_id: Option<String>,
endpoint_id: Option<String>,
key_id: Option<String>,
upstream_url: String,
handshake_headers: BTreeMap<String, String>,
/// One-way identity for the credential generation used by this socket.
///
/// A provider key id identifies a catalog row, not the secret currently
/// stored in that row. Codex decisions carry a server-owned credential
/// generation which is stable across access-token refreshes but rotates
/// when the account/static/refresh credential is replaced. Other
/// decisions conservatively fingerprint the effective authentication
/// handshake values.
credential_fingerprint: [u8; 32],
proxy: Option<ProxySnapshot>,
transport_profile: Option<ResolvedTransportProfile>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) enum UpstreamBindingIdentityError {
MissingUpstreamUrl,
InvalidUpstreamUrl,
InvalidHandshakeHeaders,
}
impl UpstreamBindingIdentity {
/// Builds an identity from the same normalized URL and headers used by
/// the WebSocket transport client.
pub(super) fn from_decision(
adapter: &'static dyn ResponsesWebSocketProtocolAdapter,
decision: &AiExecutionDecision,
) -> Result<Self, UpstreamBindingIdentityError> {
let raw_url = decision
.upstream_url
.as_deref()
.filter(|value| !value.trim().is_empty())
.ok_or(UpstreamBindingIdentityError::MissingUpstreamUrl)?;
let upstream_url = websocket_upstream_url(raw_url, "invalid")
.map_err(|_| UpstreamBindingIdentityError::InvalidUpstreamUrl)?
.to_string();
let headers = websocket_handshake_headers(&decision.provider_request_headers, "invalid")
.map_err(|_| UpstreamBindingIdentityError::InvalidHandshakeHeaders)?;
let authentication_header_names = authentication_header_names(decision);
let mut handshake_headers = BTreeMap::new();
let mut authentication_headers = BTreeMap::new();
for (name, value) in &headers {
let name = name.as_str().to_ascii_lowercase();
let value = value
.to_str()
.map_err(|_| UpstreamBindingIdentityError::InvalidHandshakeHeaders)?;
if authentication_header_names.contains(name.as_str()) {
authentication_headers.insert(name, value.to_string());
} else {
handshake_headers.insert(name, value.to_string());
}
}
let credential_fingerprint =
credential_binding_fingerprint(decision, &authentication_headers);
Ok(Self {
adapter_kind: adapter.kind(),
provider_id: decision.provider_id.clone(),
endpoint_id: decision.endpoint_id.clone(),
key_id: decision.key_id.clone(),
upstream_url,
handshake_headers,
credential_fingerprint,
proxy: effective_proxy_snapshot(decision.proxy.as_ref()),
transport_profile: decision.transport_profile.clone(),
})
}
}
/// Header names that carry credentials in the provider handshake. The
/// planner's explicit `auth_header` extends this list for provider-specific
/// schemes; unknown headers remain part of the stable handshake identity.
fn authentication_header_names(decision: &AiExecutionDecision) -> BTreeSet<String> {
let mut names = BTreeSet::from([
"authorization".to_string(),
"proxy-authorization".to_string(),
"x-api-key".to_string(),
"api-key".to_string(),
"x-goog-api-key".to_string(),
"x-azure-api-key".to_string(),
]);
if let Some(name) = decision
.auth_header
.as_deref()
.map(str::trim)
.filter(|name| !name.is_empty())
{
names.insert(name.to_ascii_lowercase());
}
names
}
fn fingerprint_headers(headers: &BTreeMap<String, String>) -> [u8; 32] {
let mut hasher = Sha256::new();
hasher.update(b"aether-responses-websocket-auth-headers-v1");
for (name, value) in headers {
hasher.update((name.len() as u64).to_be_bytes());
hasher.update(name.as_bytes());
hasher.update((value.len() as u64).to_be_bytes());
hasher.update(value.as_bytes());
}
hasher.finalize().into()
}
/// Returns the non-secret credential identity represented by a planner
/// decision. The generation is emitted by Aether's trusted Codex planner from
/// provider-key metadata; it is not sourced from the downstream request.
fn credential_binding_fingerprint(
decision: &AiExecutionDecision,
authentication_headers: &BTreeMap<String, String>,
) -> [u8; 32] {
if decision
.provider_type
.as_deref()
.is_some_and(|provider_type| provider_type.trim().eq_ignore_ascii_case("codex"))
{
if let Some(generation) = decision
.report_context
.as_ref()
.and_then(|context| context.get("codex_credential_generation"))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|generation| !generation.is_empty())
{
let mut hasher = Sha256::new();
hasher.update(b"aether-responses-websocket-codex-credential-generation-v1");
hasher.update((generation.len() as u64).to_be_bytes());
hasher.update(generation.as_bytes());
// Only a planner-owned Codex bearer access token is expected to
// rotate without changing credential generation. Compare the
// effective handshake value with the decision's original auth
// value: auth-config/routing/header overrides change only the
// former and therefore must force a rebind.
let stable_authentication_headers = authentication_headers
.iter()
.filter(|(name, value)| {
!is_planner_owned_codex_bearer(decision, name.as_str(), value.as_str())
})
.map(|(name, value)| (name.clone(), value.clone()))
.collect::<BTreeMap<_, _>>();
hasher.update(fingerprint_headers(&stable_authentication_headers));
return hasher.finalize().into();
}
}
// Fail closed when no trusted generation is available. Rebinding after an
// access-token change is preferable to sending a continuation over a
// socket authenticated with a credential that may have been replaced.
fingerprint_headers(authentication_headers)
}
fn is_planner_owned_codex_bearer(
decision: &AiExecutionDecision,
name: &str,
effective_value: &str,
) -> bool {
name.eq_ignore_ascii_case("authorization")
&& decision
.auth_header
.as_deref()
.is_some_and(|header| header.eq_ignore_ascii_case(name))
&& decision.auth_value.as_deref() == Some(effective_value)
&& effective_value
.get(.."bearer ".len())
.is_some_and(|scheme| scheme.eq_ignore_ascii_case("bearer "))
}
/// Normalize only values that are provably direct transport. Keep node/tunnel
/// fields even though the current WebSocket builder rejects those proxies: a
/// re-plan must not accidentally reuse an already-bound direct socket for a
/// decision that selected a different proxy topology.
fn effective_proxy_snapshot(proxy: Option<&ProxySnapshot>) -> Option<ProxySnapshot> {
let proxy = proxy?;
if proxy.enabled == Some(false) {
return None;
}
let mut normalized = proxy.clone();
normalized.url = normalized
.url
.take()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty());
normalized.mode = normalized
.mode
.take()
.map(|value| value.trim().to_ascii_lowercase())
.filter(|value| !value.is_empty());
normalized.node_id = normalized
.node_id
.take()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty());
normalized.label = normalized
.label
.take()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty());
let has_effective_proxy = normalized.url.is_some()
|| normalized.node_id.is_some()
|| normalized.mode.is_some()
|| normalized.extra.is_some();
has_effective_proxy.then_some(normalized)
}
impl fmt::Debug for UpstreamBindingIdentity {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("UpstreamBindingIdentity")
.field("adapter_kind", &self.adapter_kind)
.field("provider_id", &self.provider_id)
.field("endpoint_id", &self.endpoint_id)
.field("key_id", &self.key_id)
.field("upstream_url", &self.upstream_url)
.field(
"handshake_header_names",
&self.handshake_headers.keys().collect::<Vec<_>>(),
)
.field("proxy_configured", &self.proxy.is_some())
.field(
"transport_profile_id",
&self
.transport_profile
.as_ref()
.map(|profile| profile.profile_id.as_str()),
)
.finish()
}
}
#[cfg(test)]
mod tests {
use std::collections::BTreeMap;
use serde_json::json;
use super::{UpstreamBindingIdentity, UpstreamBindingIdentityError};
use crate::ai_serving::AiExecutionDecision;
use crate::handlers::proxy::websocket::responses::adapter::resolve_responses_websocket_adapter;
use crate::orchestration::ResponsesWebSocketAdapter;
fn decision() -> AiExecutionDecision {
AiExecutionDecision {
action: "execute".to_string(),
decision_kind: None,
execution_strategy: None,
conversion_mode: None,
request_id: Some("request-1".to_string()),
candidate_id: Some("candidate-1".to_string()),
provider_name: Some("provider".to_string()),
provider_type: Some("openai".to_string()),
provider_id: Some("provider-1".to_string()),
endpoint_id: Some("endpoint-1".to_string()),
key_id: Some("key-1".to_string()),
upstream_base_url: Some("https://api.example.test".to_string()),
upstream_url: Some("https://api.example.test/v1/responses".to_string()),
provider_request_method: Some("POST".to_string()),
auth_header: Some("authorization".to_string()),
auth_value: Some("Bearer secret".to_string()),
provider_api_format: Some("openai:responses".to_string()),
client_api_format: Some("openai:responses".to_string()),
provider_contract: None,
client_contract: None,
model_name: Some("gpt-5.6-sol".to_string()),
mapped_model: None,
prompt_cache_key: None,
extra_headers: BTreeMap::new(),
provider_request_headers: BTreeMap::from([
("Authorization".to_string(), "Bearer secret".to_string()),
("X-Client".to_string(), "aether".to_string()),
("Connection".to_string(), "keep-alive".to_string()),
]),
provider_request_body: Some(json!({"model": "gpt-5.6-sol"})),
provider_request_body_base64: None,
content_type: Some("application/json".to_string()),
content_encoding: None,
request_gzip: None,
proxy: None,
transport_profile: None,
timeouts: None,
upstream_is_stream: true,
report_kind: None,
report_context: None,
auth_context: None,
}
}
#[test]
fn identity_normalizes_url_and_hop_by_hop_headers() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
let identity = UpstreamBindingIdentity::from_decision(adapter, &decision()).unwrap();
assert_eq!(identity.upstream_url, "wss://api.example.test/v1/responses");
assert_eq!(
identity.handshake_headers,
BTreeMap::from([("x-client".to_string(), "aether".to_string())])
);
}
#[test]
fn identity_changes_when_physical_binding_changes() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
let base = decision();
let identity = UpstreamBindingIdentity::from_decision(adapter, &base).unwrap();
let codex_adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex);
assert_ne!(
identity,
UpstreamBindingIdentity::from_decision(codex_adapter, &base).unwrap()
);
for mutate in [
|decision: &mut AiExecutionDecision| {
decision.key_id = Some("key-2".to_string());
},
|decision: &mut AiExecutionDecision| {
decision.upstream_url = Some("https://other.example.test/v1/responses".to_string());
},
|decision: &mut AiExecutionDecision| {
decision
.provider_request_headers
.insert("X-Client".to_string(), "other".to_string());
},
|decision: &mut AiExecutionDecision| {
decision.proxy = Some(aether_contracts::ProxySnapshot {
enabled: Some(true),
url: Some("http://proxy.example.test:8080".to_string()),
..Default::default()
});
},
|decision: &mut AiExecutionDecision| {
decision.transport_profile = Some(aether_contracts::ResolvedTransportProfile {
profile_id: "chrome136".to_string(),
..Default::default()
});
},
] {
let mut changed = base.clone();
mutate(&mut changed);
let changed_identity =
UpstreamBindingIdentity::from_decision(adapter, &changed).unwrap();
assert_ne!(identity, changed_identity);
}
let mut static_secret_rotated = base.clone();
static_secret_rotated
.provider_request_headers
.insert("Authorization".to_string(), "Bearer rotated".to_string());
assert_ne!(
identity,
UpstreamBindingIdentity::from_decision(adapter, &static_secret_rotated).unwrap()
);
}
#[test]
fn stable_key_identity_rejects_custom_static_auth_value_rotation() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
let mut base = decision();
base.auth_header = Some("X-Provider-Token".to_string());
base.provider_request_headers.remove("Authorization");
base.provider_request_headers.insert(
"X-Provider-Token".to_string(),
"provider-token-1".to_string(),
);
let identity = UpstreamBindingIdentity::from_decision(adapter, &base).unwrap();
assert!(!identity.handshake_headers.contains_key("x-provider-token"));
let mut rotated = base;
rotated.provider_request_headers.insert(
"X-Provider-Token".to_string(),
"provider-token-2".to_string(),
);
assert_ne!(
identity,
UpstreamBindingIdentity::from_decision(adapter, &rotated).unwrap()
);
}
#[test]
fn codex_access_token_refresh_reuses_the_same_credential_generation() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex);
let mut first = decision();
first.provider_type = Some("codex".to_string());
first.report_context = Some(json!({
"codex_credential_generation": "credential-generation-1"
}));
let first_identity = UpstreamBindingIdentity::from_decision(adapter, &first).unwrap();
let mut access_token_refreshed = first;
access_token_refreshed.auth_value = Some("Bearer refreshed-access-token".to_string());
access_token_refreshed.provider_request_headers.insert(
"Authorization".to_string(),
"Bearer refreshed-access-token".to_string(),
);
assert_eq!(
first_identity,
UpstreamBindingIdentity::from_decision(adapter, &access_token_refreshed).unwrap()
);
}
#[test]
fn codex_authorization_override_changes_binding_with_the_same_generation() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex);
let mut first = decision();
first.provider_type = Some("codex".to_string());
first.report_context = Some(json!({
"codex_credential_generation": "credential-generation-1"
}));
let first_identity = UpstreamBindingIdentity::from_decision(adapter, &first).unwrap();
// The planner-owned auth value remains unchanged while an effective
// auth-config/header override replaces the actual handshake value.
first.provider_request_headers.insert(
"Authorization".to_string(),
"Bearer endpoint-override".to_string(),
);
assert_ne!(
first_identity,
UpstreamBindingIdentity::from_decision(adapter, &first).unwrap()
);
}
#[test]
fn codex_credential_replacement_changes_binding_for_the_same_key_id() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex);
let mut first = decision();
first.provider_type = Some("codex".to_string());
first.report_context = Some(json!({
"codex_credential_generation": "credential-generation-1"
}));
let first_identity = UpstreamBindingIdentity::from_decision(adapter, &first).unwrap();
let mut replaced = first;
replaced.provider_request_headers.insert(
"Authorization".to_string(),
"Bearer replacement-access-token".to_string(),
);
replaced.report_context = Some(json!({
"codex_credential_generation": "credential-generation-2"
}));
assert_ne!(
first_identity,
UpstreamBindingIdentity::from_decision(adapter, &replaced).unwrap()
);
}
#[test]
fn codex_custom_auth_rotation_changes_binding_with_the_same_generation() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex);
let mut first = decision();
first.provider_type = Some("codex".to_string());
first.auth_header = Some("X-Provider-Token".to_string());
first.provider_request_headers.insert(
"X-Provider-Token".to_string(),
"provider-token-1".to_string(),
);
first.report_context = Some(json!({
"codex_credential_generation": "credential-generation-1"
}));
let first_identity = UpstreamBindingIdentity::from_decision(adapter, &first).unwrap();
first.provider_request_headers.insert(
"X-Provider-Token".to_string(),
"provider-token-2".to_string(),
);
assert_ne!(
first_identity,
UpstreamBindingIdentity::from_decision(adapter, &first).unwrap()
);
}
#[test]
fn missing_codex_credential_generation_fails_closed_on_auth_rotation() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex);
let mut first = decision();
first.provider_type = Some("codex".to_string());
let first_identity = UpstreamBindingIdentity::from_decision(adapter, &first).unwrap();
first.provider_request_headers.insert(
"Authorization".to_string(),
"Bearer possibly-replaced-credential".to_string(),
);
assert_ne!(
first_identity,
UpstreamBindingIdentity::from_decision(adapter, &first).unwrap()
);
}
#[test]
fn disabled_proxy_is_equivalent_to_direct_transport() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
let direct = decision();
let direct_identity = UpstreamBindingIdentity::from_decision(adapter, &direct).unwrap();
let mut explicitly_disabled = direct;
explicitly_disabled.proxy = Some(aether_contracts::ProxySnapshot {
enabled: Some(false),
url: Some("http://ignored.example.test:8080".to_string()),
..Default::default()
});
assert_eq!(
direct_identity,
UpstreamBindingIdentity::from_decision(adapter, &explicitly_disabled).unwrap()
);
}
#[test]
fn identity_rejects_missing_or_invalid_connection_fields() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
let mut missing = decision();
missing.upstream_url = None;
assert_eq!(
UpstreamBindingIdentity::from_decision(adapter, &missing),
Err(UpstreamBindingIdentityError::MissingUpstreamUrl)
);
let mut invalid = decision();
invalid.upstream_url = Some("file:///tmp/responses".to_string());
assert_eq!(
UpstreamBindingIdentity::from_decision(adapter, &invalid),
Err(UpstreamBindingIdentityError::InvalidUpstreamUrl)
);
}
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,623 @@
//! Connection-level Responses WebSocket FSM.
use std::time::Duration;
use axum::extract::ws::{Message as AxumWsMessage, WebSocket};
use futures_util::{SinkExt, StreamExt};
use serde_json::Value;
use wreq::ws::message::Message as WreqWsMessage;
use super::adapter::ResponsesWebSocketRelayDirective;
use super::client::{adapter_drain_ready, forward_client_message, RelayDisposition};
use super::frame::{encode_opaque_websocket_event, ParsedResponsesWebSocketFrame};
use super::lifecycle::{
await_pending_adapter_observation, finalize_active_turn, queue_turn_finalization,
settle_turn_finalization, spawn_bounded_adapter_observation, PreviousAttemptSettled,
};
use super::quota::{
detach_exhausted_upstream, is_usage_limit_error_event, mark_active_response_retry_unsafe,
observe_active_response_rebind_safety, retry_active_turn_after_quota_exhaustion,
};
use super::relay_policy::{
classify_quota_relay, fatal_relay_policy, FatalRelaySignal, QuotaRelayAction, QuotaRelayFacts,
};
use super::settlement::settle_signal_for_client_delivery_failure;
use super::state::BoundResponsesConnection;
use super::turn::{
ResponsesProviderAttempt, ResponsesWebSocketTurnObservation, ResponsesWebSocketTurnOutcome,
};
use super::upstream::{close_bound_upstream, receive_optional_upstream};
use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext;
use crate::handlers::proxy::websocket::session::{
wait_for_optional_deadline, CLOSE_INTERNAL_ERROR, CLOSE_TRY_AGAIN, WEBSOCKET_LOG_TRANSPORT,
};
use crate::handlers::proxy::websocket::transport::{
close_client_socket, send_client_message, send_gateway_error_with_status,
send_responses_websocket_error, upstream_message_to_client,
};
use crate::AppState;
const LOG_TARGET: &str = "aether_gateway::handlers::proxy::responses_ws";
/// 写客户端 socket 失败时记录的投递失败原因。刻意不说「客户端在终态前断开」:
/// 供应商的终态可能已经到达,只是最后一跳没送出去。
const CLIENT_DELIVERY_FAILED_REASON: &str =
"gateway could not relay the provider event to the client";
macro_rules! debug {
($($arg:tt)*) => {
tracing::debug!(target: LOG_TARGET, $($arg)*)
};
}
macro_rules! warn {
($($arg:tt)*) => {
tracing::warn!(target: LOG_TARGET, $($arg)*)
};
}
pub(super) async fn relay_bound_connection(
client_socket: &mut WebSocket,
bound: &mut BoundResponsesConnection,
state: &AppState,
context: &WebSocketRequestContext,
) {
loop {
let active_turn_deadline = bound.turn_state.attempt().map(|turn| turn.deadline());
tokio::select! {
_ = wait_for_optional_deadline(active_turn_deadline.map(|deadline| deadline.deadline)) => {
let Some(turn_deadline) = active_turn_deadline else {
continue;
};
warn!(
event_name = "responses_websocket_turn_timeout",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
timeout_phase = ?turn_deadline.phase,
timeout_ms = turn_deadline.timeout.as_millis() as u64,
"Responses WebSocket response did not reach its configured deadline"
);
finalize_active_turn(bound, state, turn_deadline.phase.outcome()).await;
send_gateway_error_with_status(
client_socket,
504,
turn_deadline.phase.error_code(),
turn_deadline.phase.client_message(),
).await;
close_bound_upstream(bound).await;
close_client_socket(
client_socket,
CLOSE_TRY_AGAIN,
turn_deadline.phase.error_code(),
).await;
break;
}
client_message = client_socket.next() => {
let Some(client_message) = client_message else {
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::client_disconnected(),
).await;
close_bound_upstream(bound).await;
break;
};
let Ok(client_message) = client_message else {
warn!(
event_name = "responses_websocket_client_receive_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
"client WebSocket receive failed"
);
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::client_disconnected(),
).await;
close_bound_upstream(bound).await;
break;
};
match Box::pin(forward_client_message(
client_message,
bound,
client_socket,
state,
context,
))
.await
{
RelayDisposition::Continue => {}
RelayDisposition::Close => {
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::client_disconnected(),
).await;
break;
}
RelayDisposition::UpstreamError(code) => {
warn!(
event_name = "responses_websocket_upstream_send_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
error_code = code,
"Upstream WebSocket send failed"
);
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::upstream_send_failed(),
).await;
send_gateway_error_with_status(
client_socket,
502,
code,
"Gateway could not forward the WebSocket event upstream",
).await;
close_bound_upstream(bound).await;
close_client_socket(client_socket, CLOSE_INTERNAL_ERROR, code).await;
break;
}
}
}
upstream_message = receive_optional_upstream(&mut bound.upstream) => {
let Some(upstream_message) = upstream_message else {
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::upstream_closed(),
).await;
bound.upstream = None;
close_client_socket(client_socket, 1000, "upstream_closed").await;
break;
};
let Ok(upstream_message) = upstream_message else {
warn!(
event_name = "responses_websocket_upstream_receive_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
"Upstream WebSocket receive failed"
);
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::upstream_receive_failed(),
).await;
send_gateway_error_with_status(
client_socket,
502,
"responses_websocket_receive_failed",
"Provider connection closed unexpectedly",
).await;
bound.upstream = None;
close_client_socket(client_socket, CLOSE_INTERNAL_ERROR, "upstream_receive_failed").await;
break;
};
let parsed_upstream_frame = match &upstream_message {
WreqWsMessage::Text(text) => {
ParsedResponsesWebSocketFrame::parse(text.as_str()).ok()
}
_ => None,
};
let parsed_upstream_event = parsed_upstream_frame
.as_ref()
.map(ParsedResponsesWebSocketFrame::event);
if let WreqWsMessage::Text(text) = &upstream_message {
debug!(
event_name = "responses_websocket_upstream_event",
log_type = "event",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
event_type = %parsed_upstream_frame
.as_ref()
.map(ParsedResponsesWebSocketFrame::event_type_for_log)
.unwrap_or_else(|| "invalid_json".to_string()),
frame_bytes = text.len(),
chunked = parsed_upstream_frame
.as_ref()
.is_some_and(ParsedResponsesWebSocketFrame::is_chunked),
active_turn = bound.turn_state.response_in_flight(),
"gateway received Responses WebSocket event"
);
}
if matches!(&upstream_message, WreqWsMessage::Binary(_)) {
mark_active_response_retry_unsafe(bound, "upstream_binary_frame");
} else if matches!(&upstream_message, WreqWsMessage::Text(_))
&& parsed_upstream_event.is_none()
{
mark_active_response_retry_unsafe(bound, "invalid_upstream_event");
}
if let Some(event) = parsed_upstream_event {
observe_active_response_rebind_safety(bound, event);
if bound.pending_adapter_drain.is_none()
&& bound.adapter.observes_upstream_events()
{
let adapter = bound.adapter;
if let Some(observation) = adapter.observe_upstream_event(event) {
let directive = observation.drain;
await_pending_adapter_observation(bound).await;
let state_for_observation = state.clone();
let trace_id = context.trace_id.clone();
let report_context = bound.decision_template.report_context.clone();
bound.pending_adapter_observation = Some(spawn_bounded_adapter_observation(async move {
adapter
.persist_upstream_observation(
&state_for_observation,
&trace_id,
report_context.as_ref(),
observation,
)
.await;
}));
if let Some(directive) = directive {
bound.pending_adapter_drain = Some(directive);
// A definitive quota signal must be visible to
// the next planner before a transparent retry.
await_pending_adapter_observation(bound).await;
}
}
}
}
let observation = match &upstream_message {
WreqWsMessage::Text(text) => {
let adapter = bound.adapter;
match parsed_upstream_frame.as_ref() {
Some(frame) => bound
.turn_state
.attempt_mut()
.and_then(|turn| turn.observe_upstream_frame(frame, adapter)),
None => {
if let Some(turn) = bound.turn_state.attempt_mut() {
turn.observe_invalid_upstream_text(text.as_str())
}
else {
None
}
}
}
}
_ => None,
};
if matches!(
observation,
Some(ResponsesWebSocketTurnObservation::Started)
| Some(ResponsesWebSocketTurnObservation::Terminal(_))
) {
if let Some(turn) = bound.turn_state.attempt_mut() {
turn.mark_stream_started(state).await;
}
}
let terminal_outcome = match observation {
Some(ResponsesWebSocketTurnObservation::Terminal(outcome)) => Some(outcome),
_ => None,
};
if matches!(&upstream_message, WreqWsMessage::Text(_))
&& parsed_upstream_frame.is_none()
{
let policy = fatal_relay_policy(FatalRelaySignal::InvalidUpstreamText);
finalize_active_turn(
bound,
state,
terminal_outcome.unwrap_or_else(
ResponsesWebSocketTurnOutcome::upstream_receive_failed,
),
)
.await;
send_responses_websocket_error(
client_socket,
policy.status_code,
"server_error",
policy.error_code,
policy.client_message,
)
.await;
close_bound_upstream(bound).await;
close_client_socket(
client_socket,
policy.close_code,
policy.close_reason,
)
.await;
break;
}
let is_close = matches!(upstream_message, WreqWsMessage::Close(_));
let drain_for_adapter = adapter_drain_ready(
bound.pending_adapter_drain,
bound.turn_state.response_in_flight(),
observation,
is_close,
);
let quota_facts = QuotaRelayFacts {
drain_ready: drain_for_adapter,
retry_current_turn: bound
.pending_adapter_drain
.is_some_and(|directive| directive.retry_current_turn)
&& bound
.turn_state
.logical()
.is_some_and(|turn| turn.quota_retry_block_reason().is_none()),
transparent_retry_failed: false,
usage_limit_error: parsed_upstream_event.is_some_and(is_usage_limit_error_event),
upstream_closed: is_close,
};
let mut quota_relay_action = classify_quota_relay(quota_facts);
if matches!(quota_relay_action, QuotaRelayAction::AttemptTransparentRetry) {
// detach_attempt 保留 logical turn:重试是同一轮请求的下一个 attempt。
let retry_turn = bound.turn_state.detach_attempt();
// 先结算旧 attempt 并等它落地,再规划下一个 attempt。两个理由:
//
// 1. 规划要读 health / adaptive / pool 状态,而这些正是旧
// attempt 结算时才投射的。普通的新 turn 早就在 client.rs 里
// 用 await_pending_turn_finalization 挡住了「基于陈旧状态
// 规划」,透明重试这条路径原先漏了这一步。
// 2. 旧 attempt 还占着自己的 pool key lease。不先释放,重试就
// 可能因为「这把 key 仍被占用」而挑不到本该可用的替代 key,
// 或者干脆判成无可用供应商。
let settled = match retry_turn {
Some(mut turn) => {
turn.release_admission().await;
settle_turn_finalization(
bound,
state,
turn,
terminal_outcome.unwrap_or_else(
ResponsesWebSocketTurnOutcome::upstream_closed,
),
)
.await
}
None => PreviousAttemptSettled::nothing_to_settle(),
};
// Planning and binding a replacement carries the complete
// scheduler/provider state machine. Keep that large future
// off the relay task's stack; the default Tokio/test worker
// stack is otherwise easy to exhaust on this rare branch.
if Box::pin(retry_active_turn_after_quota_exhaustion(
bound, state, context, settled,
))
.await
{
continue;
}
// 重试失败。旧 attempt 已经结算,logical turn 仍停在
// Replanning,所以后面分支里的 end() / finalize_active_turn
// 只会清掉 logical turn 而不会交出 attempt——不存在重复结算。
quota_relay_action = classify_quota_relay(QuotaRelayFacts {
retry_current_turn: false,
transparent_retry_failed: true,
..quota_facts
});
}
let detach_after_forward =
matches!(quota_relay_action, QuotaRelayAction::ForwardQuotaAndDetach);
if detach_after_forward && is_close {
let directive = bound
.pending_adapter_drain
.expect("adapter drain state should be present");
finalize_active_turn(
bound,
state,
terminal_outcome
.unwrap_or_else(ResponsesWebSocketTurnOutcome::provider_quota_exhausted),
)
.await;
send_gateway_error_with_status(
client_socket,
429,
directive.error_code,
"Provider connection closed after reporting exhausted quota; send a new response.create to select another Provider connection",
)
.await;
detach_exhausted_upstream(bound, directive, &context.trace_id).await;
continue;
}
// Standard Responses frames cross the gateway byte-for-byte unless PII
// restoration has something to replace. Codex may wrap public events with
// provider-private side-channel chunks; only that explicit envelope is
// peeled, and each retained event is serialized as a complete opaque Value.
// Observation and capture continue to consume the redacted event, while the
// final client hop receives restored text.
let relay_directive = parsed_upstream_frame
.as_ref()
.map(|frame| {
bound
.adapter
.relay_directive_for_upstream_event(frame.event())
});
let mut relay_send_error = None;
let mut relay_serialization_failed = false;
match relay_directive {
Some(ResponsesWebSocketRelayDirective::ForwardOriginal) => {
let restored = parsed_upstream_frame
.as_ref()
.and_then(|frame| {
bound
.redaction_restorer
.restore_provider_frame_text(frame.event())
});
let client_frame = match restored {
Some(text) => AxumWsMessage::Text(text.into()),
None => upstream_message_to_client(upstream_message.clone()),
};
match send_client_message(client_socket, client_frame).await {
Ok(()) => {
if let (Some(turn), Some(frame)) = (
bound.turn_state.attempt_mut(),
parsed_upstream_frame.as_ref(),
) {
turn.capture_client_frame(frame.event());
}
}
Err(error) => relay_send_error = Some(error),
}
}
Some(ResponsesWebSocketRelayDirective::ForwardEvents(events)) => {
for event in events {
let text = match bound
.redaction_restorer
.restore_provider_frame_text(event)
{
Some(restored) => restored,
None => match encode_opaque_websocket_event(event) {
Ok(encoded) => encoded,
Err(_) => {
relay_serialization_failed = true;
break;
}
},
};
match send_client_message(
client_socket,
AxumWsMessage::Text(text.into()),
)
.await
{
Ok(()) => {
if let Some(turn) = bound.turn_state.attempt_mut() {
turn.capture_client_frame(event);
}
}
Err(error) => {
relay_send_error = Some(error);
break;
}
}
}
}
Some(ResponsesWebSocketRelayDirective::SuppressProviderPrivate) => {}
None => {
if let Err(error) = send_client_message(
client_socket,
upstream_message_to_client(upstream_message.clone()),
)
.await
{
relay_send_error = Some(error);
}
}
}
if relay_serialization_failed {
warn!(
event_name = "responses_websocket_provider_event_serialization_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
provider_terminal_reached = terminal_outcome.is_some(),
"gateway could not serialize an opaque provider event"
);
bound
.turn_state
.record_client_delivery_aborted(CLIENT_DELIVERY_FAILED_REASON);
finalize_active_turn(
bound,
state,
settle_signal_for_client_delivery_failure(terminal_outcome),
)
.await;
send_gateway_error_with_status(
client_socket,
502,
"responses_websocket_event_serialization_failed",
"Gateway could not relay the provider event",
)
.await;
close_bound_upstream(bound).await;
close_client_socket(
client_socket,
CLOSE_INTERNAL_ERROR,
"provider_event_serialization_failed",
)
.await;
break;
}
if let Some(error) = relay_send_error {
warn!(
event_name = "responses_websocket_client_send_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
error_code = error.as_str(),
provider_terminal_reached = terminal_outcome.is_some(),
"gateway could not relay a provider event to the client"
);
// 投递失败是独立事实,不能覆盖已经到达的 provider 终态:
// 供应商已经完成推理并消耗 token,账单按它的终态计。
bound
.turn_state
.record_client_delivery_aborted(CLIENT_DELIVERY_FAILED_REASON);
finalize_active_turn(
bound,
state,
settle_signal_for_client_delivery_failure(terminal_outcome),
).await;
close_bound_upstream(bound).await;
break;
}
if let Some(outcome) = terminal_outcome {
finalize_active_turn(bound, state, outcome).await;
} else if is_close {
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::upstream_closed(),
)
.await;
}
if detach_after_forward {
let directive = bound
.pending_adapter_drain
.expect("adapter drain state should be present");
if bound.turn_state.response_in_flight() {
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::provider_quota_exhausted(),
)
.await;
}
detach_exhausted_upstream(bound, directive, &context.trace_id).await;
continue;
}
if drain_for_adapter {
let directive = bound
.pending_adapter_drain
.expect("adapter drain state should be present");
detach_exhausted_upstream(bound, directive, &context.trace_id).await;
continue;
}
if is_close {
bound.upstream = None;
break;
}
}
}
}
}
pub(super) async fn wait_for_connection_permit_loss(
permit: Option<&aether_runtime::AdmissionPermit>,
) {
let Some(permit) = permit else {
std::future::pending::<()>().await;
return;
};
let mut health = tokio::time::interval(Duration::from_secs(1));
health.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
loop {
health.tick().await;
if !permit.is_healthy() {
return;
}
}
}
@@ -0,0 +1,173 @@
//! Per-turn control-plane refresh for long-lived Responses WebSockets.
//!
//! An Upgrade authenticates the connection, but it must not freeze API-key,
//! wallet, IP, model, or RPM policy for up to an hour. This module produces one
//! live decision and its exact strong API-key snapshot for every
//! `response.create`; the caller uses that pair consistently for rate limiting,
//! redaction, model authorization, planning, admission, balance, and retries.
use axum::http::StatusCode;
use serde_json::Value;
use crate::ai_serving::GatewayAuthApiKeySnapshot;
use crate::control::{
refresh_execution_runtime_auth_context_with_snapshot, request_model_local_rejection,
GatewayControlDecision, GatewayLocalAuthRejection,
};
use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext;
use crate::handlers::proxy::websocket::session::WEBSOCKET_LOG_TRANSPORT;
use crate::handlers::shared::ip_rules_allow;
use crate::{AppState, GatewayError};
const LOG_TARGET: &str = "aether_gateway::handlers::proxy::responses_ws";
macro_rules! warn {
($($arg:tt)*) => {
tracing::warn!(target: LOG_TARGET, $($arg)*)
};
}
#[derive(Debug, Clone)]
pub(super) struct ResponsesWebSocketTurnControl {
pub(super) decision: GatewayControlDecision,
pub(super) auth_snapshot: Option<GatewayAuthApiKeySnapshot>,
pub(super) rpm_bypassed: bool,
}
pub(super) async fn resolve_responses_websocket_turn_control(
state: &AppState,
context: &WebSocketRequestContext,
parts: &http::request::Parts,
client_event: &Value,
) -> Result<ResponsesWebSocketTurnControl, GatewayError> {
if state
.admin_security_ip_blacklisted(context.client_ip)
.await?
{
return Err(GatewayError::Client {
status: StatusCode::FORBIDDEN,
message: "The current IP is blocked".to_string(),
});
}
let mut decision = context.decision.clone();
let auth_snapshot = if let Some(auth_context) = decision.auth_context.take() {
let (refreshed, snapshot) = refresh_execution_runtime_auth_context_with_snapshot(
state,
auth_context,
decision.auth_endpoint_signature.as_deref(),
)
.await?;
decision.local_auth_rejection = refreshed.local_rejection.clone();
decision.auth_context = Some(refreshed);
snapshot
} else {
None
};
// Model-directive configuration is mutable policy too; do not retain the
// Upgrade-time snapshot for the lifetime of the socket.
decision.model_directive_policy =
crate::system_features::ModelDirectivePolicySnapshot::load(state).await;
if let Some(rejection) = decision.local_auth_rejection.clone() {
return Err(websocket_auth_rejection_error(rejection));
}
let Some(auth_context) = decision.auth_context.as_ref() else {
return Err(websocket_auth_rejection_error(
GatewayLocalAuthRejection::InvalidApiKey,
));
};
if !auth_context.access_allowed
|| auth_context.user_id.trim().is_empty()
|| auth_context.api_key_id.trim().is_empty()
{
return Err(websocket_auth_rejection_error(
GatewayLocalAuthRejection::InvalidApiKey,
));
}
if !ip_rules_allow(auth_context.ip_rules.as_deref(), context.client_ip) {
return Err(websocket_auth_rejection_error(
GatewayLocalAuthRejection::IpNotAllowed {
remote_ip: context.client_ip.to_string(),
},
));
}
let body = serde_json::to_vec(client_event)
.map(axum::body::Bytes::from)
.map_err(|error| GatewayError::Internal(error.to_string()))?;
if let Some(rejection) =
request_model_local_rejection(state, Some(&decision), &parts.uri, &parts.headers, &body)
.await?
{
return Err(websocket_auth_rejection_error(rejection));
}
let rpm_bypassed = match state.admin_security_ip_whitelisted(context.client_ip).await {
Ok(value) => value,
Err(error) => {
warn!(
event_name = "responses_websocket_turn_ip_whitelist_check_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
client_ip = %context.client_ip,
error = ?error,
"gateway applied ordinary WebSocket RPM after the live IP whitelist check failed"
);
false
}
};
Ok(ResponsesWebSocketTurnControl {
decision,
auth_snapshot,
rpm_bypassed,
})
}
fn websocket_auth_rejection_error(rejection: GatewayLocalAuthRejection) -> GatewayError {
let (status, message) = match rejection {
GatewayLocalAuthRejection::InvalidApiKey => {
(StatusCode::UNAUTHORIZED, "The API key is invalid")
}
GatewayLocalAuthRejection::LockedApiKey => (
StatusCode::FORBIDDEN,
"The API key is locked and cannot be used",
),
GatewayLocalAuthRejection::WalletUnavailable => {
(StatusCode::FORBIDDEN, "The account wallet is unavailable")
}
GatewayLocalAuthRejection::BalanceDenied { remaining } => {
let message = match remaining {
Some(remaining) => format!("Insufficient balance (remaining: ${remaining:.2})"),
None => "Insufficient balance".to_string(),
};
return GatewayError::Client {
status: StatusCode::TOO_MANY_REQUESTS,
message,
};
}
GatewayLocalAuthRejection::ProviderNotAllowed { .. } => (
StatusCode::FORBIDDEN,
"The provider is not allowed for this API key",
),
GatewayLocalAuthRejection::ApiFormatNotAllowed { .. } => (
StatusCode::FORBIDDEN,
"The API format is not allowed for this API key",
),
GatewayLocalAuthRejection::ModelNotAllowed { .. } => (
StatusCode::FORBIDDEN,
"The requested model is not allowed for this API key",
),
GatewayLocalAuthRejection::IpNotAllowed { .. } => (
StatusCode::UNAUTHORIZED,
"The current IP is not allowed for this API key",
),
};
GatewayError::Client {
status,
message: message.to_string(),
}
}
@@ -0,0 +1,570 @@
//! Parsed OpenAI Responses WebSocket text frames.
//!
//! A relay frame is parsed once and then shared by the protocol adapter, turn
//! accounting, retry safety, and connection lifecycle code. Keeping the raw
//! text as a borrow avoids copying the websocket payload while the relay is
//! processing it.
use serde_json::Value;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) struct ResponsesWebSocketFrameTerminal {
pub(super) status_code: u16,
pub(super) cancelled: bool,
}
#[derive(Debug)]
pub(super) struct ParsedResponsesWebSocketFrame<'a> {
raw_text: &'a str,
event: Value,
event_type: Option<String>,
status: Option<u16>,
started: bool,
terminal: Option<ResponsesWebSocketFrameTerminal>,
terminal_event: Option<Value>,
chunked: bool,
}
impl<'a> ParsedResponsesWebSocketFrame<'a> {
pub(super) fn parse(raw_text: &'a str) -> serde_json::Result<Self> {
let event = serde_json::from_str::<Value>(raw_text)?;
let events = protocol_events_of(&event);
let started = events.iter().copied().any(event_is_started);
// A batch carries at most one terminal in practice. Taking the first
// in document order keeps the outcome deterministic if that ever
// stops being true.
let terminal_entry = events
.iter()
.copied()
.find_map(|candidate| terminal_for_event(candidate).map(|term| (candidate, term)));
let terminal = terminal_entry.map(|(_, terminal)| terminal);
// The terminal event describes the turn's outcome, so it is the one
// worth naming in logs and recording as the terminal error body.
let event_type = terminal_entry
.map(|(candidate, _)| candidate)
.or_else(|| events.last().copied())
.and_then(event_type_of)
.map(str::to_string);
let terminal_event = terminal_entry.map(|(candidate, _)| candidate.clone());
let chunked = event.get("chunks").and_then(Value::as_array).is_some();
let status = terminal.map(|terminal| terminal.status_code);
Ok(Self {
raw_text,
event,
event_type,
status,
started,
terminal,
terminal_event,
chunked,
})
}
/// The protocol events this frame carries.
///
/// Codex batches standard `response.*` events into a `{"chunks":[...]}`
/// envelope, so one frame can carry several events — and the terminal one
/// may be buried inside the batch. Every consumer that interprets event
/// semantics must walk this rather than the envelope, or a batched
/// `response.completed` goes unnoticed and wedges the turn.
pub(super) fn protocol_events(&self) -> Vec<&Value> {
protocol_events_of(&self.event)
}
/// The individual event that ended the turn, unwrapped from its batch.
pub(super) fn terminal_event(&self) -> Option<&Value> {
self.terminal_event.as_ref()
}
pub(super) fn is_chunked(&self) -> bool {
self.chunked
}
pub(super) fn raw_text(&self) -> &'a str {
self.raw_text
}
pub(super) fn event(&self) -> &Value {
&self.event
}
pub(super) fn event_type(&self) -> Option<&str> {
self.event_type.as_deref()
}
pub(super) fn status(&self) -> Option<u16> {
self.status
}
pub(super) fn is_started(&self) -> bool {
self.started
}
pub(super) fn is_terminal(&self) -> bool {
self.terminal.is_some()
}
pub(super) fn terminal(&self) -> Option<ResponsesWebSocketFrameTerminal> {
self.terminal
}
/// Return a bounded label suitable for structured logs. Event payloads
/// are never inserted directly into a log field.
pub(super) fn event_type_for_log(&self) -> String {
self.event_type
.as_deref()
.map(safe_websocket_event_label)
.unwrap_or_else(|| "invalid_json".to_string())
}
}
/// Encodes one event peeled from a provider-private envelope without applying
/// an event-type or field projection.
///
/// Direct provider events should use [`ParsedResponsesWebSocketFrame::raw_text`]
/// so their bytes remain identical. This helper exists only for batch
/// envelopes that cannot be relayed as a whole: serializing the complete
/// [`Value`] preserves every known and future JSON member.
pub(super) fn encode_opaque_websocket_event(event: &Value) -> serde_json::Result<String> {
serde_json::to_string(event)
}
/// Flattens a frame into the events it carries. An envelope may name its own
/// `type` *and* batch further events under `chunks`; both are protocol events.
fn protocol_events_of(event: &Value) -> Vec<&Value> {
let mut events = Vec::new();
if event_type_of(event).is_some() {
events.push(event);
}
if let Some(chunks) = event.get("chunks").and_then(Value::as_array) {
events.extend(chunks.iter().filter(|chunk| event_type_of(chunk).is_some()));
}
// An unrecognized shape is still relayed and still accounted for, so it
// must not vanish from the observer's view of the stream.
if events.is_empty() {
events.push(event);
}
events
}
fn event_type_of(event: &Value) -> Option<&str> {
event.get("type").and_then(Value::as_str)
}
fn event_is_started(event: &Value) -> bool {
matches!(
event_type_of(event).unwrap_or_default(),
"response.created" | "response.in_progress" | "response.queued"
)
}
/// 读取 `response.incomplete` 携带的 `incomplete_details.reason`。
///
/// 标准位置是 `response.incomplete_details.reason`;批量封装偶尔把
/// `incomplete_details` 直接放在事件顶层,两处都要看,否则合法终态会被漏判。
fn responses_incomplete_reason(event: &Value) -> Option<&str> {
[
event.pointer("/response/incomplete_details/reason"),
event.pointer("/incomplete_details/reason"),
]
.into_iter()
.flatten()
.filter_map(Value::as_str)
.map(str::trim)
.find(|reason| !reason.is_empty())
}
fn responses_incomplete_has_explicit_error(event: &Value) -> bool {
[event.get("error"), event.pointer("/response/error")]
.into_iter()
.flatten()
.any(|error| !error.is_null())
}
/// Derives only the fallback status for `response.incomplete`.
///
/// A non-empty reason is provider-owned protocol data. Treating it as a fixed
/// allowlist would turn every future legitimate reason into a synthetic 502
/// and incorrectly penalize provider health. Missing/malformed reasons and
/// explicit error markers still fail closed; numeric status and recognized
/// error codes continue to override this fallback in
/// [`websocket_event_status_code`].
fn responses_incomplete_default_status(event: &Value) -> u16 {
match responses_incomplete_reason(event) {
None => 502,
Some(reason)
if reason.eq_ignore_ascii_case("error")
|| reason.eq_ignore_ascii_case("server_error") =>
{
502
}
Some(_) if responses_incomplete_has_explicit_error(event) => 502,
Some(_) => 200,
}
}
fn terminal_for_event(event: &Value) -> Option<ResponsesWebSocketFrameTerminal> {
match event_type_of(event).unwrap_or_default() {
"response.completed" => Some(ResponsesWebSocketFrameTerminal {
status_code: websocket_event_status_code(event, 200),
cancelled: false,
}),
// A non-empty provider reason is a normal terminal by default, including
// future reasons Aether does not yet know. Explicit status/error data
// still wins, so quota and server failures retain their failure status.
"response.incomplete" => Some(ResponsesWebSocketFrameTerminal {
status_code: websocket_event_status_code(
event,
responses_incomplete_default_status(event),
),
cancelled: false,
}),
"response.cancelled" => Some(ResponsesWebSocketFrameTerminal {
status_code: 499,
cancelled: true,
}),
"response.failed" => Some(ResponsesWebSocketFrameTerminal {
status_code: websocket_event_status_code(event, 502),
cancelled: false,
}),
"error" => Some(ResponsesWebSocketFrameTerminal {
status_code: websocket_event_status_code(event, 502),
cancelled: false,
}),
_ => None,
}
}
fn websocket_event_status_code(event: &Value, default: u16) -> u16 {
if let Some(status_code) = event
.get("status_code")
.or_else(|| event.get("status"))
.or_else(|| {
event
.get("response")
.and_then(|response| response.get("status_code"))
})
.and_then(Value::as_u64)
.and_then(|value| u16::try_from(value).ok())
.filter(|value| *value > 0)
{
return status_code;
}
let error_code = [
event.pointer("/error/type"),
event.pointer("/error/code"),
event.pointer("/response/error/type"),
event.pointer("/response/error/code"),
]
.into_iter()
.flatten()
.filter_map(Value::as_str)
.map(str::to_ascii_lowercase)
.find(|value| !value.trim().is_empty());
match error_code.as_deref() {
Some(
"usage_limit_reached" | "insufficient_quota" | "rate_limit_exceeded" | "quota_exceeded",
) => 429,
Some("invalid_api_key" | "authentication_error") => 401,
Some("invalid_request_error" | "invalid_request" | "model_not_found") => 400,
Some("overloaded" | "server_error" | "service_unavailable") => 503,
_ => default,
}
}
fn safe_websocket_event_label(value: &str) -> String {
let value = value.trim();
if value.is_empty()
|| value.len() > 80
|| !value
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'.' | b'_' | b'-'))
{
return "unknown".to_string();
}
value.to_string()
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::{encode_opaque_websocket_event, ParsedResponsesWebSocketFrame};
#[test]
fn parses_started_frame_once_with_raw_text_and_event_metadata() {
let raw = r#"{"type":"response.in_progress","response":{"status":200}}"#;
let frame = ParsedResponsesWebSocketFrame::parse(raw).expect("valid frame");
assert_eq!(frame.raw_text(), raw);
assert_eq!(frame.event_type(), Some("response.in_progress"));
assert_eq!(frame.status(), None);
assert!(frame.is_started());
assert!(!frame.is_terminal());
assert_eq!(frame.event()["response"]["status"], 200);
assert_eq!(frame.event_type_for_log(), "response.in_progress");
}
#[test]
fn future_response_event_keeps_its_exact_original_text_and_unknown_fields() {
let raw = "{ \n \"future_top_level\": {\"nested\": [1, true, null]}, \n \"type\": \"response.future_capability.delta\", \n \"delta\": {\"new_wire_shape\": \"opaque\"}\n}";
let frame = ParsedResponsesWebSocketFrame::parse(raw).expect("valid future event");
assert_eq!(frame.raw_text(), raw);
assert_eq!(
frame.event()["future_top_level"],
json!({"nested": [1, true, null]})
);
assert_eq!(frame.event()["delta"], json!({"new_wire_shape": "opaque"}));
assert!(!frame.is_terminal());
}
#[test]
fn peeled_batch_event_encoding_preserves_the_complete_opaque_value() {
let frame = ParsedResponsesWebSocketFrame::parse(
r#"{"chunks":[{"type":"response.future.done","future_capability":{"mode":"new"},"response":{"id":"resp_future","future_usage":{"novel_tokens":7}}}]}"#,
)
.expect("valid private envelope");
let events = frame.protocol_events();
let event = events.first().expect("one future response event");
let encoded = encode_opaque_websocket_event(event).expect("Value serialization succeeds");
let round_trip: serde_json::Value =
serde_json::from_str(&encoded).expect("encoded event stays valid JSON");
assert_eq!(round_trip, **event);
assert_eq!(round_trip["future_capability"], json!({"mode": "new"}));
assert_eq!(round_trip["response"]["future_usage"]["novel_tokens"], 7);
}
#[test]
fn classifies_terminal_status_and_cancellation() {
let completed = ParsedResponsesWebSocketFrame::parse(
r#"{"type":"response.completed","status_code":201}"#,
)
.expect("valid frame");
assert_eq!(completed.status(), Some(201));
assert_eq!(
completed
.terminal()
.map(|terminal| (terminal.status_code, terminal.cancelled)),
Some((201, false))
);
let cancelled = ParsedResponsesWebSocketFrame::parse(r#"{"type":"response.cancelled"}"#)
.expect("valid frame");
assert_eq!(cancelled.status(), Some(499));
assert_eq!(
cancelled
.terminal()
.map(|terminal| (terminal.status_code, terminal.cancelled)),
Some((499, true))
);
let error = ParsedResponsesWebSocketFrame::parse(
r#"{"type":"error","status_code":429,"error":{"type":"usage_limit_reached"}}"#,
)
.expect("valid frame");
assert_eq!(error.status(), Some(429));
assert!(error.is_terminal());
let failed = ParsedResponsesWebSocketFrame::parse(
r#"{"type":"response.failed","response":{"error":{"code":"rate_limit_exceeded"}}}"#,
)
.expect("valid frame");
assert_eq!(failed.status(), Some(429));
}
#[test]
fn a_legitimate_incomplete_is_a_terminal_but_not_a_provider_failure() {
for reason in [
"max_output_tokens",
"max_tokens",
"content_filter",
"tool_calls",
"function_call",
"MAX_OUTPUT_TOKENS",
] {
let raw = format!(
r#"{{"type":"response.incomplete","response":{{"status":"incomplete","incomplete_details":{{"reason":"{reason}"}}}}}}"#
);
let frame = ParsedResponsesWebSocketFrame::parse(&raw).expect("valid frame");
assert!(frame.is_terminal(), "{reason} should end the turn");
assert_eq!(
frame
.terminal()
.map(|terminal| (terminal.status_code, terminal.cancelled)),
Some((200, false)),
"{reason} is a legitimate terminal result, not a 502 provider failure"
);
}
}
#[test]
fn a_top_level_incomplete_details_reason_is_also_honored() {
let frame = ParsedResponsesWebSocketFrame::parse(
r#"{"type":"response.incomplete","incomplete_details":{"reason":"max_output_tokens"}}"#,
)
.expect("valid frame");
assert_eq!(frame.status(), Some(200));
}
#[test]
fn an_incomplete_without_a_reason_or_with_a_failure_reason_stays_a_provider_failure() {
for raw in [
r#"{"type":"response.incomplete"}"#,
r#"{"type":"response.incomplete","response":{"incomplete_details":null}}"#,
r#"{"type":"response.incomplete","response":{"incomplete_details":{"reason":""}}}"#,
r#"{"type":"response.incomplete","response":{"incomplete_details":{"reason":"error"}}}"#,
r#"{"type":"response.incomplete","response":{"incomplete_details":{"reason":"server_error"}}}"#,
] {
let frame = ParsedResponsesWebSocketFrame::parse(raw).expect("valid frame");
assert_eq!(
frame.status(),
Some(502),
"an incomplete without a usable reason must stay a provider failure: {raw}"
);
}
}
#[test]
fn a_future_incomplete_reason_is_forward_compatible_without_hiding_explicit_errors() {
let future = ParsedResponsesWebSocketFrame::parse(
r#"{"type":"response.incomplete","response":{"incomplete_details":{"reason":"future_context_boundary"}}}"#,
)
.expect("valid frame");
assert_eq!(future.status(), Some(200));
let future_with_error = ParsedResponsesWebSocketFrame::parse(
r#"{"type":"response.incomplete","response":{"error":{"code":"future_provider_error"},"incomplete_details":{"reason":"future_context_boundary"}}}"#,
)
.expect("valid frame");
assert_eq!(future_with_error.status(), Some(502));
}
#[test]
fn a_legitimate_incomplete_still_respects_an_explicit_provider_status() {
let explicit = ParsedResponsesWebSocketFrame::parse(
r#"{"type":"response.incomplete","status_code":503,"response":{"incomplete_details":{"reason":"max_output_tokens"}}}"#,
)
.expect("valid frame");
assert_eq!(explicit.status(), Some(503));
let quota = ParsedResponsesWebSocketFrame::parse(
r#"{"type":"response.incomplete","response":{"error":{"code":"rate_limit_exceeded"},"incomplete_details":{"reason":"max_output_tokens"}}}"#,
)
.expect("valid frame");
assert_eq!(quota.status(), Some(429));
}
#[test]
fn a_legitimate_incomplete_batched_inside_a_chunks_envelope_is_not_a_failure() {
let frame = ParsedResponsesWebSocketFrame::parse(
r#"{"chunks":[{"type":"response.output_text.delta","delta":"hi"},{"type":"response.incomplete","response":{"incomplete_details":{"reason":"max_output_tokens"},"usage":{"total_tokens":9}}}]}"#,
)
.expect("valid frame");
assert!(frame.is_chunked());
assert!(frame.is_terminal());
assert_eq!(frame.status(), Some(200));
assert_eq!(frame.event_type(), Some("response.incomplete"));
assert_eq!(
frame.terminal_event().and_then(|event| event
.pointer("/response/usage/total_tokens")
.and_then(serde_json::Value::as_u64)),
Some(9)
);
}
#[test]
fn detects_a_terminal_batched_inside_a_chunks_envelope() {
let frame = ParsedResponsesWebSocketFrame::parse(
r#"{"chunks":[{"type":"response.output_text.delta","delta":"hi"},{"type":"response.completed","response":{"usage":{"total_tokens":8}}}]}"#,
)
.expect("valid frame");
assert!(frame.is_chunked());
assert!(frame.is_terminal());
assert_eq!(frame.status(), Some(200));
// The label and the recorded error body must name the event that ended
// the turn, not the envelope.
assert_eq!(frame.event_type(), Some("response.completed"));
assert_eq!(
frame.terminal_event().and_then(|event| event
.pointer("/response/usage/total_tokens")
.and_then(serde_json::Value::as_u64)),
Some(8)
);
assert_eq!(frame.protocol_events().len(), 2);
}
#[test]
fn detects_a_start_event_batched_inside_a_chunks_envelope() {
let frame = ParsedResponsesWebSocketFrame::parse(
r#"{"chunks":[{"type":"codex.rate_limits"},{"type":"response.created"}]}"#,
)
.expect("valid frame");
assert!(frame.is_started());
assert!(!frame.is_terminal());
assert_eq!(frame.protocol_events().len(), 2);
}
#[test]
fn an_envelope_may_carry_its_own_type_alongside_batched_events() {
let frame = ParsedResponsesWebSocketFrame::parse(
r#"{"type":"codex.response.metadata","chunks":[{"type":"response.failed","response":{"error":{"code":"rate_limit_exceeded"}}}]}"#,
)
.expect("valid frame");
assert_eq!(frame.protocol_events().len(), 2);
assert!(frame.is_terminal());
assert_eq!(frame.status(), Some(429));
assert_eq!(frame.event_type(), Some("response.failed"));
}
#[test]
fn a_batch_without_a_terminal_does_not_end_the_turn() {
let frame = ParsedResponsesWebSocketFrame::parse(
r#"{"chunks":[{"type":"response.output_text.delta","delta":"a"},{"type":"response.output_text.delta","delta":"b"}]}"#,
)
.expect("valid frame");
assert!(!frame.is_terminal());
assert!(!frame.is_started());
assert!(frame.terminal_event().is_none());
}
#[test]
fn an_unrecognized_shape_is_still_surfaced_as_one_event() {
let frame =
ParsedResponsesWebSocketFrame::parse(r#"{"unexpected":true}"#).expect("valid frame");
assert_eq!(frame.protocol_events().len(), 1);
assert!(!frame.is_chunked());
assert!(!frame.is_terminal());
assert_eq!(frame.event_type(), None);
assert_eq!(frame.event_type_for_log(), "invalid_json");
}
#[test]
fn preserves_safe_log_label_boundaries() {
let unsafe_label =
ParsedResponsesWebSocketFrame::parse(r#"{"type":"not safe / contains spaces"}"#)
.expect("valid frame");
assert_eq!(unsafe_label.event_type_for_log(), "unknown");
let missing_label =
ParsedResponsesWebSocketFrame::parse(r#"{"message":"ok"}"#).expect("valid frame");
assert_eq!(missing_label.event_type_for_log(), "invalid_json");
}
#[test]
fn rejects_invalid_json() {
assert!(ParsedResponsesWebSocketFrame::parse("not-json").is_err());
}
}
@@ -0,0 +1,588 @@
//! Turn finalization and terminal error mapping for a Responses WebSocket.
//!
//! A connection can outlive a turn, so persistence and adapter observation
//! handles are joined in order before the next turn is planned.
use std::time::Duration;
use axum::extract::ws::WebSocket;
use axum::http::StatusCode;
use tokio::task::JoinHandle;
use tokio::time::timeout;
use super::state::BoundResponsesConnection;
use super::turn::{
begin_unowned_responses_websocket_turn, ResponsesProviderAttempt, ResponsesWebSocketTurnOutcome,
};
use crate::handlers::proxy::websocket::session::{
CLOSE_INTERNAL_ERROR, CLOSE_POLICY_VIOLATION, CLOSE_TRY_AGAIN, WEBSOCKET_LOG_TRANSPORT,
};
use crate::handlers::proxy::websocket::transport::send_responses_websocket_error;
use crate::{AppState, GatewayError};
const RESPONSES_WEBSOCKET_ADAPTER_OBSERVATION_TIMEOUT: Duration = Duration::from_secs(5);
const LOG_TARGET: &str = "aether_gateway::handlers::proxy::responses_ws";
macro_rules! warn {
($($arg:tt)*) => {
tracing::warn!(target: LOG_TARGET, $($arg)*)
};
}
/// Owns the in-flight turn so that losing the relay task still finalizes it.
///
/// Every ordinary exit path takes the turn out of here and finalizes it
/// explicitly. This guard only covers the paths that are not exit paths at all
/// — a panic in the relay loop, or the task being dropped — where the turn
/// would otherwise be discarded with its usage row left `Pending`, its
/// candidate row left `Streaming`, and its distributed pool key lease leaked
/// until the lease expires. Mirrors the HTTP path's `DirectPassthroughFinalizer`.
pub(super) struct ActiveProviderAttempt {
turn: Option<ResponsesProviderAttempt>,
state: AppState,
}
impl ActiveProviderAttempt {
pub(super) fn new(state: &AppState, turn: ResponsesProviderAttempt) -> Self {
Self {
turn: Some(turn),
state: state.clone(),
}
}
/// Hands the turn back to a caller that will finalize it explicitly.
pub(super) fn disarm(mut self) -> ResponsesProviderAttempt {
self.turn
.take()
.expect("an armed active turn always holds its turn")
}
}
impl std::ops::Deref for ActiveProviderAttempt {
type Target = ResponsesProviderAttempt;
fn deref(&self) -> &Self::Target {
self.turn
.as_ref()
.expect("an armed active turn always holds its turn")
}
}
impl std::ops::DerefMut for ActiveProviderAttempt {
fn deref_mut(&mut self) -> &mut Self::Target {
self.turn
.as_mut()
.expect("an armed active turn always holds its turn")
}
}
/// Starts a turn and arms its cancellation fallback before control returns to
/// code that can await an upstream bind or socket write.
pub(super) async fn begin_responses_websocket_turn(
state: &AppState,
trace_id: &str,
parts: http::request::Parts,
control_decision: &crate::control::GatewayControlDecision,
decision: crate::ai_serving::AiExecutionDecision,
client_event: &serde_json::Value,
) -> Result<ActiveProviderAttempt, GatewayError> {
let state = state.clone();
let trace_id = trace_id.to_string();
let owner_timeout = state
.frontdoor_runtime_guards
.local_execution_planning_timeout;
let control_decision = control_decision.clone();
let client_event = client_event.clone();
// Beginning an attempt performs several indispensable async writes before
// an `ActiveProviderAttempt` can exist (balance/admission, Pending usage,
// and candidate state). Run that whole transition in an owned task. If the
// relay/session future is cancelled while awaiting it, Tokio detaches this
// task; it still reaches either an explicitly cleaned-up error or an armed
// guard whose dropped output finalizes the attempt.
await_owned_turn_begin(
async move {
let turn = begin_unowned_responses_websocket_turn(
&state,
&parts,
&control_decision,
decision,
&client_event,
)
.await?;
Ok(ActiveProviderAttempt::new(&state, turn))
},
owner_timeout,
trace_id,
)
.await
}
async fn await_owned_turn_begin<T>(
begin: impl std::future::Future<Output = Result<T, GatewayError>> + Send + 'static,
owner_timeout: Duration,
trace_id: String,
) -> Result<T, GatewayError>
where
T: Send + 'static,
{
await_owned_turn_begin_with_timeout(begin, owner_timeout, trace_id).await
}
async fn await_owned_turn_begin_with_timeout<T>(
begin: impl std::future::Future<Output = Result<T, GatewayError>> + Send + 'static,
owner_timeout: Duration,
trace_id: String,
) -> Result<T, GatewayError>
where
T: Send + 'static,
{
tokio::spawn(async move {
tokio::time::timeout(owner_timeout, begin)
.await
.map_err(|_| GatewayError::LocalExecutionPlanningTimeout {
trace_id,
phase: "responses_websocket_turn_begin_owner",
timeout_ms: owner_timeout.as_millis() as u64,
})?
})
.await
.map_err(|error| {
GatewayError::Internal(format!(
"Responses WebSocket turn begin task failed before ownership transfer: {error}"
))
})?
}
impl Drop for ActiveProviderAttempt {
fn drop(&mut self) {
let Some(turn) = self.turn.take() else {
return;
};
let outcome = turn.abandonment_outcome();
let state = self.state.clone();
// No runtime means the process is going down; the spawn could not
// complete anyway.
if let Ok(handle) = tokio::runtime::Handle::try_current() {
warn!(
event_name = "responses_websocket_turn_abandoned",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
"gateway finalized a Responses WebSocket turn whose relay task went away"
);
handle.spawn(async move {
turn.finalize_detached(&state, outcome).await;
});
}
}
}
/// 结束当前 logical turn 并结算它的 attempt。
///
/// `end()` 同时清掉 logical turn 和 attempt,取代原来「take active_turn +
/// 在每个出口手写 `active_response_create = None`」的两步组合。
pub(super) async fn finalize_active_turn(
bound: &mut BoundResponsesConnection,
state: &AppState,
outcome: ResponsesWebSocketTurnOutcome,
) {
if let Some(turn) = bound.turn_state.end() {
queue_turn_finalization(bound, state, turn, outcome).await;
}
}
pub(super) async fn queue_turn_finalization(
bound: &mut BoundResponsesConnection,
state: &AppState,
turn: ActiveProviderAttempt,
outcome: ResponsesWebSocketTurnOutcome,
) {
await_pending_adapter_observation(bound).await;
await_pending_turn_finalization(bound).await;
bound.pending_turn_finalization = Some(spawn_guarded_turn_finalization(
state.clone(),
turn,
outcome,
));
}
/// 「上一个 attempt 已经结算完毕」的凭证。
///
/// 只能由本模块颁发,且只有在结算真正落地之后。规划下一个 attempt 的入口
/// ([`super::quota::retry_active_turn_after_quota_exhaustion`]) 要求这个参数,
/// 于是「先结算、再规划」成为签名的一部分,而不是一句注释——顺序写反连编译都
/// 过不了。
pub(super) struct PreviousAttemptSettled(());
impl PreviousAttemptSettled {
/// 没有 attempt 要结算(连接此刻不在 `Responding`)。
pub(super) const fn nothing_to_settle() -> Self {
Self(())
}
}
/// 结算一个 attempt 并等它落地。
///
/// 与 [`queue_turn_finalization`] 的区别只在于「等」:后者把 handle 挂在连接上
/// 让 relay loop 继续跑,适用于结算之后不再需要读取共享状态的出口;这个用在
/// 必须先看到结算结果才能继续的路径上——典型的就是透明重试,它紧接着要按
/// health / adaptive / pool 状态规划下一个 attempt。
pub(super) async fn settle_turn_finalization(
bound: &mut BoundResponsesConnection,
state: &AppState,
turn: ActiveProviderAttempt,
outcome: ResponsesWebSocketTurnOutcome,
) -> PreviousAttemptSettled {
queue_turn_finalization(bound, state, turn, outcome).await;
await_pending_turn_finalization(bound).await;
PreviousAttemptSettled(())
}
pub(super) fn spawn_bounded_adapter_observation(
observation: impl std::future::Future<Output = ()> + Send + 'static,
) -> JoinHandle<()> {
spawn_bounded_adapter_observation_with_timeout(
observation,
RESPONSES_WEBSOCKET_ADAPTER_OBSERVATION_TIMEOUT,
)
}
fn spawn_bounded_adapter_observation_with_timeout(
observation: impl std::future::Future<Output = ()> + Send + 'static,
owner_timeout: Duration,
) -> JoinHandle<()> {
tokio::spawn(async move {
if timeout(owner_timeout, observation).await.is_err() {
warn!(
event_name = "responses_websocket_adapter_observation_timeout",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
timeout_ms = owner_timeout.as_millis() as u64,
"gateway stopped a timed-out Responses WebSocket adapter observation"
);
}
})
}
pub(super) async fn await_pending_adapter_observation(bound: &mut BoundResponsesConnection) {
if let Some(handle) = bound.pending_adapter_observation.take() {
if let Err(error) = handle.await {
warn!(
event_name = "responses_websocket_adapter_observation_join_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
error = ?error,
"gateway Responses WebSocket adapter observation task failed"
);
}
}
}
pub(super) fn finalize_unbound_turn(
state: AppState,
turn: ActiveProviderAttempt,
outcome: ResponsesWebSocketTurnOutcome,
) -> JoinHandle<()> {
spawn_guarded_turn_finalization(state, turn, outcome)
}
fn spawn_guarded_turn_finalization(
state: AppState,
turn: ActiveProviderAttempt,
outcome: ResponsesWebSocketTurnOutcome,
) -> JoinHandle<()> {
// Spawn synchronously while the armed guard is still owned here. Caller
// cancellation cannot drop an unguarded attempt between cleanup awaits.
tokio::spawn(async move {
let mut turn = turn;
turn.release_admission().await;
turn.disarm().finalize_detached(&state, outcome).await;
})
}
pub(super) async fn await_turn_finalization_handle(handle: JoinHandle<()>) {
// Do not abort terminal persistence here. Each I/O stage inside the turn
// finalizer is independently bounded, and aborting the owner would skip
// pool-lease cleanup and leave usage/candidate state non-terminal.
match handle.await {
Ok(()) => {}
Err(error) => {
warn!(
event_name = "responses_websocket_turn_finalization_join_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
error = ?error,
"gateway Responses WebSocket turn finalizer task failed"
);
}
}
}
pub(super) async fn await_pending_turn_finalization(bound: &mut BoundResponsesConnection) {
if let Some(handle) = bound.pending_turn_finalization.take() {
await_turn_finalization_handle(handle).await;
}
}
pub(super) async fn send_responses_websocket_turn_start_error(
client_socket: &mut WebSocket,
error: &GatewayError,
) {
let status_code = responses_websocket_turn_start_http_status(error);
match error {
GatewayError::Client { status, message } => {
let (error_type, code) = if status.as_u16() == 429 {
("rate_limit_error", "gateway_request_capacity_exceeded")
} else {
("invalid_request_error", "gateway_request_not_allowed")
};
send_responses_websocket_error(client_socket, status_code, error_type, code, message)
.await;
}
GatewayError::AdmissionTimeout { .. } => {
send_responses_websocket_error(
client_socket,
status_code,
"server_error",
"gateway_admission_timeout",
"Gateway capacity is busy; retry this response",
)
.await;
}
GatewayError::LocalExecutionPlanningTimeout { .. } => {
send_responses_websocket_error(
client_socket,
status_code,
"server_error",
"gateway_planning_timeout",
"Gateway planning timed out; retry this response",
)
.await;
}
_ => {
send_responses_websocket_error(
client_socket,
status_code,
"server_error",
"responses_websocket_turn_start_failed",
"Gateway could not start this response",
)
.await;
}
}
}
fn responses_websocket_turn_start_http_status(error: &GatewayError) -> u16 {
match error {
GatewayError::Client { status, .. } => status.as_u16(),
GatewayError::AdmissionTimeout { .. } => StatusCode::TOO_MANY_REQUESTS.as_u16(),
GatewayError::LocalExecutionPlanningTimeout { .. } => StatusCode::GATEWAY_TIMEOUT.as_u16(),
_ => StatusCode::INTERNAL_SERVER_ERROR.as_u16(),
}
}
pub(super) fn responses_websocket_turn_start_close(error: &GatewayError) -> (u16, &'static str) {
match error {
GatewayError::Client { .. } => (CLOSE_POLICY_VIOLATION, "request_not_allowed"),
GatewayError::AdmissionTimeout { .. }
| GatewayError::LocalExecutionPlanningTimeout { .. } => (CLOSE_TRY_AGAIN, "gateway_busy"),
_ => (CLOSE_INTERNAL_ERROR, "turn_start_failed"),
}
}
#[cfg(test)]
mod tests {
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::Duration;
use super::{
await_owned_turn_begin, await_owned_turn_begin_with_timeout,
await_turn_finalization_handle, responses_websocket_turn_start_close,
responses_websocket_turn_start_http_status, spawn_bounded_adapter_observation_with_timeout,
};
use crate::GatewayError;
#[test]
fn admission_timeout_uses_http_429_and_keeps_the_retry_later_close_code() {
let error = GatewayError::AdmissionTimeout {
trace_id: "turn-admission".to_string(),
gate: "gateway_upstream_execution",
queue_budget_ms: 25,
};
assert_eq!(responses_websocket_turn_start_http_status(&error), 429);
assert_eq!(
responses_websocket_turn_start_close(&error),
(1013, "gateway_busy")
);
}
/// C6 依赖的性质:结算是「等到落地」而不是「排进队列」。
///
/// 透明重试在这之后立刻按 health / adaptive / pool 状态规划下一个 attempt,
/// 所以结算任务必须已经跑完——只把 handle 挂起来是不够的。
#[tokio::test]
async fn awaiting_a_finalization_handle_runs_the_settlement_to_completion() {
let settled = Arc::new(AtomicBool::new(false));
let flag = Arc::clone(&settled);
let handle = tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(60)).await;
flag.store(true, Ordering::SeqCst);
});
assert!(
!settled.load(Ordering::SeqCst),
"the settlement has not finished yet"
);
await_turn_finalization_handle(handle).await;
assert!(
settled.load(Ordering::SeqCst),
"the settlement must be complete before the caller proceeds"
);
}
/// 顺序型:结算的每一步都要排在规划之前。
///
/// 用计数器替身重放透明重试的两步——旧 attempt 结算完成写入 1,规划开始时
/// 读到的必须已经是 1。旧实现在这里先规划、再把结算排进队列,规划读到的是 0。
#[tokio::test]
async fn transparent_retry_replans_only_after_the_previous_attempt_is_settled() {
let steps = Arc::new(AtomicUsize::new(0));
// 第一步:结算旧 attempt(等到落地)。
let recorder = Arc::clone(&steps);
let settlement = tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(40)).await;
recorder.store(1, Ordering::SeqCst);
});
await_turn_finalization_handle(settlement).await;
// 第二步:规划下一个 attempt,它读到的状态必须是结算之后的。
let observed_at_planning = steps.load(Ordering::SeqCst);
assert_eq!(
observed_at_planning, 1,
"planning must observe the state projected by the settled attempt"
);
}
/// 结算任务失败(panic / cancel)也必须让调用方继续,不能把 relay loop 卡死。
#[tokio::test]
async fn a_failed_finalization_task_still_releases_the_caller() {
let handle = tokio::spawn(async { panic!("settlement task exploded") });
await_turn_finalization_handle(handle).await;
}
#[tokio::test]
async fn cancelling_the_caller_does_not_cancel_turn_begin_or_drop_an_unowned_result() {
struct DropProbe(Arc<AtomicBool>);
impl Drop for DropProbe {
fn drop(&mut self) {
self.0.store(true, Ordering::SeqCst);
}
}
let begin_finished = Arc::new(AtomicBool::new(false));
let result_dropped = Arc::new(AtomicBool::new(false));
let finished = Arc::clone(&begin_finished);
let dropped = Arc::clone(&result_dropped);
let caller = tokio::spawn(async move {
await_owned_turn_begin(
async move {
tokio::time::sleep(Duration::from_millis(60)).await;
finished.store(true, Ordering::SeqCst);
Ok(DropProbe(dropped))
},
Duration::from_secs(1),
"turn-begin-cancel".to_string(),
)
.await
});
tokio::time::sleep(Duration::from_millis(10)).await;
caller.abort();
let _ = caller.await;
tokio::time::sleep(Duration::from_millis(120)).await;
assert!(
begin_finished.load(Ordering::SeqCst),
"the owned begin task must outlive its cancelled relay caller"
);
assert!(
result_dropped.load(Ordering::SeqCst),
"an undeliverable armed result must be dropped so its cleanup guard runs"
);
}
#[tokio::test]
async fn turn_begin_owner_deadline_drops_stalled_work_and_its_guards() {
struct DropProbe(Arc<AtomicBool>);
impl Drop for DropProbe {
fn drop(&mut self) {
self.0.store(true, Ordering::SeqCst);
}
}
let dropped = Arc::new(AtomicBool::new(false));
let task_dropped = Arc::clone(&dropped);
let result: Result<(), GatewayError> = await_owned_turn_begin_with_timeout(
async move {
let _probe = DropProbe(task_dropped);
std::future::pending::<()>().await;
Ok(())
},
Duration::from_millis(20),
"turn-begin-deadline".to_string(),
)
.await;
assert!(matches!(
result,
Err(GatewayError::LocalExecutionPlanningTimeout {
trace_id,
phase: "responses_websocket_turn_begin_owner",
timeout_ms: 20,
}) if trace_id == "turn-begin-deadline"
));
assert!(
dropped.load(Ordering::SeqCst),
"owner timeout must drop the stalled begin future so RAII cleanup runs"
);
}
#[tokio::test]
async fn cancelling_observation_waiter_cannot_bypass_the_owner_timeout() {
struct DropProbe(Arc<AtomicBool>);
impl Drop for DropProbe {
fn drop(&mut self) {
self.0.store(true, Ordering::SeqCst);
}
}
let dropped = Arc::new(AtomicBool::new(false));
let task_dropped = Arc::clone(&dropped);
let observation = async move {
let _probe = DropProbe(task_dropped);
std::future::pending::<()>().await;
};
let owner =
spawn_bounded_adapter_observation_with_timeout(observation, Duration::from_millis(20));
let waiter = tokio::spawn(async move {
let _ = owner.await;
});
waiter.abort();
let _ = waiter.await;
tokio::time::timeout(Duration::from_secs(1), async {
while !dropped.load(Ordering::SeqCst) {
tokio::task::yield_now().await;
}
})
.await
.expect("detached observation owner must enforce its own timeout");
}
}
@@ -0,0 +1,66 @@
//! OpenAI Responses WebSocket protocol entry point, session engine, and adapters.
//!
//! The route is protocol-oriented. `session` bootstraps the authenticated
//! connection, `connection` owns the socket FSM, `client` and `quota` own
//! protocol/retry policy, and `lifecycle`/`turn` bridge each turn into the
//! existing usage and audit runtime. Adapters contain only provider-specific
//! connection and metadata behavior.
mod adapter;
mod adapters;
mod admission;
mod binding;
mod client;
mod connection;
mod control;
mod frame;
mod lifecycle;
mod observation;
mod ownership;
mod quota;
mod redaction;
mod relay_policy;
mod request;
mod session;
mod settlement;
mod state;
mod turn;
mod turn_state;
mod upstream;
use std::net::SocketAddr;
use axum::body::Body;
use axum::extract::ws::WebSocketUpgrade;
use axum::extract::{ConnectInfo, State};
use axum::http::{HeaderMap, Response, Uri};
use crate::handlers::proxy::websocket::ingress::{
upgrade_authenticated_ai_websocket, WebSocketIngressSpec,
};
use crate::handlers::proxy::websocket::session::RESPONSES_WEBSOCKET_SESSION_LIMITS;
use crate::{AppState, GatewayError};
pub(crate) async fn responses_websocket(
State(state): State<AppState>,
ConnectInfo(remote_addr): ConnectInfo<SocketAddr>,
ws: WebSocketUpgrade,
headers: HeaderMap,
uri: Uri,
) -> Result<Response<Body>, GatewayError> {
upgrade_authenticated_ai_websocket(
state,
remote_addr,
ws,
headers,
uri,
RESPONSES_WEBSOCKET_SESSION_LIMITS,
RESPONSES_WEBSOCKET_INGRESS_SPEC,
session::run_responses_websocket,
)
.await
}
const RESPONSES_WEBSOCKET_INGRESS_SPEC: WebSocketIngressSpec = WebSocketIngressSpec {
route_unavailable_message: "WebSocket route is unavailable",
};
@@ -0,0 +1,220 @@
//! Responses WebSocket 的终态观测入口。
//!
//! 这条传输收到的本来就是结构化的 Responses 协议事件。之前为了复用面向 SSE 的
//! `push_line`,观测路径要先把每个事件序列化成 `data: {json}\n\n`,解析器再把它
//! 解码回 `Value`——一次纯粹的往返,而且这个「伪 SSE」形状是随手拼的,一旦
//! 上游事件里出现需要转义的内容,或者以后有人给拼装函数加了换行/分块逻辑,
//! 观测结果就会和真实事件悄悄分叉。
//!
//! 现在观测走 [`StreamingStandardTerminalObserver::push_event`],直接吃
//! `frame.protocol_events()` 借出的事件,不再序列化、不再解码。
//!
//! **body capture 不走这条路,仍然保持 SSE 形状**(`data: {json}\n\n`):
//! `aether_usage_runtime::report` 用 `line.strip_prefix("data:")` 解析被捕获的
//! body 来判定 `StreamCapturedTerminalState`,而它是 `stream_report_represents_failure`
//! 的一个 OR 项。把捕获内容换成结构化 JSON 会让终态判定恒为 Missing。
//! 也就是说这一层只换「观测」,不换「捕获」——见
//! [`super::turn::ResponsesProviderAttempt::capture_client_frame`] 一侧仍在用
//! SSE 编码。
use serde_json::Value;
use crate::ai_serving::api::StreamingStandardTerminalObserver;
use aether_contracts::ExecutionStreamTerminalSummary;
/// 包一层 [`StreamingStandardTerminalObserver`],只暴露结构化入口。
///
/// 存在的意义是让「WS 不再拼 SSE」成为类型层面的事实:这里没有任何接受字节的
/// 方法,所以不可能有人不小心把观测路径改回 `push_line`。
#[derive(Default)]
pub(super) struct ResponsesStructuredTerminalObserver {
inner: StreamingStandardTerminalObserver,
}
impl ResponsesStructuredTerminalObserver {
/// 观测一帧里的全部协议事件。
///
/// 第一个被拒绝的事件就停止推进并把摘要标成 parser_error:解析器的状态机是
/// 有顺序的,跳过一个事件继续喂后面的只会得到更没意义的摘要。
pub(super) fn observe_events(&mut self, report_context: &Value, events: &[&Value]) {
for event in events
.iter()
.copied()
.filter(|event| event_is_relevant_to_terminal_observation(event))
{
if let Err(error) = self.inner.push_event(report_context, event) {
self.inner.disable_with_error(error.to_string());
break;
}
}
}
pub(super) fn disable_with_error(&mut self, parser_error: impl Into<String>) {
self.inner.disable_with_error(parser_error);
}
pub(super) fn finish(&mut self, report_context: &Value) -> ExecutionStreamTerminalSummary {
match self.inner.finish(report_context) {
Ok(Some(summary)) => summary,
Ok(None) => ExecutionStreamTerminalSummary::default(),
Err(error) => {
self.inner.disable_with_error(error.to_string());
self.inner.latest_summary().cloned().unwrap_or_default()
}
}
}
}
/// The WebSocket relay is not a Responses schema gateway. It forwards all
/// events opaquely, while this observer consumes only identity/terminal
/// snapshots needed for usage and settlement. In particular, a future
/// `response.*` delta must not become an observation failure merely because
/// Aether's canonical streaming parser does not know it yet.
fn event_is_relevant_to_terminal_observation(event: &Value) -> bool {
matches!(
event.get("type").and_then(Value::as_str),
Some(
"response.created"
| "response.in_progress"
| "response.queued"
| "response.completed"
| "response.done"
| "response.failed"
| "response.incomplete"
| "response.cancelled"
| "error"
)
)
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::ResponsesStructuredTerminalObserver;
fn report_context() -> serde_json::Value {
json!({
"provider_api_format": "openai:responses",
"client_api_format": "openai:responses",
})
}
#[test]
fn structured_events_reach_the_terminal_summary_without_sse_text() {
let context = report_context();
let created = json!({"type": "response.created", "response": {"id": "resp_ws", "model": "gpt-5-codex"}});
let completed = json!({
"type": "response.completed",
"response": {
"id": "resp_ws",
"model": "gpt-5-codex",
"status": "completed",
"usage": {"input_tokens": 9, "output_tokens": 4, "total_tokens": 13},
},
});
let mut observer = ResponsesStructuredTerminalObserver::default();
observer.observe_events(&context, &[&created, &completed]);
let summary = observer.finish(&context);
assert!(summary.observed_finish);
assert_eq!(summary.response_id.as_deref(), Some("resp_ws"));
let usage = summary
.standardized_usage
.as_ref()
.expect("a completed response carries usage");
assert_eq!(usage.input_tokens, 9);
assert_eq!(usage.output_tokens, 4);
assert!(summary.parser_error.is_none());
}
/// 批量帧里的多个事件按顺序喂入,usage 不能因为批量而丢失。
#[test]
fn a_batched_frame_keeps_the_usage_of_its_last_event() {
let context = report_context();
let events = [
json!({"type": "response.created", "response": {"id": "resp_ws", "model": "m"}}),
json!({
"type": "response.output_text.delta",
"item_id": "msg",
"output_index": 0,
"content_index": 0,
"delta": "hi",
}),
json!({
"type": "response.completed",
"response": {
"id": "resp_ws",
"model": "m",
"status": "completed",
"usage": {"input_tokens": 3, "output_tokens": 1, "total_tokens": 4},
},
}),
];
let borrowed: Vec<&serde_json::Value> = events.iter().collect();
let mut observer = ResponsesStructuredTerminalObserver::default();
observer.observe_events(&context, &borrowed);
let summary = observer.finish(&context);
let usage = summary
.standardized_usage
.as_ref()
.expect("usage survives batching");
assert_eq!(usage.input_tokens, 3);
assert_eq!(usage.output_tokens, 1);
assert_eq!(usage.dimensions.get("total_tokens"), Some(&json!(4)));
}
#[test]
fn future_and_provider_private_events_are_ignored_only_by_the_side_observer() {
let context = report_context();
let private = json!({
"type": "codex.response.metadata",
"private_future_field": {"shape": "unknown"},
});
let future = json!({
"type": "response.future_capability.delta",
"future_capability": {"nested": [1, 2, 3]},
});
let completed = json!({
"type": "response.completed",
"response": {
"id": "resp_future",
"model": "future-model",
"status": "completed",
"future_response_field": {"also": "unknown"},
"usage": {"input_tokens": 5, "output_tokens": 2, "total_tokens": 7},
},
});
let mut observer = ResponsesStructuredTerminalObserver::default();
observer.observe_events(&context, &[&private, &future, &completed]);
let summary = observer.finish(&context);
assert!(summary.observed_finish);
assert_eq!(summary.response_id.as_deref(), Some("resp_future"));
assert_eq!(summary.unknown_event_count, 0);
assert_eq!(
summary
.standardized_usage
.as_ref()
.map(|usage| (usage.input_tokens, usage.output_tokens)),
Some((5, 2))
);
assert!(summary.parser_error.is_none());
}
#[test]
fn a_disabled_observer_reports_the_parser_error() {
let context = report_context();
let mut observer = ResponsesStructuredTerminalObserver::default();
observer.disable_with_error("upstream event was not valid JSON");
let summary = observer.finish(&context);
assert_eq!(
summary.parser_error.as_deref(),
Some("upstream event was not valid JSON")
);
}
}
@@ -0,0 +1,245 @@
//! Cancellation-safe ownership handoff for WebSocket planning leases.
//!
//! The relay races every turn against connection and response deadlines. A
//! planner future therefore cannot directly own a distributed pool-key lease:
//! losing the race would drop the future between scheduler selection and turn
//! startup, leaving that key unavailable until the lease TTL elapsed.
use std::collections::BTreeSet;
use std::time::Duration;
use serde_json::Value;
use tokio::task::JoinHandle;
use super::lifecycle::{begin_responses_websocket_turn, ActiveProviderAttempt};
use crate::ai_serving::{
maybe_build_responses_websocket_decision, AiExecutionDecision, GatewayAuthApiKeySnapshot,
ResponsesWebSocketDecision, ResponsesWebSocketPinnedCandidate,
};
use crate::control::GatewayControlDecision;
use crate::orchestration::release_pool_key_lease_from_report_context;
use crate::{AppState, GatewayError};
/// Owns a selected pool-key lease until the attempt lifecycle has taken over
/// the decision report context.
pub(super) struct PlannedPoolKeyLeaseGuard {
state: AppState,
report_context: Option<Value>,
}
/// Planner output coupled to both its request parts and lease guard.
pub(super) struct OwnedResponsesWebSocketDecision {
pub(super) planned: ResponsesWebSocketDecision,
pub(super) planning_parts: http::request::Parts,
pub(super) planned_lease: PlannedPoolKeyLeaseGuard,
}
/// Runs planning in an owner task. Dropping the caller's waiter detaches this
/// task; an unobserved successful output drops its guard and releases the
/// selected pool-key lease.
#[allow(clippy::too_many_arguments)]
pub(super) fn spawn_owned_responses_websocket_plan(
state: AppState,
parts: http::request::Parts,
trace_id: String,
control_decision: GatewayControlDecision,
auth_snapshot: Option<GatewayAuthApiKeySnapshot>,
client_event: Value,
excluded_key_ids: Option<BTreeSet<String>>,
excluded_codex_account_ids: Option<BTreeSet<String>>,
pinned_candidate: Option<ResponsesWebSocketPinnedCandidate>,
) -> JoinHandle<Result<Option<OwnedResponsesWebSocketDecision>, GatewayError>> {
let owner_timeout = state
.frontdoor_runtime_guards
.local_execution_planning_timeout;
tokio::spawn(async move {
let planned = await_owned_planning_deadline(
maybe_build_responses_websocket_decision(
&state,
&parts,
&trace_id,
&control_decision,
auth_snapshot.as_ref(),
&client_event,
excluded_key_ids.as_ref(),
excluded_codex_account_ids.as_ref(),
pinned_candidate.as_ref(),
),
owner_timeout,
)
.await
.map_err(|_| GatewayError::LocalExecutionPlanningTimeout {
trace_id: trace_id.clone(),
phase: "responses_websocket_plan_owner",
timeout_ms: owner_timeout.as_millis() as u64,
})??;
Ok(planned.map(|planned| {
let planned_lease =
PlannedPoolKeyLeaseGuard::new(&state, planned.execution.report_context.as_ref());
OwnedResponsesWebSocketDecision {
planned,
planning_parts: parts,
planned_lease,
}
}))
})
}
async fn await_owned_planning_deadline<F, T>(
planning: F,
deadline: Duration,
) -> Result<T, tokio::time::error::Elapsed>
where
F: std::future::Future<Output = T>,
{
tokio::time::timeout(deadline, planning).await
}
pub(super) async fn await_owned_responses_websocket_plan(
handle: JoinHandle<Result<Option<OwnedResponsesWebSocketDecision>, GatewayError>>,
) -> Result<Option<OwnedResponsesWebSocketDecision>, GatewayError> {
handle.await.map_err(|error| {
GatewayError::Internal(format!(
"Responses WebSocket planning task failed before ownership transfer: {error}"
))
})?
}
impl PlannedPoolKeyLeaseGuard {
fn new(state: &AppState, report_context: Option<&Value>) -> Self {
Self {
state: state.clone(),
report_context: report_context.cloned(),
}
}
pub(super) async fn release(mut self) {
release_pool_key_lease_from_report_context(&self.state, self.report_context.as_ref()).await;
self.report_context = None;
}
fn disarm(&mut self) {
self.report_context = None;
}
}
impl Drop for PlannedPoolKeyLeaseGuard {
fn drop(&mut self) {
let Some(report_context) = self.report_context.take() else {
return;
};
let state = self.state.clone();
if let Ok(handle) = tokio::runtime::Handle::try_current() {
handle.spawn(async move {
release_pool_key_lease_from_report_context(&state, Some(&report_context)).await;
});
}
}
}
/// Keeps the planning guard in the same detached owner task as lifecycle
/// startup. If the relay loses a deadline race while awaiting startup, the
/// task completes the handoff (or releases the lease on failure) without a
/// cancellation gap.
pub(super) async fn begin_responses_websocket_turn_with_planned_lease(
state: &AppState,
trace_id: &str,
parts: http::request::Parts,
control_decision: &GatewayControlDecision,
decision: AiExecutionDecision,
client_event: &Value,
mut planned_lease: PlannedPoolKeyLeaseGuard,
) -> Result<ActiveProviderAttempt, GatewayError> {
let state = state.clone();
let trace_id = trace_id.to_string();
let control_decision = control_decision.clone();
let client_event = client_event.clone();
tokio::spawn(async move {
let turn = begin_responses_websocket_turn(
&state,
&trace_id,
parts,
&control_decision,
decision,
&client_event,
)
.await?;
// ActiveProviderAttempt now owns the report context containing the
// lease. No await occurs between that handoff and disarming the guard.
planned_lease.disarm();
Ok(turn)
})
.await
.map_err(|error| {
GatewayError::Internal(format!(
"Responses WebSocket guarded turn startup task failed: {error}"
))
})?
}
#[cfg(test)]
mod tests {
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::Duration;
struct DropProbe(Arc<AtomicUsize>);
impl Drop for DropProbe {
fn drop(&mut self) {
self.0.fetch_add(1, Ordering::SeqCst);
}
}
#[tokio::test]
async fn dropping_a_planning_waiter_detaches_the_owner_and_drops_its_output() {
let started = Arc::new(tokio::sync::Notify::new());
let release = Arc::new(tokio::sync::Notify::new());
let dropped = Arc::new(AtomicUsize::new(0));
let task_started = Arc::clone(&started);
let task_release = Arc::clone(&release);
let task_dropped = Arc::clone(&dropped);
let owner = tokio::spawn(async move {
task_started.notify_one();
task_release.notified().await;
DropProbe(task_dropped)
});
started.notified().await;
let waiter = tokio::spawn(async move {
let _ = owner.await;
});
waiter.abort();
let _ = waiter.await;
release.notify_one();
tokio::time::timeout(Duration::from_secs(1), async {
while dropped.load(Ordering::SeqCst) == 0 {
tokio::task::yield_now().await;
}
})
.await
.expect("detached owner output should be dropped after it finishes");
assert_eq!(dropped.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn planning_owner_deadline_drops_stalled_work_and_its_guards() {
let dropped = Arc::new(AtomicUsize::new(0));
let task_dropped = Arc::clone(&dropped);
let planning = async move {
let _probe = DropProbe(task_dropped);
std::future::pending::<()>().await;
};
let result =
super::await_owned_planning_deadline(planning, Duration::from_millis(20)).await;
assert!(
result.is_err(),
"stalled planning must hit its owner deadline"
);
assert_eq!(dropped.load(Ordering::SeqCst), 1);
}
}
@@ -0,0 +1,377 @@
//! Quota exhaustion, replay safety, and upstream replacement policy.
use serde_json::Value;
use uuid::Uuid;
use wreq::ws::message::Message as WreqWsMessage;
use super::adapter::{
resolve_responses_websocket_adapter, ResponsesWebSocketDrainDirective,
ResponsesWebSocketRebindSafety,
};
use super::lifecycle::{queue_turn_finalization, PreviousAttemptSettled};
use super::ownership::{
await_owned_responses_websocket_plan, begin_responses_websocket_turn_with_planned_lease,
spawn_owned_responses_websocket_plan, OwnedResponsesWebSocketDecision,
};
use super::request::{build_planning_parts, planned_response_create_event};
use super::state::BoundResponsesConnection;
use super::turn::{prepare_responses_websocket_turn_decision, ResponsesWebSocketTurnOutcome};
use super::upstream::{bind_responses_upstream, close_bound_upstream};
use crate::clock::current_unix_secs;
use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext;
use crate::handlers::proxy::websocket::session::WEBSOCKET_LOG_TRANSPORT;
use crate::handlers::proxy::websocket::transport::close_upstream_socket;
use crate::AppState;
const LOG_TARGET: &str = "aether_gateway::handlers::proxy::responses_ws";
macro_rules! debug {
($($arg:tt)*) => {
tracing::debug!(target: LOG_TARGET, $($arg)*)
};
}
macro_rules! warn {
($($arg:tt)*) => {
tracing::warn!(target: LOG_TARGET, $($arg)*)
};
}
pub(super) async fn detach_exhausted_upstream(
bound: &mut BoundResponsesConnection,
directive: ResponsesWebSocketDrainDirective,
trace_id: &str,
) {
let exclusion = record_exhausted_bound_key(bound, directive.retry_exclusion_until_unix_secs);
close_bound_upstream(bound).await;
// 调用方必须先结束当前 logical turn 再 detach:拆掉上游后 attempt 已经不可能
// 收到终态,留着它只会等 deadline 或 drop guard 兜底。
debug_assert!(
!bound.turn_state.response_in_flight(),
"an exhausted upstream must be detached after its logical turn ended"
);
bound.pending_adapter_drain = None;
let now_unix_secs = current_unix_secs();
debug!(
event_name = "responses_websocket_upstream_detached",
log_type = "event",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %trace_id,
reason = directive.error_code,
exhausted_key_id = ?exclusion.as_ref().map(|(key_id, _)| key_id),
retry_exclusion_until_unix_secs = ?exclusion.as_ref().map(|(_, until)| until),
exhausted_exclusion_count = bound.exhausted_exclusions.len(now_unix_secs),
"gateway detached an exhausted Responses WebSocket upstream while preserving the client socket"
);
}
pub(super) fn record_exhausted_bound_key(
bound: &mut BoundResponsesConnection,
reset_at_unix_secs: Option<u64>,
) -> Option<(String, u64)> {
let key_id = bound
.decision_template
.key_id
.as_deref()
.map(str::trim)
.filter(|key_id| !key_id.is_empty())?
.to_string();
let provider_account_id = bound
.adapter
.exhaustion_exclusion_identity(&bound.decision_template)
.and_then(|identity| identity.account_id);
let exclusion_until = bound.exhausted_exclusions.exclude(
key_id.clone(),
provider_account_id,
reset_at_unix_secs,
current_unix_secs(),
);
Some((key_id, exclusion_until))
}
/// 为同一个 logical turn 规划并绑定下一个 attempt。
///
/// `_previous_settled` 不被使用,它只是把「上一个 attempt 已经结算完毕」这个
/// 前置条件写进签名:规划要读 health / adaptive / pool 状态,而这些是上一个
/// attempt 结算时才投射的;它的 pool key lease 也要先释放,否则替代 key 的挑选
/// 会看到一把仍被占用的 key。
pub(super) async fn retry_active_turn_after_quota_exhaustion(
bound: &mut BoundResponsesConnection,
state: &AppState,
context: &WebSocketRequestContext,
_previous_settled: PreviousAttemptSettled,
) -> bool {
let Some(active) = bound.turn_state.logical_mut() else {
return false;
};
if let Some(reason) = active.quota_retry_block_reason() {
debug!(
event_name = "responses_websocket_quota_retry_skipped",
log_type = "event",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
turn_index = active.turn_index,
logical_turn_id = %active.logical_turn_id,
turn_attempt = active.turn_attempt,
reason,
"gateway will not transparently replay an unsafe Responses WebSocket turn"
);
return false;
}
active.retry_attempted = true;
active.turn_attempt = active.turn_attempt.saturating_add(1);
let client_event = active.client_event.clone();
let Some(turn_control) = active.turn_control.clone() else {
warn!(
event_name = "responses_websocket_quota_retry_control_missing",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
"gateway refused to retry a WebSocket turn without its live authorization snapshot"
);
return false;
};
let turn_index = active.turn_index;
let logical_turn_id = active.logical_turn_id.clone();
let turn_attempt = active.turn_attempt;
let retry_exclusion_until_unix_secs = bound
.pending_adapter_drain
.and_then(|directive| directive.retry_exclusion_until_unix_secs);
let exhausted_key = record_exhausted_bound_key(bound, retry_exclusion_until_unix_secs);
let exhausted_key_id = exhausted_key.as_ref().map(|(key_id, _)| key_id.clone());
let planning_parts = build_planning_parts(context);
let turn_request_id = Uuid::new_v4().to_string();
let now_unix_secs = current_unix_secs();
let excluded_key_ids = bound.exhausted_exclusions.key_ids(now_unix_secs);
let excluded_codex_account_ids = bound.exhausted_exclusions.codex_account_ids(now_unix_secs);
let excluded_key_ids = (!excluded_key_ids.is_empty()).then_some(excluded_key_ids);
let excluded_codex_account_ids =
(!excluded_codex_account_ids.is_empty()).then_some(excluded_codex_account_ids);
let planned = match await_owned_responses_websocket_plan(spawn_owned_responses_websocket_plan(
state.clone(),
planning_parts,
turn_request_id.clone(),
turn_control.decision.clone(),
turn_control.auth_snapshot.clone(),
client_event.clone(),
excluded_key_ids,
excluded_codex_account_ids,
None,
))
.await
{
Ok(Some(decision)) => decision,
Ok(None) => {
warn!(
event_name = "responses_websocket_quota_retry_provider_unavailable",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
exhausted_key_id = ?exhausted_key_id,
"gateway could not find an alternate Responses WebSocket provider after quota exhaustion"
);
return false;
}
Err(error) => {
warn!(
event_name = "responses_websocket_quota_retry_planning_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
exhausted_key_id = ?exhausted_key_id,
error = ?error,
"gateway could not plan an alternate Responses WebSocket provider after quota exhaustion"
);
return false;
}
};
let OwnedResponsesWebSocketDecision {
planned,
planning_parts,
planned_lease,
} = planned;
let adapter = resolve_responses_websocket_adapter(planned.adapter);
let normalization = planned.normalization;
let decision = planned.execution;
if exhausted_key_id.as_deref() == decision.key_id.as_deref() {
planned_lease.release().await;
warn!(
event_name = "responses_websocket_quota_retry_selected_exhausted_key",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
key_id = ?decision.key_id,
"gateway rejected an alternate Responses WebSocket plan that reused the exhausted key"
);
return false;
}
let provider_event = match planned_response_create_event(&decision, &client_event).and_then(
|event| {
serde_json::from_str::<Value>(&event)
.map_err(|_| "response_create_serialization_failed")
},
) {
Ok(event) => event,
Err(code) => {
planned_lease.release().await;
warn!(
event_name = "responses_websocket_quota_retry_normalization_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
error_code = code,
"gateway could not rebuild a Responses response.create for transparent quota retry"
);
return false;
}
};
let turn_decision = prepare_responses_websocket_turn_decision(
&decision,
turn_request_id,
true,
&client_event,
&provider_event,
&context.trace_id,
turn_index,
&logical_turn_id,
turn_attempt,
);
let mut turn = match begin_responses_websocket_turn_with_planned_lease(
state,
&context.trace_id,
planning_parts,
&turn_control.decision,
turn_decision,
&client_event,
planned_lease,
)
.await
{
Ok(turn) => turn,
Err(error) => {
warn!(
event_name = "responses_websocket_quota_retry_reporting_unavailable",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
error = ?error,
"gateway could not start usage and audit tracking for transparent quota retry"
);
return false;
}
};
let mut replacement = match bind_responses_upstream(
&decision,
normalization,
&client_event,
adapter,
)
.await
{
Ok(connection) => connection,
Err(code) => {
queue_turn_finalization(
bound,
state,
turn,
ResponsesWebSocketTurnOutcome::upstream_connect_failed(code),
)
.await;
warn!(
event_name = "responses_websocket_quota_retry_rebind_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
error_code = code,
"gateway could not bind an alternate Responses WebSocket provider after quota exhaustion"
);
return false;
}
};
turn.mark_upstream_request_sent();
turn.set_provider_response_headers(replacement.upstream_response_headers.clone());
let replacement_upstream = replacement
.upstream
.take()
.expect("newly bound Responses upstream should be present");
if let Some(mut previous_upstream) = bound.upstream.replace(replacement_upstream) {
close_upstream_socket(&mut previous_upstream, None).await;
}
let previous_key_id = bound.decision_template.key_id.clone();
bound.adapter = replacement.adapter;
bound.client_model = replacement.client_model;
bound.provider_model = replacement.provider_model;
bound.decision_template = replacement.decision_template;
bound.body_normalization = replacement.body_normalization;
bound.binding_identity = replacement.binding_identity;
// 同一个 logical turn 的下一个 attempt 就位。状态不符时把 attempt 交回
// drop guard 结算并让调用方走「透明重试失败」分支,不静默丢弃一条已经写了
// pending usage 行、占着 candidate 和 pool key lease 的 attempt。
if let Err(orphan) = bound.turn_state.resume(turn) {
drop(orphan);
return false;
}
bound.upstream_response_headers = replacement.upstream_response_headers;
bound.pending_adapter_drain = None;
debug!(
event_name = "responses_websocket_quota_retry_rebound",
log_type = "event",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
turn_index,
logical_turn_id = %logical_turn_id,
turn_attempt,
previous_key_id = ?previous_key_id,
key_id = ?bound.decision_template.key_id,
"gateway transparently rebound a Responses WebSocket turn after quota exhaustion"
);
true
}
pub(super) fn is_usage_limit_error_event(event: &Value) -> bool {
let is_error = |value: &Value| {
value.get("type").and_then(Value::as_str) == Some("error")
&& value.pointer("/error/type").and_then(Value::as_str) == Some("usage_limit_reached")
};
is_error(event)
|| event
.get("chunks")
.and_then(Value::as_array)
.is_some_and(|chunks| chunks.iter().any(is_error))
}
pub(super) fn observe_active_response_rebind_safety(
bound: &mut BoundResponsesConnection,
event: &Value,
) {
let ResponsesWebSocketRebindSafety::Unsafe { reason } =
bound.adapter.rebind_safety_for_upstream_event(event)
else {
return;
};
if let Some(active) = bound.turn_state.logical_mut() {
active.mark_retry_unsafe(reason);
}
}
pub(super) fn mark_active_response_retry_unsafe(
bound: &mut BoundResponsesConnection,
reason: &'static str,
) {
if let Some(active) = bound.turn_state.logical_mut() {
active.mark_retry_unsafe(reason);
}
}
@@ -0,0 +1,850 @@
//! Responses WebSocket 两侧的 PII 脱敏:请求侧 mask + 响应侧 restore。
//!
//! HTTP 路径在前门建 `RedactionSessionSlot` 并塞进 `parts.extensions`,planner
//! 只有拿到这个 slot 才会脱敏。WS 的 planning Parts 是合成的:四个规划入口
//! (首轮、换模型 re-plan、独立轮、配额透明重试)靠 `build_planning_parts` 注入
//! slot 就能复用 planner 的脱敏;但复用已绑定 upstream 的 continuation 根本不进
//! planner,必须在这里先把客户端事件脱敏,再交给协议归一化、上游发送和审计。
//!
//! 因此约定:**进入任何下游用途之前,客户端 `response.create` 只在这里脱敏一次**,
//! 之后所有路径都只看脱敏后的事件。
//!
//! # 响应侧
//!
//! 只 mask 不 restore 是半个实现:HTTP 在把响应交给客户端之前会把占位符换回真实值
//! (`privacy::restore_sync_response_body` / `privacy::StreamingResponseRestorer`),
//! WS 少了这一步,客户端就会直接看到 `<AETHER:EMAIL:...>`。
//! [`ResponsesWebSocketRedactionRestorer`] 补上这一跳,语义与 HTTP 完全一致:
//! 复用 `privacy::restore_json_strings`,只还原本连接自己 mask 出来的映射,
//! 未映射的占位符原样透传。
//!
//! ## session 为什么活在连接上而不是活在这一轮里
//!
//! mask session 由 planner 写进 per-turn 的 slot,而 slot 随 planning Parts 在
//! 规划结束时就被丢弃,响应帧到达时已经无处可取。可选的存活范围有两个:
//!
//! * 挂在 `LogicalTurn` 上:这一轮结束即释放,是 HTTP「一个请求一个 session」的
//! 直译。但 WS 的会话历史留在上游:continuation 只发增量输入,第 1 轮的
//! `input` 不会在第 3 轮重发。于是第 3 轮的响应里若回显了第 1 轮的占位符
//! ("你刚才给我的邮箱是……"),本轮 session 里没有这条映射,占位符就漏给客户端。
//! HTTP 不会漏,是因为它每次都重发整段历史,重新 mask 同一个值会派生出同一个
//! sentinel(HMAC over 规则 + bucket + 值),所以映射天然齐备。
//! * 挂在连接上(当前实现):每轮仍然各自 mask、各自持有独立 session
//! (per-turn 语义不变),连接只是把最近若干轮的 session 留下来一起参与还原,
//! 凑出的映射集合正好等于「等价 HTTP 请求会拥有的那一份」。
//!
//! 选后者。代价是每帧最多对 [`MAX_RETAINED_TURN_REDACTION_SESSIONS`] 个 session
//! 各扫一遍,以及这些 session 的映射会驻留到连接结束;用有界 FIFO 兜住上限。
//! 窗口不够用或每帧成本变高时,正确的下一步是在 `privacy` 侧提供跨 session 的
//! 合并匹配器,而不是把这个窗口调大。
use std::collections::VecDeque;
use serde_json::Value;
use crate::ai_serving::{
resolve_local_decision_execution_runtime_auth_context, resolve_provider_chat_pii_redaction,
};
use crate::control::GatewayControlDecision;
use crate::privacy::{restore_json_strings, RedactionSession, RedactionSessionSlot};
use crate::{AppState, GatewayError};
/// Responses WebSocket 只承载 `openai:responses`,脱敏规则按这个客户端格式选取。
const RESPONSES_WEBSOCKET_CLIENT_API_FORMAT: &str = "openai:responses";
/// WS 在选出候选之前就要脱敏,所以脱敏 session 先记在这个固定 key 下。
///
/// slot 是 per-turn 的(见 `build_planning_parts`),这一轮之后即随 slot 一起丢弃;
/// planner 后续用真实 candidate_id 再取一次配置时,body 已是脱敏态、不会重复写入。
const WEBSOCKET_TURN_REDACTION_CANDIDATE_ID: &str = "responses_websocket_turn";
/// 一条连接最多留几轮的 mask session 用于响应侧还原。
///
/// 取值权衡见模块文档:调大会线性增加每帧还原成本和常驻映射量,调小则更容易漏还原
/// 上游历史里更早那几轮的占位符。8 覆盖的是「上游最可能回显的最近窗口」。
const MAX_RETAINED_TURN_REDACTION_SESSIONS: usize = 8;
/// 一轮客户端 `response.create` 的请求侧脱敏结果。
#[derive(Debug)]
pub(super) struct ResponsesWebSocketTurnRedaction {
/// 脱敏后的客户端事件;这一轮之后所有下游路径都只看它。
pub(super) client_event: Value,
/// 这一轮 mask 出来的映射表,响应侧还原只能靠它。
pub(super) session: RedactionSession,
}
/// 对一条客户端 `response.create` 做请求侧脱敏。
///
/// 返回 `Some(..)` 仅当脱敏真正命中;`None` 表示未启用或没有命中,调用方
/// 继续用原事件即可(避免未开启脱敏时多一次整包 clone)。
///
/// 脱敏只改写 `instructions` / `input`(见 `privacy::mask_openai_responses_request_value`),
/// `type` / `model` / `previous_response_id` / `generate` 等协议字段原样保留,所以脱敏后的
/// 事件仍可直接用于协议归一化和上游发送。
///
/// 出错必须让这一轮失败:脱敏已启用却读不到配置或加密密钥时,把原文发上游就是
/// 静默旁路,正是本次要修的问题。
pub(super) async fn redact_responses_websocket_client_event(
state: &AppState,
parts: &http::request::Parts,
control_decision: &GatewayControlDecision,
client_event: &Value,
) -> Result<Option<ResponsesWebSocketTurnRedaction>, GatewayError> {
let Some(auth_context) =
resolve_local_decision_execution_runtime_auth_context(control_decision)
else {
return Ok(None);
};
let redaction = resolve_provider_chat_pii_redaction(
state,
parts,
client_event,
&auth_context,
RESPONSES_WEBSOCKET_CLIENT_API_FORMAT,
WEBSOCKET_TURN_REDACTION_CANDIDATE_ID,
)
.await?;
if !redaction.redacted {
return Ok(None);
}
// mask 命中时 `resolve_provider_chat_pii_redaction` 必定把 session 写进 slot。
// 取不到就是内部契约被破坏了,此时继续下发意味着这一轮的响应无法还原、占位符
// 会漏给客户端;按本模块既有的「脱敏链路出错就让这一轮失败」处理,不做降级。
let Some(session) = parts
.extensions
.get::<RedactionSessionSlot>()
.and_then(|slot| slot.take_for_candidate(Some(WEBSOCKET_TURN_REDACTION_CANDIDATE_ID)))
else {
return Err(GatewayError::Internal(
"chat pii redaction masked a Responses WebSocket turn without retaining its session"
.to_string(),
));
};
Ok(Some(ResponsesWebSocketTurnRedaction {
client_event: redaction.body_json.into_owned(),
session,
}))
}
/// 一条连接上「我们 mask 过哪些映射」的留存集合,供响应侧还原使用。
///
/// 每轮一个独立 session(per-turn mask 语义不变),连接按 FIFO 留最近
/// [`MAX_RETAINED_TURN_REDACTION_SESSIONS`] 轮。上游重绑不清空:客户端仍在同一段
/// 对话里,旧占位符可能随重发的输入再次出现。
#[derive(Default)]
pub(super) struct ResponsesWebSocketRedactionRestorer {
sessions: VecDeque<RedactionSession>,
}
impl ResponsesWebSocketRedactionRestorer {
/// 登记这一轮的 mask session。
pub(super) fn register(&mut self, session: RedactionSession) {
if session.mapping_count() == 0 {
return;
}
self.sessions.push_back(session);
while self.sessions.len() > MAX_RETAINED_TURN_REDACTION_SESSIONS {
self.sessions.pop_front();
}
}
/// 把一帧 provider 事件里的占位符换回真实值,返回要发给客户端的帧文本。
///
/// `None` 表示这一帧没有任何东西要还原,调用方必须原样转发上游字节:未启用
/// 脱敏(没有任何 session)时连 clone 都不做。
///
/// 入参只读:审计与终态观测继续消费脱敏态的事件,还原只作用于发往客户端的
/// 那一份拷贝,和 HTTP 侧「审计存脱敏体、线上还原」保持一致。
pub(super) fn restore_provider_frame_text(&self, event: &Value) -> Option<String> {
if self.sessions.is_empty() {
return None;
}
let mut restored_event = event.clone();
let mut restored = false;
for session in &self.sessions {
// 逐 session 还原而不是合并映射:每个 session 只认自己 mask 过的
// sentinel(`RedactionSession::restore_text`),跨 session 合并会绕开
// 这条边界。同一个值在不同轮派生出的 sentinel 相同,所以顺序无关。
restored |= restore_json_strings(&mut restored_event, session);
}
if !restored {
return None;
}
// 刚从 JSON 解析出来的 Value 再序列化不会失败;真失败时宁可让客户端看到
// 占位符,也不能丢掉这一帧——丢帧会让客户端的协议状态机卡死。
serde_json::to_string(&restored_event).ok()
}
}
#[cfg(test)]
mod tests {
use std::net::SocketAddr;
use std::sync::Arc;
use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY;
use aether_data::repository::auth::{
InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeyExportRecord,
};
use axum::http::{HeaderMap, Uri};
use serde_json::{json, Value};
use super::super::request::{
build_planning_parts, normalize_followup_response_create, planned_response_create_event,
};
use super::super::turn::prepare_responses_websocket_turn_decision;
use super::super::turn_state::LogicalTurn;
use super::{
redact_responses_websocket_client_event, ResponsesWebSocketRedactionRestorer,
ResponsesWebSocketTurnRedaction, MAX_RETAINED_TURN_REDACTION_SESSIONS,
};
use crate::ai_serving::{AiExecutionDecision, ResponsesWebSocketBodyNormalization};
use crate::control::{GatewayControlAuthContext, GatewayControlDecision};
use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext;
use crate::AppState;
const TEST_USER_ID: &str = "user-responses-ws-redaction";
const TEST_API_KEY_ID: &str = "api-key-responses-ws-redaction";
const TEST_EMAIL: &str = "[email protected]";
/// 另一轮用的 PII,用来证明连接级还原覆盖到更早的轮次。
const OTHER_TEST_EMAIL: &str = "[email protected]";
/// 不是本连接 mask 出来的占位符:格式合法(符合 sentinel 正则),但没有任何
/// session 记过它,必须原样透传。
const FOREIGN_SENTINEL: &str = "<AETHER:EMAIL:AAAAAAAAAAAAAAAAAAAA>";
fn auth_export_record() -> StoredAuthApiKeyExportRecord {
StoredAuthApiKeyExportRecord::new(
TEST_USER_ID.to_string(),
TEST_API_KEY_ID.to_string(),
"hash-responses-ws-redaction".to_string(),
None,
Some("ws".to_string()),
None,
None,
None,
None,
None,
None,
true,
None,
false,
0,
0,
0.0,
false,
)
.expect("auth api key export record should build")
.with_feature_settings(Some(json!({
"chat_pii_redaction": {"enabled": true}
})))
}
/// 只装脱敏真正需要的东西:系统配置开关 + 规则、加密密钥、带 feature settings
/// 的 API Key 导出记录。候选/上游都不需要,这条链路在 planner 之前。
fn redaction_enabled_state() -> AppState {
let auth_repository = Arc::new(
InMemoryAuthApiKeySnapshotRepository::seed(vec![])
.with_export_records(vec![auth_export_record()]),
);
let data_state =
crate::data::GatewayDataState::with_auth_api_key_reader_for_tests(auth_repository)
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY)
.with_system_config_values_for_tests(vec![
("module.chat_pii_redaction.enabled".to_string(), json!(true)),
(
"module.chat_pii_redaction.rules".to_string(),
json!([{
"id": "email",
"name": "邮箱",
"pattern": r"(?i)[A-Z0-9._%+-]{1,64}@[A-Z0-9.-]{1,253}\.[A-Z]{2,63}",
"enabled": true,
"features": {"validator": "email"},
"system": true
}]),
),
(
"module.chat_pii_redaction.cache_ttl_seconds".to_string(),
json!(300),
),
]);
AppState::new()
.expect("gateway state should build")
.with_data_state_for_tests(data_state)
}
fn control_decision() -> GatewayControlDecision {
let mut decision = GatewayControlDecision::synthetic(
"/v1/responses".to_string(),
Some("ai_public".to_string()),
Some("openai".to_string()),
Some("responses_websocket".to_string()),
Some("openai:responses".to_string()),
);
decision.auth_context = Some(GatewayControlAuthContext {
user_id: TEST_USER_ID.to_string(),
api_key_id: TEST_API_KEY_ID.to_string(),
username: Some("ws".to_string()),
api_key_name: Some("ws".to_string()),
balance_remaining: None,
access_allowed: true,
user_rate_limit: None,
api_key_rate_limit: None,
api_key_is_standalone: false,
admin_bypass_limits: false,
local_rejection: None,
allowed_models: None,
ip_rules: None,
});
decision
}
fn websocket_context(decision: GatewayControlDecision) -> WebSocketRequestContext {
WebSocketRequestContext {
trace_id: "trace-responses-ws-redaction".to_string(),
headers: HeaderMap::new(),
uri: Uri::from_static("/v1/responses"),
remote_addr: "127.0.0.1:65000"
.parse::<SocketAddr>()
.expect("remote address should parse"),
client_ip: "127.0.0.1".parse().expect("client IP should parse"),
decision,
websocket_connection_permit: None,
}
}
fn client_event() -> Value {
client_event_with_email(TEST_EMAIL)
}
fn client_event_with_email(email: &str) -> Value {
json!({
"type": "response.create",
"model": "public-model",
"previous_response_id": "resp-previous",
"generate": false,
"input": [{
"role": "user",
"content": [{"type": "input_text", "text": format!("mail {email}")}]
}]
})
}
/// 真跑一遍请求侧脱敏,拿到这一轮的生效事件和 mask session。
async fn turn_redaction(
state: &AppState,
decision: &GatewayControlDecision,
email: &str,
) -> ResponsesWebSocketTurnRedaction {
let context = websocket_context(decision.clone());
let parts = build_planning_parts(&context);
let event = client_event_with_email(email);
redact_responses_websocket_client_event(state, &parts, &context.decision, &event)
.await
.expect("redaction should resolve")
.expect("an email in the request should be redacted")
}
/// 这一轮为 `email` 派生出的占位符。
fn sentinel_for(redaction: &ResponsesWebSocketTurnRedaction, email: &str) -> String {
redaction
.session
.sentinel_for_original(email)
.expect("a masked email must have a sentinel")
.to_string()
}
/// 上游回显占位符的一帧 provider 事件。
fn provider_delta_frame(text: &str) -> Value {
json!({
"type": "response.output_text.delta",
"item_id": "msg_ws",
"output_index": 0,
"content_index": 0,
"delta": text,
})
}
#[tokio::test]
async fn websocket_client_event_is_redacted_without_losing_protocol_fields() {
let state = redaction_enabled_state();
let context = websocket_context(control_decision());
let parts = build_planning_parts(&context);
let event = client_event();
let redacted =
redact_responses_websocket_client_event(&state, &parts, &context.decision, &event)
.await
.expect("redaction should resolve")
.expect("an email in the request should be redacted")
.client_event;
let serialized = serde_json::to_string(&redacted).expect("event should serialize");
assert!(!serialized.contains(TEST_EMAIL), "{serialized}");
assert!(serialized.contains("<AETHER:EMAIL:"), "{serialized}");
// 协议字段必须原样保留,否则 continuation 链路会断。
assert_eq!(redacted["type"], "response.create");
assert_eq!(redacted["model"], "public-model");
assert_eq!(redacted["previous_response_id"], "resp-previous");
assert_eq!(redacted["generate"], false);
}
#[tokio::test]
async fn redacting_an_already_redacted_event_is_a_no_op() {
// re-plan 与配额重试路径会把已脱敏的事件再交给 planner,planner 内部会对
// 同一个 body 再跑一遍 mask。占位符本身不该被任何规则命中,否则会被二次
// 替换、破坏与上游已有 previous_response_id 链的一致性。
let state = redaction_enabled_state();
let context = websocket_context(control_decision());
let parts = build_planning_parts(&context);
let event = client_event();
let redacted =
redact_responses_websocket_client_event(&state, &parts, &context.decision, &event)
.await
.expect("redaction should resolve")
.expect("an email in the request should be redacted")
.client_event;
// 复用同一个 parts/slot,和 re-plan 在同一 turn 内二次脱敏的情形一致。
let second_pass =
redact_responses_websocket_client_event(&state, &parts, &context.decision, &redacted)
.await
.expect("second redaction pass should resolve");
assert!(
second_pass.is_none(),
"already redacted event should stay byte-identical: {second_pass:?}"
);
}
#[tokio::test]
async fn redaction_is_skipped_without_a_local_auth_context() {
let state = redaction_enabled_state();
let mut decision = control_decision();
decision.auth_context = None;
let context = websocket_context(decision);
let parts = build_planning_parts(&context);
let event = client_event();
let redacted =
redact_responses_websocket_client_event(&state, &parts, &context.decision, &event)
.await
.expect("redaction should resolve");
assert!(redacted.is_none());
}
/// 真跑一遍脱敏,拿到这一轮的「生效事件」。
async fn redacted_client_event(state: &AppState, decision: &GatewayControlDecision) -> Value {
turn_redaction(state, decision, TEST_EMAIL)
.await
.client_event
}
/// 只有 `action` 没有 serde 默认值,其余字段都能省略。
fn decision_template(
provider_request_body: Value,
report_context: Value,
) -> AiExecutionDecision {
serde_json::from_value(json!({
"action": "local",
"candidate_id": "candidate-responses-ws",
"provider_request_body": provider_request_body,
"report_context": report_context,
}))
.expect("decision template should deserialize")
}
/// planner 在脱敏 body 上做模型映射后的 provider body。
fn provider_body_from(effective_event: &Value) -> Value {
let mut provider_body = effective_event.clone();
provider_body["model"] = json!("provider-model");
provider_body
}
/// 绑定那一轮留下的 report_context seed:故意带上原始 PII,用来证明这一轮
/// 会用脱敏后的 body 覆盖它,而不是把原文带进审计。
fn seed_report_context_with_raw_pii() -> Value {
json!({
"request_id": "connection",
"candidate_id": "candidate-responses-ws",
"original_request_body": {
"type": "response.create",
"model": "public-model",
"input": format!("mail {TEST_EMAIL}")
}
})
}
fn assert_redacted_json(value: &Value, label: &str) {
let serialized = serde_json::to_string(value).expect("value should serialize");
assert!(
!serialized.contains(TEST_EMAIL),
"{label} must not carry raw PII: {serialized}"
);
assert!(
serialized.contains("<AETHER:EMAIL:"),
"{label} must carry the redaction sentinel: {serialized}"
);
}
#[tokio::test]
async fn first_turn_upstream_and_audit_bodies_are_redacted() {
let state = redaction_enabled_state();
let decision = control_decision();
let effective_event = redacted_client_event(&state, &decision).await;
let template = decision_template(
provider_body_from(&effective_event),
seed_report_context_with_raw_pii(),
);
// 首轮实际发上游的事件由 decision.provider_request_body 派生。
let provider_event: Value = serde_json::from_str(
&planned_response_create_event(&template, &effective_event)
.expect("first provider event should serialize"),
)
.expect("first provider event should parse");
let turn_decision = prepare_responses_websocket_turn_decision(
&template,
"turn-1".to_string(),
true,
&effective_event,
&provider_event,
"connection",
1,
"logical-turn-1",
1,
);
assert_redacted_json(&provider_event, "first turn upstream event");
assert_redacted_json(
turn_decision
.provider_request_body
.as_ref()
.expect("turn decision should carry a provider body"),
"first turn provider request body",
);
let report_context = turn_decision
.report_context
.as_ref()
.expect("turn decision should carry a report context");
assert_redacted_json(
&report_context["original_request_body"],
"first turn audit body",
);
// 整个 report_context 都不该残留原文(seed 里的原始 body 必须被覆盖)。
assert_redacted_json(report_context, "first turn report context");
assert_eq!(provider_event["type"], "response.create");
assert_eq!(provider_event["model"], "provider-model");
}
#[tokio::test]
async fn continuation_upstream_and_audit_bodies_are_redacted() {
let state = redaction_enabled_state();
let decision = control_decision();
let effective_event = redacted_client_event(&state, &decision).await;
// continuation 复用已绑定的 upstream:不再规划,直接重放归一化器。
let outbound = normalize_followup_response_create(
&effective_event,
"provider-model",
&ResponsesWebSocketBodyNormalization::for_tests("provider-model"),
)
.expect("continuation should normalize");
let provider_event: Value =
serde_json::from_str(&outbound).expect("continuation event should parse");
let template = decision_template(
provider_body_from(&effective_event),
seed_report_context_with_raw_pii(),
);
let turn_decision = prepare_responses_websocket_turn_decision(
&template,
"turn-2".to_string(),
false,
&effective_event,
&provider_event,
"connection",
2,
"logical-turn-2",
1,
);
assert!(
!outbound.contains(TEST_EMAIL),
"continuation upstream frame must not carry raw PII: {outbound}"
);
assert!(
outbound.contains("<AETHER:EMAIL:"),
"continuation upstream frame must carry the sentinel: {outbound}"
);
assert_eq!(provider_event["previous_response_id"], "resp-previous");
let report_context = turn_decision
.report_context
.as_ref()
.expect("turn decision should carry a report context");
assert_redacted_json(
&report_context["original_request_body"],
"continuation audit body",
);
assert_redacted_json(report_context, "continuation report context");
}
#[tokio::test]
async fn quota_retry_replays_the_redacted_event() {
let state = redaction_enabled_state();
let decision = control_decision();
let effective_event = redacted_client_event(&state, &decision).await;
// 配额透明重试重放 LogicalTurn 里保存的事件,所以保存的必须
// 已经是脱敏版,否则重试会把原文发给新的上游账号。
let active = LogicalTurn::new(effective_event.clone(), 2, "logical-turn-2".to_string());
assert_redacted_json(&active.client_event, "quota retry replay event");
let template = decision_template(
provider_body_from(&active.client_event),
seed_report_context_with_raw_pii(),
);
let provider_event: Value = serde_json::from_str(
&planned_response_create_event(&template, &active.client_event)
.expect("retry provider event should serialize"),
)
.expect("retry provider event should parse");
let turn_decision = prepare_responses_websocket_turn_decision(
&template,
"turn-2-retry".to_string(),
true,
&active.client_event,
&provider_event,
"connection",
active.turn_index,
"logical-turn-2",
2,
);
assert_redacted_json(&provider_event, "quota retry upstream event");
let report_context = turn_decision
.report_context
.as_ref()
.expect("turn decision should carry a report context");
assert_eq!(report_context["websocket_turn_attempt"], 2);
assert_redacted_json(
&report_context["original_request_body"],
"quota retry audit body",
);
assert_redacted_json(report_context, "quota retry report context");
}
// -----------------------------------------------------------------------
// 响应侧还原
// -----------------------------------------------------------------------
/// 本次修复的核心:上游把占位符回显在事件里,客户端必须拿到真实值。
#[tokio::test]
async fn provider_frame_placeholders_are_restored_before_client_delivery() {
let state = redaction_enabled_state();
let redaction = turn_redaction(&state, &control_decision(), TEST_EMAIL).await;
let sentinel = sentinel_for(&redaction, TEST_EMAIL);
let mut restorer = ResponsesWebSocketRedactionRestorer::default();
restorer.register(redaction.session);
let frame = provider_delta_frame(&format!("your mail is {sentinel}"));
let restored = restorer
.restore_provider_frame_text(&frame)
.expect("a frame echoing this turn's sentinel must be restored");
assert!(
restored.contains(TEST_EMAIL),
"the client must receive the real value: {restored}"
);
assert!(
!restored.contains(&sentinel),
"no sentinel may survive to the client: {restored}"
);
// 协议字段不受影响,客户端的状态机照旧。
let restored: Value = serde_json::from_str(&restored).expect("restored frame is JSON");
assert_eq!(restored["type"], "response.output_text.delta");
assert_eq!(restored["item_id"], "msg_ws");
assert_eq!(restored["output_index"], 0);
}
/// Codex 把多个事件批量塞进 `{"chunks":[...]}`,还原必须走进批量里。
#[tokio::test]
async fn placeholders_batched_inside_a_chunks_envelope_are_restored() {
let state = redaction_enabled_state();
let redaction = turn_redaction(&state, &control_decision(), TEST_EMAIL).await;
let sentinel = sentinel_for(&redaction, TEST_EMAIL);
let mut restorer = ResponsesWebSocketRedactionRestorer::default();
restorer.register(redaction.session);
let frame = json!({
"chunks": [
provider_delta_frame("plain delta"),
provider_delta_frame(&format!("mail {sentinel}")),
]
});
let restored = restorer
.restore_provider_frame_text(&frame)
.expect("a batched sentinel must be restored");
assert!(restored.contains(TEST_EMAIL), "{restored}");
assert!(!restored.contains(&sentinel), "{restored}");
}
/// 只还原本连接 mask 过的映射,和 `RedactionSession::restore_text` 一致:
/// 别处来的占位符(比如客户端自己发的、或上一条连接的)保持原样。
#[tokio::test]
async fn an_unmapped_placeholder_is_left_untouched() {
let state = redaction_enabled_state();
let redaction = turn_redaction(&state, &control_decision(), TEST_EMAIL).await;
let sentinel = sentinel_for(&redaction, TEST_EMAIL);
let mut restorer = ResponsesWebSocketRedactionRestorer::default();
restorer.register(redaction.session);
let frame = provider_delta_frame(&format!("{FOREIGN_SENTINEL} and {sentinel}"));
let restored = restorer
.restore_provider_frame_text(&frame)
.expect("the mapped sentinel is still restored");
assert!(restored.contains(TEST_EMAIL), "{restored}");
assert!(
restored.contains(FOREIGN_SENTINEL),
"an unmapped placeholder must survive verbatim: {restored}"
);
}
/// 没有命中还原时必须让调用方原样转发上游字节。
#[tokio::test]
async fn a_frame_without_known_placeholders_is_not_rewritten() {
let state = redaction_enabled_state();
let redaction = turn_redaction(&state, &control_decision(), TEST_EMAIL).await;
let mut restorer = ResponsesWebSocketRedactionRestorer::default();
restorer.register(redaction.session);
assert!(
restorer
.restore_provider_frame_text(&provider_delta_frame("nothing to restore"))
.is_none(),
"a frame with no mapped sentinel must be relayed byte-for-byte"
);
assert!(
restorer
.restore_provider_frame_text(&provider_delta_frame(FOREIGN_SENTINEL))
.is_none(),
"a frame that only carries unmapped placeholders must not be rewritten"
);
}
/// 未启用脱敏(或这条连接从没 mask 到东西)时,还原器必须完全不介入:
/// 连 clone 都不做,输出就是上游原字节。
#[tokio::test]
async fn a_restorer_without_sessions_never_rewrites_a_frame() {
let restorer = ResponsesWebSocketRedactionRestorer::default();
assert!(restorer
.restore_provider_frame_text(&provider_delta_frame(FOREIGN_SENTINEL))
.is_none());
assert!(restorer
.restore_provider_frame_text(&provider_delta_frame(TEST_EMAIL))
.is_none());
}
/// 空 session(启用了脱敏但这一轮没命中任何规则)不该被留下来白扫每一帧。
#[tokio::test]
async fn a_session_without_mappings_is_not_retained() {
let state = redaction_enabled_state();
let hmac_key = state
.encryption_key()
.expect("the test state carries an encryption key")
.as_bytes()
.to_vec();
let empty_session = crate::privacy::RedactionSession::new(
crate::privacy::RedactionSessionConfig::default_ttl(hmac_key, 0),
);
let mut restorer = ResponsesWebSocketRedactionRestorer::default();
restorer.register(empty_session);
assert!(restorer
.restore_provider_frame_text(&provider_delta_frame(FOREIGN_SENTINEL))
.is_none());
}
/// 还原只作用于发给客户端的那一份拷贝:审计和终态观测消费的事件必须保持脱敏态。
#[tokio::test]
async fn restoring_does_not_mutate_the_event_the_audit_path_keeps() {
let state = redaction_enabled_state();
let redaction = turn_redaction(&state, &control_decision(), TEST_EMAIL).await;
let sentinel = sentinel_for(&redaction, TEST_EMAIL);
let mut restorer = ResponsesWebSocketRedactionRestorer::default();
restorer.register(redaction.session);
let frame = provider_delta_frame(&format!("mail {sentinel}"));
let before = frame.clone();
let _ = restorer
.restore_provider_frame_text(&frame)
.expect("the frame is restored for the client");
assert_eq!(
frame, before,
"capture_client_frame / 终态观测拿到的事件必须仍是脱敏态"
);
}
/// 连接级持有的意义:WS 的会话历史留在上游,continuation 只发增量输入,
/// 所以第 2 轮的响应可能回显第 1 轮的占位符。per-turn 持有会漏掉这一条。
#[tokio::test]
async fn a_later_turn_restores_a_placeholder_first_masked_by_an_earlier_turn() {
let state = redaction_enabled_state();
let decision = control_decision();
let first = turn_redaction(&state, &decision, TEST_EMAIL).await;
let second = turn_redaction(&state, &decision, OTHER_TEST_EMAIL).await;
let first_sentinel = sentinel_for(&first, TEST_EMAIL);
let second_sentinel = sentinel_for(&second, OTHER_TEST_EMAIL);
assert_ne!(first_sentinel, second_sentinel);
let mut restorer = ResponsesWebSocketRedactionRestorer::default();
restorer.register(first.session);
restorer.register(second.session);
let frame = provider_delta_frame(&format!("{first_sentinel} then {second_sentinel}"));
let restored = restorer
.restore_provider_frame_text(&frame)
.expect("both turns' sentinels are restorable on this connection");
assert!(restored.contains(TEST_EMAIL), "{restored}");
assert!(restored.contains(OTHER_TEST_EMAIL), "{restored}");
assert!(!restored.contains(&first_sentinel), "{restored}");
assert!(!restored.contains(&second_sentinel), "{restored}");
}
/// 留存窗口是有界的:长连接不能无限累积映射,代价是更早的轮次会退回
/// 「占位符原样透传」而不是被错误还原成别的值。
#[tokio::test]
async fn the_retained_session_window_is_bounded() {
let state = redaction_enabled_state();
let decision = control_decision();
let oldest = turn_redaction(&state, &decision, TEST_EMAIL).await;
let oldest_sentinel = sentinel_for(&oldest, TEST_EMAIL);
let mut restorer = ResponsesWebSocketRedactionRestorer::default();
restorer.register(oldest.session);
// 再灌满整个窗口,最老的那一轮必须被挤出去。
let mut newest_sentinel = String::new();
for index in 0..MAX_RETAINED_TURN_REDACTION_SESSIONS {
let email = format!("ws.turn{index}@example.com");
let redaction = turn_redaction(&state, &decision, &email).await;
newest_sentinel = sentinel_for(&redaction, &email);
restorer.register(redaction.session);
}
assert!(
restorer
.restore_provider_frame_text(&provider_delta_frame(&oldest_sentinel))
.is_none(),
"the evicted turn's sentinel is relayed verbatim, never mis-restored"
);
assert!(
restorer
.restore_provider_frame_text(&provider_delta_frame(&newest_sentinel))
.is_some(),
"the most recent turns stay restorable"
);
}
}

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