From 1c5ee5228c34c711c3337889001cab8cae7ce374 Mon Sep 17 00:00:00 2001 From: AAEE86 Date: Fri, 31 Jul 2026 17:31:17 +0800 Subject: [PATCH] =?UTF-8?q?refactor(ws):=20=E7=94=A8=20ResponsesTurnState?= =?UTF-8?q?=20=E6=94=B6=E6=95=9B=E8=BF=9E=E6=8E=A5=20turn=20=E7=8A=B6?= =?UTF-8?q?=E6=80=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 评审第 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。 --- .../proxy/websocket/responses/client.rs | 40 +-- .../proxy/websocket/responses/connection.rs | 77 ++--- .../proxy/websocket/responses/lifecycle.rs | 6 +- .../handlers/proxy/websocket/responses/mod.rs | 1 + .../proxy/websocket/responses/quota.rs | 27 +- .../proxy/websocket/responses/redaction.rs | 6 +- .../proxy/websocket/responses/session.rs | 98 +++--- .../proxy/websocket/responses/state.rs | 47 +-- .../proxy/websocket/responses/turn_state.rs | 303 ++++++++++++++++++ .../proxy/websocket/responses/upstream.rs | 7 +- 10 files changed, 428 insertions(+), 184 deletions(-) create mode 100644 apps/aether-gateway/src/handlers/proxy/websocket/responses/turn_state.rs diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/client.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/client.rs index 4c42cfba3..e3e447f24 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/client.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/client.rs @@ -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; diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/connection.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/connection.rs index 43b6d0bc2..8ed84d208 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/connection.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/connection.rs @@ -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 => {} - } -} diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/lifecycle.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/lifecycle.rs index f5967648d..15b41f417 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/lifecycle.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/lifecycle.rs @@ -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; } } diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/mod.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/mod.rs index ec1f059c1..d16929ed9 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/mod.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/mod.rs @@ -21,6 +21,7 @@ mod request; mod session; mod state; mod turn; +mod turn_state; mod upstream; use std::net::SocketAddr; diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/quota.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/quota.rs index 591c2e23a..e2501ff6e 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/quota.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/quota.rs @@ -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); } } diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/redaction.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/redaction.rs index 1c7394ce4..d3ae81540 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/redaction.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/redaction.rs @@ -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(), diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/session.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/session.rs index 1f75ece9c..990825c63 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/session.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/session.rs @@ -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, diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/state.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/state.rs index 74bc168bd..35851e981 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/state.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/state.rs @@ -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, - pub(super) active_response_create: Option, + /// 这条连接上「有没有正在进行的 logical turn」的唯一事实来源。 + pub(super) turn_state: ResponsesTurnState, pub(super) next_turn_index: u64, pub(super) upstream_response_headers: BTreeMap, pub(super) pending_adapter_drain: Option, @@ -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); - } -} diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/turn_state.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/turn_state.rs new file mode 100644 index 000000000..d2cdd9812 --- /dev/null +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/turn_state.rs @@ -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 { + /// 没有进行中的 logical turn。上游可能仍绑定,也可能已被 detach。 + Idle, + /// 一个 logical turn 正在等待 provider 终态:logical 与当前 attempt 同时存在。 + Responding { logical: LogicalTurn, attempt: A }, + /// logical turn 仍在,但当前 attempt 已被取走去结算或重绑,新 attempt 未就位。 + /// 配额透明重试期间就处于这个状态。 + Replanning { logical: LogicalTurn }, +} + +impl ResponsesTurnState { + /// 上游是否有一个正在进行的 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 { + 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 { + 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::::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") + ); + } +} diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/upstream.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/upstream.rs index d7238164f..97b018ba5 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/upstream.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/upstream.rs @@ -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,