fix: harden websocket and soap adapters

This commit is contained in:
a.tolmachev
2026-04-11 01:37:24 +03:00
parent 2770c5935f
commit e60d848293
6 changed files with 365 additions and 60 deletions
+119 -24
View File
@@ -19,6 +19,11 @@ use crate::{
WebsocketWindowResponse,
};
enum WindowCollectionStatus {
WindowExpired,
MaxItemsReached,
}
#[derive(Clone, Debug)]
pub struct WebsocketAdapter;
@@ -60,33 +65,51 @@ impl WebsocketAdapter {
send_subscribe_message(&mut stream, target).await?;
let completed = collect_window(
let status = match collect_window(
&mut stream,
request.max_items,
deadline,
heartbeat.as_ref(),
&mut items,
)
.await?;
.await
{
Ok(status) => status,
Err(WebsocketAdapterError::ClosedEarly) => {
if attempts >= reconnect.max_attempts {
return Err(WebsocketAdapterError::ReconnectExhausted);
}
attempts = attempts.saturating_add(1);
reconnect_if_needed(&reconnect, attempts).await;
continue;
}
Err(error) => return Err(error),
};
if completed {
send_unsubscribe_message(&mut stream, target).await?;
return Ok(WebsocketWindowResponse {
status_code: 101,
headers: connected_headers,
body: serde_json::json!({
"items": items,
"done": true,
}),
});
match status {
WindowCollectionStatus::WindowExpired => {
send_unsubscribe_message(&mut stream, target).await?;
return Ok(WebsocketWindowResponse {
status_code: 101,
headers: connected_headers,
body: serde_json::json!({
"items": items,
"done": false,
}),
});
}
WindowCollectionStatus::MaxItemsReached => {
send_unsubscribe_message(&mut stream, target).await?;
return Ok(WebsocketWindowResponse {
status_code: 101,
headers: connected_headers,
body: serde_json::json!({
"items": items,
"done": true,
}),
});
}
}
if attempts >= reconnect.max_attempts {
return Err(WebsocketAdapterError::ReconnectExhausted);
}
attempts = attempts.saturating_add(1);
reconnect_if_needed(&reconnect, attempts).await;
}
}
}
@@ -189,12 +212,12 @@ async fn collect_window(
deadline: Instant,
heartbeat: Option<&HeartbeatPolicy>,
items: &mut Vec<Value>,
) -> Result<bool, WebsocketAdapterError> {
) -> Result<WindowCollectionStatus, WebsocketAdapterError> {
let mut heartbeat_deadline = heartbeat.map(|policy| Instant::now() + policy.interval);
loop {
if Instant::now() >= deadline {
return Ok(!items.is_empty());
return Ok(WindowCollectionStatus::WindowExpired);
}
let now = Instant::now();
@@ -212,17 +235,17 @@ async fn collect_window(
continue;
}
return Ok(!items.is_empty());
return Ok(WindowCollectionStatus::WindowExpired);
}
frame = read_next_frame(stream) => {
match frame? {
Some(value) => {
items.push(value);
if max_items.is_some_and(|limit| items.len() as u32 >= limit) {
return Ok(true);
return Ok(WindowCollectionStatus::MaxItemsReached);
}
}
None => return Ok(!items.is_empty()),
None => return Err(WebsocketAdapterError::ClosedEarly),
}
heartbeat_deadline = heartbeat.map(|policy| Instant::now() + policy.interval);
}
@@ -374,6 +397,39 @@ mod tests {
assert_eq!(response.body["items"].as_array().unwrap().len(), 3);
assert_eq!(response.body["items"][2]["seq"], 3);
assert_eq!(response.body["done"], true);
}
#[tokio::test]
async fn reconnects_after_partial_close_without_marking_done_early() {
let target_url = spawn_partial_close_server().await;
let adapter = WebsocketAdapter::new();
let target = WebsocketTarget {
url: target_url,
subprotocols: Vec::new(),
subscribe_message_template: None,
unsubscribe_message_template: None,
static_headers: BTreeMap::new(),
};
let response = adapter
.execute_window(
&target,
&WebsocketWindowRequest {
headers: BTreeMap::new(),
window_duration_ms: 1_000,
max_items: Some(3),
heartbeat_interval_ms: None,
reconnect_max_attempts: 2,
reconnect_backoff_ms: 10,
},
)
.await
.unwrap();
assert_eq!(response.body["items"].as_array().unwrap().len(), 3);
assert_eq!(response.body["items"][0]["seq"], 1);
assert_eq!(response.body["items"][2]["seq"], 3);
assert_eq!(response.body["done"], true);
}
async fn spawn_server(received: Arc<Mutex<Vec<Value>>>, close_early: bool) -> String {
@@ -439,4 +495,43 @@ mod tests {
format!("ws://{}", addr)
}
async fn spawn_partial_close_server() -> String {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
let mut accepted = 0_u32;
loop {
let (stream, _) = listener.accept().await.unwrap();
accepted = accepted.saturating_add(1);
tokio::spawn(async move {
let websocket =
accept_hdr_async(stream, |_request: &Request, response: Response| {
Ok(response)
})
.await
.unwrap();
let (mut sink, _source) = websocket.split();
let payloads = if accepted == 1 {
vec![json!({"seq": 1})]
} else {
vec![json!({"seq": 2}), json!({"seq": 3})]
};
for payload in payloads {
sink.send(tokio_tungstenite::tungstenite::Message::Text(
payload.to_string().into(),
))
.await
.unwrap();
}
let _ = sink.close().await;
});
}
});
format!("ws://{}", addr)
}
}