mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-04 16:37:46 +08:00
fix(ws): enforce absolute upstream handshake and initial-message deadlines
This commit is contained in:
@@ -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"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user