fix: harden websocket and soap adapters
This commit is contained in:
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user