Files
Aether/apps/aether-gateway/src/handlers/proxy/websocket/responses/quota.rs
T
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

394 lines
14 KiB
Rust

//! Quota exhaustion, replay safety, and upstream replacement policy.
use axum::extract::ws::WebSocket;
use futures_util::SinkExt;
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, ActiveResponsesWebSocketTurn};
use super::request::{
build_planning_parts, planned_response_create_event, response_create_has_previous_response_id,
};
use super::state::BoundResponsesConnection;
use super::turn::{
begin_responses_websocket_turn, prepare_responses_websocket_turn_decision,
ResponsesWebSocketTurnOutcome,
};
use super::upstream::{bind_responses_upstream, close_bound_upstream};
use crate::ai_serving::maybe_build_responses_websocket_decision;
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, send_responses_websocket_error,
};
use crate::orchestration::release_pool_key_lease_from_report_context;
use crate::AppState;
const PREVIOUS_RESPONSE_NOT_FOUND_MESSAGE: &str =
"Previous response was not found. Retrying the full request.";
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))
}
pub(super) async fn retry_active_turn_after_quota_exhaustion(
bound: &mut BoundResponsesConnection,
state: &AppState,
context: &WebSocketRequestContext,
) -> 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 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 maybe_build_responses_websocket_decision(
state,
&planning_parts,
&turn_request_id,
&context.decision,
&client_event,
excluded_key_ids,
excluded_codex_account_ids,
)
.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 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() {
release_pool_key_lease_from_report_context(state, decision.report_context.as_ref()).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) => {
release_pool_key_lease_from_report_context(state, decision.report_context.as_ref())
.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(
state,
&planning_parts,
&context.decision,
turn_decision,
&client_event,
)
.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(ActiveResponsesWebSocketTurn::new(state, 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 active_continuation_can_retry_from_full_input(
bound: &BoundResponsesConnection,
) -> bool {
bound.turn_state.logical().is_some_and(|active| {
response_create_has_previous_response_id(&active.client_event)
&& active.retry_unsafe_reason.is_none()
})
}
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 should_request_full_continuation_retry(
bound: &BoundResponsesConnection,
retry_current_turn: bool,
upstream_event: Option<&Value>,
) -> bool {
retry_current_turn
&& active_continuation_can_retry_from_full_input(bound)
&& upstream_event.is_some_and(is_usage_limit_error_event)
}
pub(super) async fn send_previous_response_not_found(client_socket: &mut WebSocket) {
send_responses_websocket_error(
client_socket,
400,
"invalid_request_error",
"previous_response_not_found",
PREVIOUS_RESPONSE_NOT_FOUND_MESSAGE,
)
.await;
}
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);
}
}