Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -130,4 +130,4 @@ http-client.private.env.json
*.vsix
# ------- </VSCode> -------

specs/*
specs/*
57 changes: 36 additions & 21 deletions src/ws/connection.rs
Original file line number Diff line number Diff line change
Expand Up @@ -248,28 +248,35 @@ where
// Handle incoming messages
Some(msg) = read.next() => {
match msg {
Ok(Message::Text(text)) if text == "PONG" => {
_ = pong_tx.send(Instant::now());
}
Ok(Message::Text(text)) => {
#[cfg(feature = "tracing")]
tracing::trace!(%text, "Received WebSocket text message");

// Parse messages using the provided parser
match parser.parse(text.as_bytes()) {
Ok(messages) => {
for message in messages {
if text == "PONG" {
if pong_tx.send(Instant::now()).is_err() {
heartbeat_handle.abort();
return Err(Error::with_source(
Kind::WebSocket,
WsError::Timeout,
));
}
} else {
#[cfg(feature = "tracing")]
tracing::trace!(%text, "Received WebSocket text message");

// Parse messages using the provided parser
match parser.parse(text.as_bytes()) {
Ok(messages) => {
for message in messages {
#[cfg(feature = "tracing")]
tracing::trace!(?message, "Parsed WebSocket message");
_ = broadcast_tx.send(message);
}
}
Err(e) => {
#[cfg(feature = "tracing")]
tracing::trace!(?message, "Parsed WebSocket message");
_ = broadcast_tx.send(message);
tracing::warn!(%text, error = %e, "Failed to parse WebSocket message");
#[cfg(not(feature = "tracing"))]
let _: (&_, &_) = (&text, &e);
}
}
Err(e) => {
#[cfg(feature = "tracing")]
tracing::warn!(%text, error = %e, "Failed to parse WebSocket message");
#[cfg(not(feature = "tracing"))]
let _: (&_, &_) = (&text, &e);
}
}
}
Ok(Message::Close(_)) => {
Expand Down Expand Up @@ -300,9 +307,17 @@ where
}

// Handle PING requests from heartbeat loop
Some(()) = ping_rx.recv() => {
if write.send(Message::Text("PING".into())).await.is_err() {
break;
ping = ping_rx.recv() => {
match ping {
Some(()) => {
if write.send(Message::Text("PING".into())).await.is_err() {
heartbeat_handle.abort();
return Err(Error::with_source(Kind::WebSocket, WsError::ConnectionClosed));
}
}
None => {
return Err(Error::with_source(Kind::WebSocket, WsError::Timeout));
}
}
}

Expand Down
38 changes: 38 additions & 0 deletions tests/websocket.rs
Original file line number Diff line number Diff line change
Expand Up @@ -849,6 +849,44 @@ mod reconnection {
config
}

fn heartbeat_config() -> Config {
let mut config = config();
config.heartbeat_interval = Duration::from_millis(20);
config.heartbeat_timeout = Duration::from_millis(20);
config.reconnect.initial_backoff = Duration::from_millis(20);
config.reconnect.max_backoff = Duration::from_millis(50);
config
}

#[tokio::test]
async fn reconnects_when_heartbeat_pong_is_missing() {
let mut server = ReconnectableMockServer::start().await;
let endpoint = server.ws_url("/ws/market");

let client = Client::new(&endpoint, heartbeat_config()).unwrap();

let asset_id = payloads::asset_id();
let stream = client.subscribe_orderbook(vec![asset_id]).unwrap();
let mut stream = Box::pin(stream);

let sub_request = server.recv_subscription().await.unwrap();
assert!(sub_request.contains(&asset_id.to_string()));

let resub = server.recv_subscription().await;
assert!(
resub.is_some(),
"heartbeat timeout should force reconnect and re-subscription"
);
assert!(resub.unwrap().contains(&asset_id.to_string()));

server.send(&payloads::book().to_string());
let msg = timeout(Duration::from_secs(2), stream.next()).await;
assert!(
msg.is_ok() && msg.unwrap().is_some(),
"stream should receive messages after heartbeat reconnect"
);
}

#[tokio::test]
async fn resubscribes_and_receives_messages_after_reconnect() {
let mut server = ReconnectableMockServer::start().await;
Expand Down