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.
This commit is contained in:
ZheFox
2026-08-17 18:50:29 +08:00
parent 4a0775c4ea
commit c8118edf36
42 changed files with 4017 additions and 1276 deletions
@@ -6,13 +6,13 @@
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::{
spawn_responses_websocket_turn_finalization, ResponsesProviderAttempt,
ResponsesWebSocketTurnOutcome,
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,
@@ -76,11 +76,90 @@ impl std::ops::DerefMut for ActiveProviderAttempt {
}
}
/// 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.
@@ -93,11 +172,7 @@ impl Drop for ActiveProviderAttempt {
"gateway finalized a Responses WebSocket turn whose relay task went away"
);
handle.spawn(async move {
turn.finalize_detached(
&state,
ResponsesWebSocketTurnOutcome::relay_task_abandoned(),
)
.await;
turn.finalize_detached(&state, outcome).await;
});
}
}
@@ -113,20 +188,23 @@ pub(super) async fn finalize_active_turn(
outcome: ResponsesWebSocketTurnOutcome,
) {
if let Some(turn) = bound.turn_state.end() {
queue_turn_finalization(bound, state, turn.disarm(), outcome).await;
queue_turn_finalization(bound, state, turn, outcome).await;
}
}
pub(super) async fn queue_turn_finalization(
bound: &mut BoundResponsesConnection,
state: &AppState,
turn: ResponsesProviderAttempt,
turn: ActiveProviderAttempt,
outcome: ResponsesWebSocketTurnOutcome,
) {
await_pending_adapter_observation(bound).await;
await_pending_turn_finalization(bound).await;
bound.pending_turn_finalization =
Some(spawn_responses_websocket_turn_finalization(state.clone(), turn, outcome).await);
bound.pending_turn_finalization = Some(spawn_guarded_turn_finalization(
state.clone(),
turn,
outcome,
));
}
/// 「上一个 attempt 已经结算完毕」的凭证。
@@ -153,7 +231,7 @@ impl PreviousAttemptSettled {
pub(super) async fn settle_turn_finalization(
bound: &mut BoundResponsesConnection,
state: &AppState,
turn: ResponsesProviderAttempt,
turn: ActiveProviderAttempt,
outcome: ResponsesWebSocketTurnOutcome,
) -> PreviousAttemptSettled {
queue_turn_finalization(bound, state, turn, outcome).await;
@@ -161,42 +239,68 @@ pub(super) async fn settle_turn_finalization(
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(mut handle) = bound.pending_adapter_observation.take() {
match timeout(RESPONSES_WEBSOCKET_ADAPTER_OBSERVATION_TIMEOUT, &mut handle).await {
Ok(Err(error)) => {
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"
);
}
Ok(Ok(())) => {}
Err(_) => {
handle.abort();
let _ = handle.await;
warn!(
event_name = "responses_websocket_adapter_observation_timeout",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
timeout_ms = RESPONSES_WEBSOCKET_ADAPTER_OBSERVATION_TIMEOUT.as_millis() as u64,
"gateway stopped waiting for a Responses WebSocket adapter observation"
);
}
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) async fn finalize_unbound_turn(
pub(super) fn finalize_unbound_turn(
state: AppState,
turn: ResponsesProviderAttempt,
turn: ActiveProviderAttempt,
outcome: ResponsesWebSocketTurnOutcome,
) -> JoinHandle<()> {
spawn_responses_websocket_turn_finalization(state, turn, outcome).await
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<()>) {
@@ -228,6 +332,7 @@ 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 {
@@ -235,19 +340,13 @@ pub(super) async fn send_responses_websocket_turn_start_error(
} else {
("invalid_request_error", "gateway_request_not_allowed")
};
send_responses_websocket_error(
client_socket,
status.as_u16(),
error_type,
code,
message,
)
.await;
send_responses_websocket_error(client_socket, status_code, error_type, code, message)
.await;
}
GatewayError::AdmissionTimeout { .. } => {
send_responses_websocket_error(
client_socket,
503,
status_code,
"server_error",
"gateway_admission_timeout",
"Gateway capacity is busy; retry this response",
@@ -257,7 +356,7 @@ pub(super) async fn send_responses_websocket_turn_start_error(
GatewayError::LocalExecutionPlanningTimeout { .. } => {
send_responses_websocket_error(
client_socket,
504,
status_code,
"server_error",
"gateway_planning_timeout",
"Gateway planning timed out; retry this response",
@@ -267,7 +366,7 @@ pub(super) async fn send_responses_websocket_turn_start_error(
_ => {
send_responses_websocket_error(
client_socket,
500,
status_code,
"server_error",
"responses_websocket_turn_start_failed",
"Gateway could not start this response",
@@ -277,6 +376,15 @@ pub(super) async fn send_responses_websocket_turn_start_error(
}
}
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"),
@@ -292,7 +400,27 @@ mod tests {
use std::sync::Arc;
use std::time::Duration;
use super::await_turn_finalization_handle;
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 依赖的性质:结算是「等到落地」而不是「排进队列」。
///
@@ -348,4 +476,113 @@ mod tests {
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");
}
}