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。
This commit is contained in:
AAEE86
2026-08-17 14:52:57 +08:00
committed by ZheFox
parent 9d80281b53
commit 1c5ee5228c
10 changed files with 428 additions and 184 deletions
@@ -20,11 +20,12 @@ use super::request::{
planned_response_create_event, provider_model_from_decision,
response_create_has_previous_response_id, response_create_model_or_current,
};
use super::state::{ActiveResponsesWebSocketRequest, BoundResponsesConnection};
use super::state::BoundResponsesConnection;
use super::turn::{
begin_responses_websocket_turn, prepare_responses_websocket_turn_decision,
ResponsesWebSocketTurnObservation, ResponsesWebSocketTurnOutcome,
};
use super::turn_state::LogicalTurn;
use super::upstream::{bind_responses_upstream, decision_reuses_bound_upstream};
use crate::ai_serving::maybe_build_responses_websocket_decision;
use crate::clock::current_unix_secs;
@@ -118,7 +119,7 @@ pub(super) async fn forward_client_message(
));
}
if bound.response_in_flight {
if !bound.turn_state.accepts_new_response_create() {
send_gateway_error(
client_socket,
"response_already_in_progress",
@@ -366,21 +367,18 @@ pub(super) async fn forward_client_message(
}
};
turn.set_provider_response_headers(bound.upstream_response_headers.clone());
bound.active_turn = Some(ActiveResponsesWebSocketTurn::new(state, turn));
bound.active_response_create = Some(ActiveResponsesWebSocketRequest::new(
client_event.clone(),
turn_index,
logical_turn_id,
));
bound.turn_state.begin(
LogicalTurn::new(client_event.clone(), turn_index, logical_turn_id),
ActiveResponsesWebSocketTurn::new(state, turn),
);
bound.next_turn_index = bound.next_turn_index.saturating_add(1);
bound.response_in_flight = true;
let Some(upstream) = bound.upstream.as_mut() else {
return RelayDisposition::UpstreamError("responses_websocket_send_failed");
};
match send_upstream_message(upstream, WreqWsMessage::text(outbound)).await {
Ok(()) => {
if let Some(turn) = bound.active_turn.as_mut() {
if let Some(turn) = bound.turn_state.attempt_mut() {
turn.mark_upstream_request_sent();
}
RelayDisposition::Continue
@@ -697,14 +695,11 @@ async fn forward_replanned_response_create(
// The re-plan keeps this upstream but resolved a new model, so later
// continuations must normalize against the new plan, not the old one.
bound.body_normalization = normalization;
bound.active_turn = Some(ActiveResponsesWebSocketTurn::new(state, turn));
bound.active_response_create = Some(ActiveResponsesWebSocketRequest::new(
client_event.clone(),
turn_index,
logical_turn_id.clone(),
));
bound.turn_state.begin(
LogicalTurn::new(client_event.clone(), turn_index, logical_turn_id.clone()),
ActiveResponsesWebSocketTurn::new(state, turn),
);
bound.next_turn_index = bound.next_turn_index.saturating_add(1);
bound.response_in_flight = true;
debug!(
event_name = "responses_websocket_followup_model_replanned",
log_type = "event",
@@ -769,16 +764,13 @@ async fn forward_replanned_response_create(
bound.adapter = replacement.adapter;
bound.client_model = replacement.client_model;
bound.provider_model = replacement.provider_model;
bound.response_in_flight = true;
bound.decision_template = replacement.decision_template;
bound.body_normalization = replacement.body_normalization;
bound.binding_identity = replacement.binding_identity;
bound.active_turn = Some(ActiveResponsesWebSocketTurn::new(state, turn));
bound.active_response_create = Some(ActiveResponsesWebSocketRequest::new(
client_event,
turn_index,
logical_turn_id,
));
bound.turn_state.begin(
LogicalTurn::new(client_event, turn_index, logical_turn_id),
ActiveResponsesWebSocketTurn::new(state, turn),
);
bound.next_turn_index = bound.next_turn_index.saturating_add(1);
bound.upstream_response_headers = replacement.upstream_response_headers;
bound.pending_adapter_drain = replacement.pending_adapter_drain;
@@ -21,8 +21,7 @@ use super::quota::{
send_previous_response_not_found, should_request_full_continuation_retry,
};
use super::relay_policy::{
classify_quota_relay, classify_upstream_frame, fatal_relay_policy, FatalRelaySignal,
QuotaRelayAction, QuotaRelayFacts, UpstreamFrameAction, UpstreamFrameKind,
classify_quota_relay, fatal_relay_policy, FatalRelaySignal, QuotaRelayAction, QuotaRelayFacts,
};
use super::state::BoundResponsesConnection;
use super::turn::{
@@ -65,7 +64,7 @@ pub(super) async fn relay_bound_connection(
tokio::pin!(connection_deadline);
loop {
let active_turn_deadline = bound.active_turn.as_ref().map(|turn| turn.deadline());
let active_turn_deadline = bound.turn_state.attempt().map(|turn| turn.deadline());
tokio::select! {
_ = &mut connection_deadline => {
finalize_active_turn(
@@ -79,7 +78,6 @@ pub(super) async fn relay_bound_connection(
"websocket_connection_limit_reached",
"WebSocket connection duration limit reached; reconnect to continue",
).await;
bound.active_response_create = None;
close_bound_upstream(bound).await;
close_client_socket(client_socket, CLOSE_TRY_AGAIN, "connection_limit_reached").await;
break;
@@ -105,7 +103,6 @@ pub(super) async fn relay_bound_connection(
turn_deadline.phase.error_code(),
turn_deadline.phase.client_message(),
).await;
bound.active_response_create = None;
close_bound_upstream(bound).await;
close_client_socket(
client_socket,
@@ -129,7 +126,6 @@ pub(super) async fn relay_bound_connection(
state,
ResponsesWebSocketTurnOutcome::connection_admission_lost(),
).await;
bound.active_response_create = None;
close_bound_upstream(bound).await;
send_gateway_error_with_status(
client_socket,
@@ -147,7 +143,6 @@ pub(super) async fn relay_bound_connection(
state,
ResponsesWebSocketTurnOutcome::client_disconnected(),
).await;
bound.active_response_create = None;
close_bound_upstream(bound).await;
break;
};
@@ -165,7 +160,6 @@ pub(super) async fn relay_bound_connection(
state,
ResponsesWebSocketTurnOutcome::client_disconnected(),
).await;
bound.active_response_create = None;
close_bound_upstream(bound).await;
break;
};
@@ -200,7 +194,6 @@ pub(super) async fn relay_bound_connection(
code,
"Gateway could not forward the WebSocket event upstream",
).await;
bound.active_response_create = None;
close_bound_upstream(bound).await;
close_client_socket(client_socket, CLOSE_INTERNAL_ERROR, code).await;
break;
@@ -214,7 +207,6 @@ pub(super) async fn relay_bound_connection(
state,
ResponsesWebSocketTurnOutcome::upstream_closed(),
).await;
bound.active_response_create = None;
bound.upstream = None;
close_client_socket(client_socket, 1000, "upstream_closed").await;
break;
@@ -239,7 +231,6 @@ pub(super) async fn relay_bound_connection(
"responses_websocket_receive_failed",
"Provider connection closed unexpectedly",
).await;
bound.active_response_create = None;
bound.upstream = None;
close_client_socket(client_socket, CLOSE_INTERNAL_ERROR, "upstream_receive_failed").await;
break;
@@ -268,7 +259,7 @@ pub(super) async fn relay_bound_connection(
chunked = parsed_upstream_frame
.as_ref()
.is_some_and(ParsedResponsesWebSocketFrame::is_chunked),
active_turn = bound.active_turn.is_some(),
active_turn = bound.turn_state.response_in_flight(),
"gateway received Responses WebSocket event"
);
}
@@ -315,11 +306,11 @@ pub(super) async fn relay_bound_connection(
let adapter = bound.adapter;
match parsed_upstream_frame.as_ref() {
Some(frame) => bound
.active_turn
.as_mut()
.turn_state
.attempt_mut()
.and_then(|turn| turn.observe_upstream_frame(frame, adapter)),
None => {
if let Some(turn) = bound.active_turn.as_mut() {
if let Some(turn) = bound.turn_state.attempt_mut() {
turn.observe_invalid_upstream_text(text.as_str())
}
else {
@@ -330,13 +321,12 @@ pub(super) async fn relay_bound_connection(
}
_ => None,
};
update_response_in_flight(bound, parsed_upstream_frame.as_ref());
if matches!(
observation,
Some(ResponsesWebSocketTurnObservation::Started)
| Some(ResponsesWebSocketTurnObservation::Terminal(_))
) {
if let Some(turn) = bound.active_turn.as_mut() {
if let Some(turn) = bound.turn_state.attempt_mut() {
turn.mark_stream_started(state).await;
}
}
@@ -344,9 +334,6 @@ pub(super) async fn relay_bound_connection(
Some(ResponsesWebSocketTurnObservation::Terminal(outcome)) => Some(outcome),
_ => None,
};
if terminal_outcome.is_some() {
bound.response_in_flight = false;
}
if matches!(&upstream_message, WreqWsMessage::Text(_))
&& parsed_upstream_frame.is_none()
{
@@ -359,7 +346,6 @@ pub(super) async fn relay_bound_connection(
),
)
.await;
bound.active_response_create = None;
send_responses_websocket_error(
client_socket,
policy.status_code,
@@ -380,7 +366,7 @@ pub(super) async fn relay_bound_connection(
let is_close = matches!(upstream_message, WreqWsMessage::Close(_));
let drain_for_adapter = adapter_drain_ready(
bound.pending_adapter_drain,
bound.response_in_flight,
bound.turn_state.response_in_flight(),
observation,
is_close,
);
@@ -396,7 +382,8 @@ pub(super) async fn relay_bound_connection(
};
let mut quota_relay_action = classify_quota_relay(quota_facts);
if matches!(quota_relay_action, QuotaRelayAction::AttemptTransparentRetry) {
let mut retry_turn = bound.active_turn.take().map(ActiveResponsesWebSocketTurn::disarm);
// detach_attempt 保留 logical turn:重试是同一轮请求的下一个 attempt。
let mut retry_turn = bound.turn_state.detach_attempt().map(ActiveResponsesWebSocketTurn::disarm);
if let Some(turn) = retry_turn.as_mut() {
turn.release_admission().await;
}
@@ -414,7 +401,16 @@ pub(super) async fn relay_bound_connection(
}
continue;
}
bound.active_turn = retry_turn.map(|turn| ActiveResponsesWebSocketTurn::new(state, turn));
if let Some(turn) = retry_turn {
let restored = bound
.turn_state
.resume(ActiveResponsesWebSocketTurn::new(state, turn));
debug_assert!(
restored.is_ok(),
"a failed transparent retry must be able to restore its attempt"
);
drop(restored);
}
quota_relay_action = classify_quota_relay(QuotaRelayFacts {
retry_current_turn: false,
transparent_retry_failed: true,
@@ -437,7 +433,7 @@ pub(super) async fn relay_bound_connection(
error_code = "previous_response_not_found",
"gateway will ask the client to retry the continuation with complete input"
);
let mut turn = bound.active_turn.take().map(ActiveResponsesWebSocketTurn::disarm);
let mut turn = bound.turn_state.end().map(ActiveResponsesWebSocketTurn::disarm);
if let Some(active_turn) = turn.as_mut() {
active_turn.release_admission().await;
}
@@ -453,7 +449,6 @@ pub(super) async fn relay_bound_connection(
)
.await;
}
bound.active_response_create = None;
detach_exhausted_upstream(bound, directive, &context.trace_id).await;
continue;
}
@@ -475,7 +470,6 @@ pub(super) async fn relay_bound_connection(
"Provider connection closed after reporting exhausted quota; send a new response.create to select another Provider connection",
)
.await;
bound.active_response_create = None;
detach_exhausted_upstream(bound, directive, &context.trace_id).await;
continue;
}
@@ -497,18 +491,16 @@ pub(super) async fn relay_bound_connection(
state,
ResponsesWebSocketTurnOutcome::client_disconnected(),
).await;
bound.active_response_create = None;
close_bound_upstream(bound).await;
break;
}
if let (Some(turn), Some(frame)) =
(bound.active_turn.as_mut(), parsed_upstream_frame.as_ref())
(bound.turn_state.attempt_mut(), parsed_upstream_frame.as_ref())
{
turn.capture_client_frame(frame.event());
}
if let Some(outcome) = terminal_outcome {
finalize_active_turn(bound, state, outcome).await;
bound.active_response_create = None;
} else if is_close {
finalize_active_turn(
bound,
@@ -521,7 +513,6 @@ pub(super) async fn relay_bound_connection(
let directive = bound
.pending_adapter_drain
.expect("adapter drain state should be present");
bound.active_response_create = None;
detach_exhausted_upstream(bound, directive, &context.trace_id).await;
continue;
}
@@ -549,27 +540,3 @@ async fn wait_for_connection_permit_loss(permit: Option<&aether_runtime::Admissi
}
}
fn update_response_in_flight(
bound: &mut BoundResponsesConnection,
frame: Option<&ParsedResponsesWebSocketFrame<'_>>,
) {
let Some(frame) = frame else {
return;
};
let frame_kind = if frame.is_terminal() {
UpstreamFrameKind::Terminal
} else if frame.is_started() {
UpstreamFrameKind::Started
} else {
UpstreamFrameKind::Other
};
match classify_upstream_frame(frame_kind) {
UpstreamFrameAction::Continue if frame.is_started() => {
bound.response_in_flight = true;
}
UpstreamFrameAction::FinalizeTurn => {
bound.response_in_flight = false;
}
UpstreamFrameAction::Continue | UpstreamFrameAction::FinalizeAndClose => {}
}
}
@@ -103,12 +103,16 @@ impl Drop for ActiveResponsesWebSocketTurn {
}
}
/// 结束当前 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.active_turn.take() {
if let Some(turn) = bound.turn_state.end() {
queue_turn_finalization(bound, state, turn.disarm(), outcome).await;
}
}
@@ -21,6 +21,7 @@ mod request;
mod session;
mod state;
mod turn;
mod turn_state;
mod upstream;
use std::net::SocketAddr;
@@ -53,7 +53,12 @@ pub(super) async fn detach_exhausted_upstream(
) {
let exclusion = record_exhausted_bound_key(bound, directive.retry_exclusion_until_unix_secs);
close_bound_upstream(bound).await;
bound.response_in_flight = false;
// 调用方必须先结束当前 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!(
@@ -99,7 +104,7 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion(
state: &AppState,
context: &WebSocketRequestContext,
) -> bool {
let Some(active) = bound.active_response_create.as_mut() else {
let Some(active) = bound.turn_state.logical_mut() else {
return false;
};
if let Some(reason) = active.quota_retry_block_reason() {
@@ -291,11 +296,19 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion(
bound.adapter = replacement.adapter;
bound.client_model = replacement.client_model;
bound.provider_model = replacement.provider_model;
bound.response_in_flight = true;
bound.decision_template = replacement.decision_template;
bound.body_normalization = replacement.body_normalization;
bound.binding_identity = replacement.binding_identity;
bound.active_turn = Some(ActiveResponsesWebSocketTurn::new(state, turn));
// 同一个 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!(
@@ -317,7 +330,7 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion(
pub(super) fn active_continuation_can_retry_from_full_input(
bound: &BoundResponsesConnection,
) -> bool {
bound.active_response_create.as_ref().is_some_and(|active| {
bound.turn_state.logical().is_some_and(|active| {
response_create_has_previous_response_id(&active.client_event)
&& active.retry_unsafe_reason.is_none()
})
@@ -365,7 +378,7 @@ pub(super) fn observe_active_response_rebind_safety(
else {
return;
};
if let Some(active) = bound.active_response_create.as_mut() {
if let Some(active) = bound.turn_state.logical_mut() {
active.mark_retry_unsafe(reason);
}
}
@@ -374,7 +387,7 @@ pub(super) fn mark_active_response_retry_unsafe(
bound: &mut BoundResponsesConnection,
reason: &'static str,
) {
if let Some(active) = bound.active_response_create.as_mut() {
if let Some(active) = bound.turn_state.logical_mut() {
active.mark_retry_unsafe(reason);
}
}
@@ -78,7 +78,7 @@ mod tests {
use super::super::request::{
build_planning_parts, normalize_followup_response_create, planned_response_create_event,
};
use super::super::state::ActiveResponsesWebSocketRequest;
use super::super::turn_state::LogicalTurn;
use super::super::turn::prepare_responses_websocket_turn_decision;
use super::redact_responses_websocket_client_event;
use crate::ai_serving::{AiExecutionDecision, ResponsesWebSocketBodyNormalization};
@@ -435,9 +435,9 @@ mod tests {
let state = redaction_enabled_state();
let decision = control_decision();
let effective_event = redacted_client_event(&state, &decision).await;
// 配额透明重试重放 active_response_create 里保存的事件,所以保存的必须
// 配额透明重试重放 LogicalTurn 里保存的事件,所以保存的必须
// 已经是脱敏版,否则重试会把原文发给新的上游账号。
let active = ActiveResponsesWebSocketRequest::new(
let active = LogicalTurn::new(
effective_event.clone(),
2,
"logical-turn-2".to_string(),
@@ -24,11 +24,11 @@ use super::lifecycle::{
};
use super::redaction::redact_responses_websocket_client_event;
use super::request::{build_planning_parts, planned_response_create_event};
use super::state::ActiveResponsesWebSocketRequest;
use super::turn::{
begin_responses_websocket_turn, prepare_responses_websocket_turn_decision,
ResponsesWebSocketTurnOutcome,
};
use super::turn_state::LogicalTurn;
use super::upstream::bind_responses_upstream;
use crate::ai_serving::maybe_build_responses_websocket_decision;
@@ -413,12 +413,10 @@ pub(super) async fn run_responses_websocket(
};
first_turn.mark_upstream_request_sent();
first_turn.set_provider_response_headers(bound.upstream_response_headers.clone());
bound.active_turn = Some(ActiveResponsesWebSocketTurn::new(&state, first_turn));
bound.active_response_create = Some(ActiveResponsesWebSocketRequest::new(
first_event,
1,
first_logical_turn_id,
));
bound.turn_state.begin(
LogicalTurn::new(first_event, 1, first_logical_turn_id),
ActiveResponsesWebSocketTurn::new(&state, first_turn),
);
relay_bound_connection(
&mut client_socket,
@@ -532,9 +530,9 @@ mod tests {
response_create_model_or_current,
};
use super::super::state::{
ActiveResponsesWebSocketRequest, BoundResponsesConnection,
ExhaustedResponsesWebSocketExclusions,
BoundResponsesConnection, ExhaustedResponsesWebSocketExclusions,
};
use super::super::turn_state::{LogicalTurn, ResponsesTurnState};
use super::super::turn::{
ResponsesWebSocketTurnDeadline, ResponsesWebSocketTurnObservation,
ResponsesWebSocketTurnOutcome, ResponsesWebSocketTurnTimeoutPhase,
@@ -734,19 +732,21 @@ mod tests {
#[test]
fn quota_error_can_request_a_full_retry_only_before_public_response_state() {
let mut bound = sample_bound_for_rebind_safety();
bound.active_response_create = Some(ActiveResponsesWebSocketRequest::new(
json!({
"type": "response.create",
"previous_response_id": "resp-previous",
}),
2,
"logical-turn".to_string(),
));
bound.turn_state = ResponsesTurnState::Replanning {
logical: LogicalTurn::new(
json!({
"type": "response.create",
"previous_response_id": "resp-previous",
}),
2,
"logical-turn".to_string(),
),
};
assert!(active_continuation_can_retry_from_full_input(&bound));
bound
.active_response_create
.as_mut()
.turn_state
.logical_mut()
.expect("active request")
.mark_retry_unsafe("standard_response_event");
assert!(!active_continuation_can_retry_from_full_input(&bound));
@@ -772,14 +772,16 @@ mod tests {
#[test]
fn full_continuation_retry_does_not_consume_a_successful_terminal_event() {
let mut bound = sample_bound_for_rebind_safety();
bound.active_response_create = Some(ActiveResponsesWebSocketRequest::new(
json!({
"type": "response.create",
"previous_response_id": "resp-previous",
}),
2,
"logical-turn".to_string(),
));
bound.turn_state = ResponsesTurnState::Replanning {
logical: LogicalTurn::new(
json!({
"type": "response.create",
"previous_response_id": "resp-previous",
}),
2,
"logical-turn".to_string(),
),
};
assert!(should_request_full_continuation_retry(
&bound,
@@ -879,7 +881,7 @@ mod tests {
#[test]
fn quota_retry_requires_an_explicitly_replay_safe_turn() {
let mut request = ActiveResponsesWebSocketRequest::new(
let mut request = LogicalTurn::new(
json!({"type": "response.create", "model": "gpt-5.6-sol"}),
2,
"logical-turn".to_string(),
@@ -892,7 +894,7 @@ mod tests {
Some("standard_response_event")
);
let mut retried = ActiveResponsesWebSocketRequest::new(
let mut retried = LogicalTurn::new(
json!({"type": "response.create", "model": "gpt-5.6-sol"}),
2,
"logical-turn".to_string(),
@@ -903,7 +905,7 @@ mod tests {
Some("quota_retry_already_attempted")
);
let mut client_control = ActiveResponsesWebSocketRequest::new(
let mut client_control = LogicalTurn::new(
json!({"type": "response.create", "model": "gpt-5.6-sol"}),
2,
"logical-turn".to_string(),
@@ -914,7 +916,7 @@ mod tests {
Some("client_control_event")
);
let continuation = ActiveResponsesWebSocketRequest::new(
let continuation = LogicalTurn::new(
json!({
"type": "response.create",
"model": "gpt-5.6-sol",
@@ -941,18 +943,18 @@ mod tests {
);
assert_eq!(
bound
.active_response_create
.as_ref()
.and_then(ActiveResponsesWebSocketRequest::quota_retry_block_reason),
.turn_state
.logical()
.and_then(LogicalTurn::quota_retry_block_reason),
None
);
observe_active_response_rebind_safety(&mut bound, &json!({"type": "response.created"}));
assert_eq!(
bound
.active_response_create
.as_ref()
.and_then(ActiveResponsesWebSocketRequest::quota_retry_block_reason),
.turn_state
.logical()
.and_then(LogicalTurn::quota_retry_block_reason),
Some("standard_response_event")
);
@@ -960,9 +962,9 @@ mod tests {
observe_active_response_rebind_safety(&mut unknown, &json!({"type": "codex.unknown"}));
assert_eq!(
unknown
.active_response_create
.as_ref()
.and_then(ActiveResponsesWebSocketRequest::quota_retry_block_reason),
.turn_state
.logical()
.and_then(LogicalTurn::quota_retry_block_reason),
Some("unrecognized_upstream_event")
);
}
@@ -1192,16 +1194,18 @@ mod tests {
adapter,
client_model: "gpt-5.6-sol".to_string(),
provider_model: "gpt-5.6-sol".to_string(),
response_in_flight: true,
decision_template: decision,
body_normalization: ResponsesWebSocketBodyNormalization::for_tests("gpt-5.6-sol"),
binding_identity,
active_turn: None,
active_response_create: Some(ActiveResponsesWebSocketRequest::new(
json!({"type": "response.create", "model": "gpt-5.6-sol"}),
1,
"logical-turn".to_string(),
)),
// Replanning:logical turn 在、attempt 不在。重放安全与配额排除都只看
// logical turn,所以这些用例不需要真实 socket 或真实 attempt。
turn_state: ResponsesTurnState::Replanning {
logical: LogicalTurn::new(
json!({"type": "response.create", "model": "gpt-5.6-sol"}),
1,
"logical-turn".to_string(),
),
},
next_turn_index: 2,
upstream_response_headers: BTreeMap::new(),
pending_adapter_drain: None,
@@ -4,14 +4,12 @@
//! connection may survive many `response.create` turns, while the turn
//! lifecycle and upstream binding are replaced independently.
use serde_json::Value;
use std::collections::{BTreeMap, BTreeSet};
use tokio::task::JoinHandle;
use super::adapter::{ResponsesWebSocketDrainDirective, ResponsesWebSocketProtocolAdapter};
use super::binding::UpstreamBindingIdentity;
use super::lifecycle::ActiveResponsesWebSocketTurn;
use super::request::response_create_has_previous_response_id;
use super::turn_state::ResponsesTurnState;
use crate::ai_serving::{AiExecutionDecision, ResponsesWebSocketBodyNormalization};
const EXHAUSTED_KEY_EXCLUSION_FALLBACK_SECONDS: u64 = 300;
@@ -22,15 +20,14 @@ pub(super) struct BoundResponsesConnection {
pub(super) adapter: &'static dyn ResponsesWebSocketProtocolAdapter,
pub(super) client_model: String,
pub(super) provider_model: String,
pub(super) response_in_flight: bool,
pub(super) decision_template: AiExecutionDecision,
/// Reproduces this binding's provider-body normalization for continuation
/// turns, which must not re-enter the planner. Replaced whenever the
/// binding or its decision is replaced.
pub(super) body_normalization: ResponsesWebSocketBodyNormalization,
pub(super) binding_identity: UpstreamBindingIdentity,
pub(super) active_turn: Option<ActiveResponsesWebSocketTurn>,
pub(super) active_response_create: Option<ActiveResponsesWebSocketRequest>,
/// 这条连接上「有没有正在进行的 logical turn」的唯一事实来源。
pub(super) turn_state: ResponsesTurnState,
pub(super) next_turn_index: u64,
pub(super) upstream_response_headers: BTreeMap<String, String>,
pub(super) pending_adapter_drain: Option<ResponsesWebSocketDrainDirective>,
@@ -101,41 +98,3 @@ impl ExhaustedResponsesWebSocketExclusions {
}
}
#[derive(Debug, Clone)]
pub(super) struct ActiveResponsesWebSocketRequest {
pub(super) client_event: Value,
pub(super) turn_index: u64,
pub(super) logical_turn_id: String,
pub(super) turn_attempt: u32,
pub(super) retry_attempted: bool,
pub(super) retry_unsafe_reason: Option<&'static str>,
}
impl ActiveResponsesWebSocketRequest {
pub(super) fn new(client_event: Value, turn_index: u64, logical_turn_id: String) -> Self {
Self {
client_event,
turn_index,
logical_turn_id,
turn_attempt: 1,
retry_attempted: false,
retry_unsafe_reason: None,
}
}
pub(super) fn quota_retry_block_reason(&self) -> Option<&'static str> {
if self.retry_attempted {
Some("quota_retry_already_attempted")
} else if let Some(reason) = self.retry_unsafe_reason {
Some(reason)
} else if response_create_has_previous_response_id(&self.client_event) {
Some("previous_response_id")
} else {
None
}
}
pub(super) fn mark_retry_unsafe(&mut self, reason: &'static str) {
self.retry_unsafe_reason.get_or_insert(reason);
}
}
@@ -0,0 +1,303 @@
//! 一条 Responses WebSocket 连接上的 turn 状态机。
//!
//! 现状把「有没有正在进行的 turn」拆成 `response_in_flight`、`active_turn`、
//! `active_response_create` 三个可独立变化的字段,8 种组合里只有 3 种合法,
//! 非法组合只能靠调用点的 if 和「记得同时改另外两个字段」来避免。这里把它收敛成
//! 一个枚举:合法组合由类型保证,转换只能走受控 API。
use serde_json::Value;
use super::lifecycle::ActiveResponsesWebSocketTurn;
use super::request::response_create_has_previous_response_id;
/// 客户端一次 `response.create` 对应的 logical turn。
///
/// 一个 logical turn 可能经历多个 provider attempt:配额透明重试会换一把 key、
/// 换一条上游连接重放同一份客户端事件,但对客户端始终是同一轮请求。
/// `client_event` 保存的必须是**已脱敏**的事件(见 `super::redaction`),
/// 因为透明重试直接重放它。
#[derive(Debug, Clone)]
pub(super) struct LogicalTurn {
pub(super) client_event: Value,
pub(super) turn_index: u64,
pub(super) logical_turn_id: String,
pub(super) turn_attempt: u32,
pub(super) retry_attempted: bool,
pub(super) retry_unsafe_reason: Option<&'static str>,
}
impl LogicalTurn {
pub(super) fn new(client_event: Value, turn_index: u64, logical_turn_id: String) -> Self {
Self {
client_event,
turn_index,
logical_turn_id,
turn_attempt: 1,
retry_attempted: false,
retry_unsafe_reason: None,
}
}
pub(super) fn quota_retry_block_reason(&self) -> Option<&'static str> {
if self.retry_attempted {
Some("quota_retry_already_attempted")
} else if let Some(reason) = self.retry_unsafe_reason {
Some(reason)
} else if response_create_has_previous_response_id(&self.client_event) {
Some("previous_response_id")
} else {
None
}
}
pub(super) fn mark_retry_unsafe(&mut self, reason: &'static str) {
self.retry_unsafe_reason.get_or_insert(reason);
}
}
/// 连接上「有没有正在进行的 logical turn」这一唯一事实。
///
/// 类型参数只为测试留出注入点:生产代码一律用默认的
/// [`ActiveResponsesWebSocketTurn`],测试用轻量替身驱动同一套转换逻辑,
/// 不必构造 `AppState` 和真实 socket。
pub(super) enum ResponsesTurnState<A = ActiveResponsesWebSocketTurn> {
/// 没有进行中的 logical turn。上游可能仍绑定,也可能已被 detach。
Idle,
/// 一个 logical turn 正在等待 provider 终态:logical 与当前 attempt 同时存在。
Responding { logical: LogicalTurn, attempt: A },
/// logical turn 仍在,但当前 attempt 已被取走去结算或重绑,新 attempt 未就位。
/// 配额透明重试期间就处于这个状态。
Replanning { logical: LogicalTurn },
}
impl<A> ResponsesTurnState<A> {
/// 上游是否有一个正在进行的 response。取代原来的 `response_in_flight` 字段。
pub(super) const fn response_in_flight(&self) -> bool {
matches!(self, Self::Responding { .. })
}
/// 是否可以接受一条新的客户端 `response.create`。
///
/// `Replanning` 也要拒绝:那时旧 attempt 的结算/重绑还没收尾。实际上
/// 透明重试整段都在 relay loop 的上游分支里同步完成,此时不会读客户端帧,
/// 所以这条相对原来的 `response_in_flight` 判断没有行为差异。
pub(super) const fn accepts_new_response_create(&self) -> bool {
matches!(self, Self::Idle)
}
pub(super) const fn logical(&self) -> Option<&LogicalTurn> {
match self {
Self::Idle => None,
Self::Responding { logical, .. } | Self::Replanning { logical } => Some(logical),
}
}
pub(super) fn logical_mut(&mut self) -> Option<&mut LogicalTurn> {
match self {
Self::Idle => None,
Self::Responding { logical, .. } | Self::Replanning { logical } => Some(logical),
}
}
pub(super) const fn attempt(&self) -> Option<&A> {
match self {
Self::Responding { attempt, .. } => Some(attempt),
Self::Idle | Self::Replanning { .. } => None,
}
}
pub(super) fn attempt_mut(&mut self) -> Option<&mut A> {
match self {
Self::Responding { attempt, .. } => Some(attempt),
Self::Idle | Self::Replanning { .. } => None,
}
}
/// `Idle` → `Responding`:装上一个新 logical turn 及其首个 attempt。
///
/// 与原来的 `active_turn = Some(..)` 语义一致:若此刻竟持有旧 attempt,
/// 它会被丢弃并由自身的 drop guard 兜底结算,而不是静默泄漏。
pub(super) fn begin(&mut self, logical: LogicalTurn, attempt: A) {
debug_assert!(
self.accepts_new_response_create(),
"a new logical turn must only begin on an idle connection"
);
*self = Self::Responding { logical, attempt };
}
/// `Responding` → `Replanning`:把当前 attempt 交给调用方结算,保留 logical turn。
pub(super) fn detach_attempt(&mut self) -> Option<A> {
match std::mem::replace(self, Self::Idle) {
Self::Responding { logical, attempt } => {
*self = Self::Replanning { logical };
Some(attempt)
}
state @ (Self::Idle | Self::Replanning { .. }) => {
*self = state;
None
}
}
}
/// `Replanning` → `Responding`:同一 logical turn 的下一个 attempt 就位。
///
/// 状态不符时把 attempt 交还调用方,避免静默丢弃一条已经写了 pending usage
/// 行、占着 candidate 和 pool key lease 的 attempt。
pub(super) fn resume(&mut self, attempt: A) -> Result<(), A> {
match std::mem::replace(self, Self::Idle) {
Self::Replanning { logical } => {
*self = Self::Responding { logical, attempt };
Ok(())
}
state => {
*self = state;
Err(attempt)
}
}
}
/// `Responding`/`Replanning` → `Idle`:logical turn 结束,交出待结算的 attempt。
///
/// 取代原来「`active_turn.take()` + 在每个出口手写 `active_response_create = None`」
/// 的组合:清理只有这一个出口,漏清不再可能。
pub(super) fn end(&mut self) -> Option<A> {
match std::mem::replace(self, Self::Idle) {
Self::Responding { attempt, .. } => Some(attempt),
Self::Idle | Self::Replanning { .. } => None,
}
}
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::{LogicalTurn, ResponsesTurnState};
/// attempt 的测试替身:只需要能被 move,不需要 AppState 或真实 socket。
#[derive(Debug, PartialEq, Eq)]
struct FakeAttempt(u32);
fn logical() -> LogicalTurn {
LogicalTurn::new(
json!({"type": "response.create", "model": "gpt-5.6-sol"}),
7,
"logical-turn".to_string(),
)
}
#[test]
fn idle_has_no_turn_and_accepts_a_new_response_create() {
let state = ResponsesTurnState::<FakeAttempt>::Idle;
assert!(!state.response_in_flight());
assert!(state.accepts_new_response_create());
assert!(state.logical().is_none());
assert!(state.attempt().is_none());
}
#[test]
fn beginning_a_turn_makes_the_response_in_flight_and_blocks_a_second_one() {
let mut state = ResponsesTurnState::Idle;
state.begin(logical(), FakeAttempt(1));
assert!(state.response_in_flight());
assert!(!state.accepts_new_response_create());
assert_eq!(state.logical().map(|logical| logical.turn_index), Some(7));
assert_eq!(state.attempt(), Some(&FakeAttempt(1)));
}
#[test]
fn detaching_an_attempt_keeps_the_logical_turn_but_ends_the_in_flight_response() {
let mut state = ResponsesTurnState::Idle;
state.begin(logical(), FakeAttempt(1));
assert_eq!(state.detach_attempt(), Some(FakeAttempt(1)));
assert!(!state.response_in_flight());
assert!(!state.accepts_new_response_create());
assert!(state.attempt().is_none());
assert_eq!(
state.logical().map(|logical| logical.logical_turn_id.clone()),
Some("logical-turn".to_string())
);
// 幂等:已经没有 attempt 了,再取一次不会伪造一个出来。
assert_eq!(state.detach_attempt(), None);
}
#[test]
fn resuming_replaces_the_attempt_of_the_same_logical_turn() {
let mut state = ResponsesTurnState::Idle;
state.begin(logical(), FakeAttempt(1));
state
.logical_mut()
.expect("a responding turn always has its logical turn")
.turn_attempt = 2;
let _ = state.detach_attempt();
assert_eq!(state.resume(FakeAttempt(2)), Ok(()));
assert!(state.response_in_flight());
assert_eq!(state.attempt(), Some(&FakeAttempt(2)));
assert_eq!(state.logical().map(|logical| logical.turn_attempt), Some(2));
}
#[test]
fn resuming_without_a_logical_turn_hands_the_attempt_back() {
let mut state = ResponsesTurnState::Idle;
assert_eq!(state.resume(FakeAttempt(9)), Err(FakeAttempt(9)));
assert!(state.accepts_new_response_create());
state.begin(logical(), FakeAttempt(1));
assert_eq!(state.resume(FakeAttempt(9)), Err(FakeAttempt(9)));
assert_eq!(state.attempt(), Some(&FakeAttempt(1)));
}
#[test]
fn ending_a_turn_clears_both_the_logical_turn_and_the_attempt() {
let mut state = ResponsesTurnState::Idle;
state.begin(logical(), FakeAttempt(1));
assert_eq!(state.end(), Some(FakeAttempt(1)));
assert!(state.accepts_new_response_create());
assert!(state.logical().is_none());
// 从 Replanning 结束时没有 attempt 要交出,但 logical turn 同样必须清掉。
state.begin(logical(), FakeAttempt(2));
let _ = state.detach_attempt();
assert_eq!(state.end(), None);
assert!(state.accepts_new_response_create());
assert!(state.logical().is_none());
}
#[test]
fn quota_retry_safety_lives_on_the_logical_turn() {
let mut state = ResponsesTurnState::Idle;
state.begin(logical(), FakeAttempt(1));
assert_eq!(
state
.logical()
.and_then(LogicalTurn::quota_retry_block_reason),
None
);
state
.logical_mut()
.expect("logical turn")
.mark_retry_unsafe("standard_response_event");
assert_eq!(
state
.logical()
.and_then(LogicalTurn::quota_retry_block_reason),
Some("standard_response_event")
);
// 重绑后仍是同一个 logical turn,重放安全结论不能被 attempt 轮换洗掉。
let _ = state.detach_attempt();
assert_eq!(state.resume(FakeAttempt(2)), Ok(()));
assert_eq!(
state
.logical()
.and_then(LogicalTurn::quota_retry_block_reason),
Some("standard_response_event")
);
}
}
@@ -9,6 +9,7 @@ use super::adapter::ResponsesWebSocketProtocolAdapter;
use super::binding::{UpstreamBindingIdentity, UpstreamBindingIdentityError};
use super::request::planned_response_create_event;
use super::state::{BoundResponsesConnection, ExhaustedResponsesWebSocketExclusions};
use super::turn_state::ResponsesTurnState;
use crate::ai_serving::{AiExecutionDecision, ResponsesWebSocketBodyNormalization};
use crate::handlers::proxy::websocket::session::RESPONSES_WEBSOCKET_SESSION_LIMITS;
use crate::handlers::proxy::websocket::transport::{
@@ -115,12 +116,12 @@ async fn bind_responses_upstream_inner(
adapter,
client_model,
provider_model,
response_in_flight: true,
decision_template: decision.clone(),
body_normalization: normalization,
binding_identity,
active_turn: None,
active_response_create: None,
// 首条 response.create 已经发出,但这一轮的 logical turn 和 attempt 由调用方
// 通过 `ResponsesTurnState::begin` 装上:绑定本身不持有记账状态。
turn_state: ResponsesTurnState::Idle,
next_turn_index: 2,
upstream_response_headers: upstream.response_headers,
pending_adapter_drain: None,