fix(ws): enforce absolute upstream handshake and initial-message deadlines

This commit is contained in:
AAEE86
2026-08-17 14:52:39 +08:00
committed by ZheFox
parent f70ae68273
commit 3b036299d4
2 changed files with 402 additions and 9 deletions
@@ -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<S>(
socket: &mut S,
deadline_budget: std::time::Duration,
) -> Result<(String, Value), InitialMessageError>
where
S: futures_util::Stream<Item = Result<AxumWsMessage, axum::Error>>
+ futures_util::Sink<AxumWsMessage, Error = axum::Error>
+ 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<axum::extract::ws::Message>,
pong_tx: tokio::sync::mpsc::UnboundedSender<axum::extract::ws::Message>,
}
struct FakeSocketPair {
tx: tokio::sync::mpsc::Sender<axum::extract::ws::Message>,
pong_rx: tokio::sync::mpsc::UnboundedReceiver<axum::extract::ws::Message>,
}
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<axum::extract::ws::Message, axum::Error>;
fn poll_next(
mut self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<Self::Item>> {
self.rx.poll_recv(cx).map(|opt| opt.map(Ok))
}
}
impl futures_util::Sink<axum::extract::ws::Message> for FakeSocket {
type Error = axum::Error;
fn poll_ready(
self: std::pin::Pin<&mut Self>,
_cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), Self::Error>> {
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<Result<(), Self::Error>> {
std::task::Poll::Ready(Ok(()))
}
fn poll_close(
self: std::pin::Pin<&mut Self>,
_cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), Self::Error>> {
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");
}
}
@@ -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<BoundResponsesConnection, &'static str> {
// 绝对 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<BoundResponsesConnection, &'static str> {
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"
);
}
}