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 5feb17bf4..2bc8dc055 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/session.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/session.rs @@ -12,7 +12,6 @@ use axum::body::Bytes; use axum::extract::ws::{Message as AxumWsMessage, WebSocket}; use futures_util::{SinkExt, StreamExt}; use serde_json::Value; -use tokio::time::timeout; use uuid::Uuid; use super::adapter::resolve_responses_websocket_adapter; @@ -41,7 +40,7 @@ use crate::handlers::proxy::websocket::session::{ RESPONSES_WEBSOCKET_SESSION_LIMITS, WEBSOCKET_LOG_TRANSPORT, }; use crate::handlers::proxy::websocket::transport::{ - close_client_socket, send_client_message, send_gateway_error, send_gateway_error_with_status, + close_client_socket, send_gateway_error, send_gateway_error_with_status, }; use crate::orchestration::release_pool_key_lease_from_report_context; use crate::AppState; @@ -392,23 +391,50 @@ pub(super) async fn run_responses_websocket( await_pending_adapter_observation(&mut bound).await; } +/// 等待客户端发送第一条 response.create 事件。 +/// 使用绝对 deadline:从函数入口起计算一次截止时间,Ping/Pong 只会被正常回复, +/// 但不会重置计时器。防止客户端通过周期性 Ping 无限占用 connection permit。 async fn receive_initial_response_create( client_socket: &mut WebSocket, ) -> Result<(String, Value), InitialMessageError> { + receive_initial_response_create_with_deadline( + client_socket, + RESPONSES_WEBSOCKET_SESSION_LIMITS.initial_message_timeout, + ) + .await +} + +/// 核心循环:在绝对 deadline 内等待客户端发送 response.create。 +/// 泛型约束允许测试注入 fake socket,驱动真实逻辑。 +/// +/// - `deadline_budget`:从调用时刻起的最长等待时间,全循环共享同一截止时刻。 +/// - Ping 帧被回复 Pong 但不重置计时器。 +/// - Pong / 非法帧 / Close 按协议处理。 +async fn receive_initial_response_create_with_deadline( + socket: &mut S, + deadline_budget: std::time::Duration, +) -> Result<(String, Value), InitialMessageError> +where + S: futures_util::Stream> + + futures_util::Sink + + Unpin, +{ + use futures_util::{SinkExt as _, StreamExt as _}; + + // 绝对 deadline:入口计算一次,后续所有迭代共享,Ping/Pong 不会重启 + let deadline = tokio::time::Instant::now() + deadline_budget; loop { - let message = timeout( - RESPONSES_WEBSOCKET_SESSION_LIMITS.initial_message_timeout, - client_socket.next(), - ) - .await - .map_err(|_| InitialMessageError::TimedOut)?; + let message = tokio::time::timeout_at(deadline, socket.next()) + .await + .map_err(|_| InitialMessageError::TimedOut)?; let Some(message) = message else { return Err(InitialMessageError::ClientClosed); }; let message = message.map_err(|_| InitialMessageError::ClientRead)?; match message { AxumWsMessage::Ping(payload) => { - send_client_message(client_socket, AxumWsMessage::Pong(payload)) + socket + .send(AxumWsMessage::Pong(payload)) .await .map_err(|_| InitialMessageError::ClientRead)?; } @@ -1143,4 +1169,160 @@ mod tests { pending_turn_finalization: None, } } + + /// 用 mpsc 驱动的 FakeSocket,实现 Stream + Sink 两个 trait。 + /// 测试侧通过 tx 注入消息,通过 pong_rx 观察 Pong 回包。 + struct FakeSocket { + rx: tokio::sync::mpsc::Receiver, + pong_tx: tokio::sync::mpsc::UnboundedSender, + } + + struct FakeSocketPair { + tx: tokio::sync::mpsc::Sender, + pong_rx: tokio::sync::mpsc::UnboundedReceiver, + } + + fn fake_socket() -> (FakeSocket, FakeSocketPair) { + let (tx, rx) = tokio::sync::mpsc::channel(16); + let (pong_tx, pong_rx) = tokio::sync::mpsc::unbounded_channel(); + (FakeSocket { rx, pong_tx }, FakeSocketPair { tx, pong_rx }) + } + + impl futures_util::Stream for FakeSocket { + type Item = Result; + + fn poll_next( + mut self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + self.rx.poll_recv(cx).map(|opt| opt.map(Ok)) + } + } + + impl futures_util::Sink for FakeSocket { + type Error = axum::Error; + + fn poll_ready( + self: std::pin::Pin<&mut Self>, + _cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + std::task::Poll::Ready(Ok(())) + } + + fn start_send( + self: std::pin::Pin<&mut Self>, + item: axum::extract::ws::Message, + ) -> Result<(), Self::Error> { + let _ = self.pong_tx.send(item); + Ok(()) + } + + fn poll_flush( + self: std::pin::Pin<&mut Self>, + _cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + std::task::Poll::Ready(Ok(())) + } + + fn poll_close( + self: std::pin::Pin<&mut Self>, + _cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + std::task::Poll::Ready(Ok(())) + } + } + + /// 验证 receive_initial_response_create_with_deadline 的绝对 deadline: + /// 客户端周期性发送 Ping 帧不会重置计时器,deadline 到期后返回 TimedOut。 + /// 这直接驱动真实的循环逻辑,如果改回每次迭代 timeout(budget, ...) 则会变红。 + /// + /// 设计思路:deadline_budget = 80ms,Ping 每 30ms 发一次且永不停止。 + /// - 绝对 deadline:~80ms 后函数返回 TimedOut(即使 Ping 仍在到来)。 + /// - 每次迭代 timeout(80ms, ...):每个 Ping 在 30ms 内到达 < 80ms,函数 + /// 永不超时,500ms 后外层 timeout 判定测试失败。 + #[tokio::test] + async fn initial_message_times_out_despite_periodic_pings() { + use super::{receive_initial_response_create_with_deadline, InitialMessageError}; + + let (mut fake, pair) = fake_socket(); + + let handle = tokio::spawn(async move { + receive_initial_response_create_with_deadline(&mut fake, Duration::from_millis(80)) + .await + }); + + // 持续发送 Ping,间隔 30ms,永不停止(直到被测函数返回导致 rx drop) + let ping_task = tokio::spawn(async move { + let mut i = 0u8; + loop { + tokio::time::sleep(Duration::from_millis(30)).await; + if pair + .tx + .send(axum::extract::ws::Message::Ping(vec![i].into())) + .await + .is_err() + { + break; + } + i = i.wrapping_add(1); + } + }); + + // 绝对 deadline 应在 ~80ms 后触发;给 500ms 宽限等待结果。 + let result = tokio::time::timeout(Duration::from_millis(500), handle) + .await + .expect("function should return within 500ms (absolute deadline = 80ms)") + .expect("task should not panic"); + + ping_task.abort(); + + assert!( + matches!(result, Err(InitialMessageError::TimedOut)), + "expected TimedOut after absolute deadline, got: {result:?}" + ); + } + + /// 验证 deadline 内收到合法 response.create 时正常返回,Ping 被正确回复 Pong。 + #[tokio::test] + async fn initial_message_succeeds_within_deadline() { + use super::{receive_initial_response_create_with_deadline, InitialMessageError}; + + let (mut fake, mut pair) = fake_socket(); + + let handle = tokio::spawn(async move { + receive_initial_response_create_with_deadline(&mut fake, Duration::from_secs(5)).await + }); + + // 先发一个 Ping,验证 Pong 回包且不影响后续解析 + pair.tx + .send(axum::extract::ws::Message::Ping(vec![42].into())) + .await + .unwrap(); + let pong = tokio::time::timeout(Duration::from_secs(1), pair.pong_rx.recv()) + .await + .expect("should receive pong within 1s") + .expect("pong channel should not close"); + assert!( + matches!(pong, axum::extract::ws::Message::Pong(ref data) if data.as_ref() == [42]), + "expected Pong([42]), got: {pong:?}" + ); + + // 发送合法的 response.create + let event_text = r#"{"type":"response.create","model":"gpt-4o"}"#; + pair.tx + .send(axum::extract::ws::Message::Text( + event_text.to_string().into(), + )) + .await + .unwrap(); + + let result = tokio::time::timeout(Duration::from_secs(2), handle) + .await + .expect("handle should finish within 2s") + .expect("task should not panic"); + let (text, event) = result.expect("should return Ok"); + assert_eq!(text, event_text); + assert_eq!(event["type"], "response.create"); + assert_eq!(event["model"], "gpt-4o"); + } } 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 c760c5cfb..d7238164f 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/upstream.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/upstream.rs @@ -1,5 +1,7 @@ //! Physical upstream WebSocket binding and transport helpers. +use std::time::Duration; + use serde_json::Value; use wreq::ws::message::Message as WreqWsMessage; @@ -13,11 +15,53 @@ use crate::handlers::proxy::websocket::transport::{ close_upstream_socket, connect_upstream_websocket, send_upstream_message, }; +/// 上游 WebSocket 握手的默认绝对 deadline(30 秒)。 +/// 覆盖 DNS → TCP connect → TLS → HTTP 101 Upgrade → 发送首条 event 的完整链路。 +/// 如果 decision 配置了更短的 first_byte_ms 或 total_ms,取其与此值的较小者。 +const DEFAULT_UPSTREAM_HANDSHAKE_DEADLINE_MS: u64 = 30_000; + +/// 从 decision.timeouts 推导实际 handshake 绝对 deadline。 +/// 取 first_byte_ms / total_ms / DEFAULT 三者中的最小正值。 +pub(super) fn resolve_upstream_handshake_deadline( + decision: &AiExecutionDecision, +) -> Duration { + let mut deadline_ms = DEFAULT_UPSTREAM_HANDSHAKE_DEADLINE_MS; + if let Some(timeouts) = decision.timeouts.as_ref() { + if let Some(first_byte_ms) = timeouts.first_byte_ms.filter(|v| *v > 0) { + deadline_ms = deadline_ms.min(first_byte_ms); + } + if let Some(total_ms) = timeouts.total_ms.filter(|v| *v > 0) { + deadline_ms = deadline_ms.min(total_ms); + } + } + Duration::from_millis(deadline_ms) +} + pub(super) async fn bind_responses_upstream( decision: &AiExecutionDecision, normalization: ResponsesWebSocketBodyNormalization, initial_event: &Value, adapter: &'static dyn ResponsesWebSocketProtocolAdapter, +) -> Result { + // 绝对 deadline:从此刻起必须在限定时间内完成握手 + 首条事件发送, + // 防止慢 TLS / 慢 HTTP Upgrade 无限占用 connection permit。 + let handshake_deadline = resolve_upstream_handshake_deadline(decision); + tokio::time::timeout(handshake_deadline, bind_responses_upstream_inner( + decision, + normalization, + initial_event, + adapter, + )) + .await + .map_err(|_| "responses_websocket_upstream_handshake_timeout")? +} + +/// 实际执行握手 + 首条事件发送的内部函数,由外层 timeout 包裹。 +async fn bind_responses_upstream_inner( + decision: &AiExecutionDecision, + normalization: ResponsesWebSocketBodyNormalization, + initial_event: &Value, + adapter: &'static dyn ResponsesWebSocketProtocolAdapter, ) -> Result { let binding_identity = UpstreamBindingIdentity::from_decision(adapter, decision).map_err(|error| match error { @@ -111,3 +155,170 @@ pub(super) fn decision_reuses_bound_upstream( .map(|identity| bound.binding_identity == identity) .unwrap_or(false) } + +#[cfg(test)] +mod tests { + use std::time::Duration; + + use aether_contracts::ExecutionTimeouts; + + use crate::ai_serving::AiExecutionDecision; + + use super::{resolve_upstream_handshake_deadline, DEFAULT_UPSTREAM_HANDSHAKE_DEADLINE_MS}; + + fn sample_decision() -> AiExecutionDecision { + AiExecutionDecision { + action: "local".to_string(), + decision_kind: None, + execution_strategy: None, + conversion_mode: None, + request_id: None, + candidate_id: None, + provider_name: None, + provider_type: Some("custom".to_string()), + provider_id: None, + endpoint_id: None, + key_id: None, + upstream_base_url: None, + upstream_url: Some("https://example.test/v1/responses".to_string()), + provider_request_method: None, + auth_header: None, + auth_value: None, + provider_api_format: Some("openai:responses".to_string()), + client_api_format: Some("openai:responses".to_string()), + provider_contract: None, + client_contract: None, + model_name: None, + mapped_model: Some("provider-model".to_string()), + prompt_cache_key: None, + extra_headers: std::collections::BTreeMap::new(), + provider_request_headers: std::collections::BTreeMap::new(), + provider_request_body: None, + provider_request_body_base64: None, + content_type: None, + content_encoding: None, + request_gzip: None, + proxy: None, + transport_profile: None, + timeouts: None, + upstream_is_stream: true, + report_kind: None, + report_context: None, + auth_context: None, + } + } + + #[test] + fn handshake_deadline_defaults_to_30s_without_configured_timeouts() { + let decision = sample_decision(); + let deadline = resolve_upstream_handshake_deadline(&decision); + assert_eq!( + deadline, + Duration::from_millis(DEFAULT_UPSTREAM_HANDSHAKE_DEADLINE_MS) + ); + } + + #[test] + fn handshake_deadline_uses_first_byte_ms_when_shorter_than_default() { + let mut decision = sample_decision(); + decision.timeouts = Some(ExecutionTimeouts { + first_byte_ms: Some(10_000), + total_ms: Some(60_000), + ..ExecutionTimeouts::default() + }); + let deadline = resolve_upstream_handshake_deadline(&decision); + assert_eq!(deadline, Duration::from_millis(10_000)); + } + + #[test] + fn handshake_deadline_uses_total_ms_when_shorter_than_first_byte_and_default() { + let mut decision = sample_decision(); + decision.timeouts = Some(ExecutionTimeouts { + first_byte_ms: Some(25_000), + total_ms: Some(8_000), + ..ExecutionTimeouts::default() + }); + let deadline = resolve_upstream_handshake_deadline(&decision); + assert_eq!(deadline, Duration::from_millis(8_000)); + } + + #[test] + fn handshake_deadline_ignores_zero_values() { + let mut decision = sample_decision(); + decision.timeouts = Some(ExecutionTimeouts { + first_byte_ms: Some(0), + total_ms: Some(0), + ..ExecutionTimeouts::default() + }); + let deadline = resolve_upstream_handshake_deadline(&decision); + assert_eq!( + deadline, + Duration::from_millis(DEFAULT_UPSTREAM_HANDSHAKE_DEADLINE_MS) + ); + } + + #[test] + fn handshake_deadline_does_not_exceed_default_even_with_larger_configured_values() { + let mut decision = sample_decision(); + decision.timeouts = Some(ExecutionTimeouts { + first_byte_ms: Some(120_000), + total_ms: Some(600_000), + ..ExecutionTimeouts::default() + }); + let deadline = resolve_upstream_handshake_deadline(&decision); + assert_eq!( + deadline, + Duration::from_millis(DEFAULT_UPSTREAM_HANDSHAKE_DEADLINE_MS) + ); + } + + #[tokio::test] + async fn bind_responses_upstream_times_out_against_stalled_server() { + use super::bind_responses_upstream; + use crate::ai_serving::ResponsesWebSocketBodyNormalization; + use crate::handlers::proxy::websocket::responses::adapter::resolve_responses_websocket_adapter; + use serde_json::json; + + // 启动一个接受 TCP 连接但永不完成 HTTP Upgrade 的 mock 服务器 + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("mock listener should bind"); + let addr = listener.local_addr().expect("should have local addr"); + let _server = tokio::spawn(async move { + loop { + let (socket, _) = listener.accept().await.unwrap(); + // 接受连接但不发送任何 HTTP 响应,模拟 stalled handshake + tokio::spawn(async move { + let _hold = socket; + tokio::time::sleep(Duration::from_secs(300)).await; + }); + } + }); + + let mut decision = sample_decision(); + decision.upstream_url = Some(format!("http://{addr}/v1/responses")); + // 设置极短的 deadline 以便测试快速完成 + decision.timeouts = Some(ExecutionTimeouts { + first_byte_ms: Some(100), + total_ms: Some(200), + ..ExecutionTimeouts::default() + }); + decision.provider_request_body = Some(json!({"model": "test-model"})); + + let adapter = resolve_responses_websocket_adapter( + crate::orchestration::ResponsesWebSocketAdapter::Standard, + ); + let result = bind_responses_upstream( + &decision, + ResponsesWebSocketBodyNormalization::for_tests("test-model"), + &json!({"type": "response.create", "model": "test-model"}), + adapter, + ) + .await; + + assert_eq!( + result.err().expect("bind should fail with timeout"), + "responses_websocket_upstream_handshake_timeout" + ); + } +}