use futures_util::{SinkExt, StreamExt};
use tokio_tungstenite::tungstenite::Message;
use rs_teststand_bridge::{Ack, Command, Error, MessageEvent};
mod backoff;
mod inbound;
pub use backoff::Backoff;
pub use inbound::Inbound;
fn transport(error: &tokio_tungstenite::tungstenite::Error) -> Error {
Error::Transport(std::io::Error::other(error.to_string()))
}
pub const MAX_CONTROL_PAYLOAD: usize = 125;
#[derive(Debug)]
pub struct Client {
socket: tokio_tungstenite::WebSocketStream<
tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>,
>,
}
impl Client {
pub async fn connect(address: &str) -> Result<Self, Error> {
let (socket, _) = tokio_tungstenite::connect_async(address)
.await
.map_err(|error| transport(&error))?;
Ok(Self { socket })
}
pub async fn connect_with_backoff(address: &str, backoff: Backoff) -> Result<Self, Error> {
tokio::time::sleep(Backoff::first_delay()).await;
let mut last = None;
for attempt in 0..backoff.attempts.max(1) {
match Self::connect(address).await {
Ok(client) => return Ok(client),
Err(error) => last = Some(error),
}
if attempt + 1 < backoff.attempts {
tokio::time::sleep(backoff.delay(attempt)).await;
}
}
Err(last.unwrap_or_else(|| {
Error::Transport(std::io::Error::other("no connection attempt was made"))
}))
}
pub async fn send(&mut self, command: &Command) -> Result<(), Error> {
let text = serde_json::to_string(command)?;
self.socket
.send(Message::Text(text.into()))
.await
.map_err(|error| transport(&error))
}
pub async fn next(&mut self) -> Result<Option<Inbound>, Error> {
while let Some(frame) = self.socket.next().await {
match frame.map_err(|error| transport(&error))? {
Message::Text(text) => return Inbound::parse(&text).map(Some),
Message::Close(frame) => {
let _ = self.socket.send(Message::Close(frame)).await;
return Ok(None);
}
_ => (),
}
}
Ok(None)
}
pub async fn close(mut self, mut observe: impl FnMut(Inbound)) -> Result<(), Error> {
self.socket
.send(Message::Close(None))
.await
.map_err(|error| transport(&error))?;
while let Some(frame) = self.socket.next().await {
match frame {
Ok(Message::Close(_)) | Err(_) => break,
Ok(Message::Text(text)) => {
if let Ok(inbound) = Inbound::parse(&text) {
observe(inbound);
}
}
Ok(_) => (),
}
}
Ok(())
}
pub async fn ping(&mut self, payload: Vec<u8>) -> Result<(), Error> {
if payload.len() > MAX_CONTROL_PAYLOAD {
return Err(Error::ControlFrameTooLarge {
bytes: payload.len(),
limit: MAX_CONTROL_PAYLOAD,
});
}
self.socket
.send(Message::Ping(payload.into()))
.await
.map_err(|error| transport(&error))
}
pub async fn request(
&mut self,
command: &Command,
mut observe: impl FnMut(&MessageEvent),
) -> Result<Ack, Error> {
self.send(command).await?;
while let Some(inbound) = self.next().await? {
match inbound {
Inbound::Ack(ack) => return Ok(ack),
Inbound::Event(event) => observe(&event),
}
}
Err(Error::Transport(std::io::Error::new(
std::io::ErrorKind::UnexpectedEof,
"the host closed before acknowledging",
)))
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use super::{Backoff, Inbound};
use rs_teststand_bridge::{Ack, MessageEvent};
#[test]
fn an_acknowledgement_is_recognized_by_its_command_field() {
let text = serde_json::to_string(&Ack::ok("start", "started")).unwrap_or_default();
let parsed = Inbound::parse(&text).ok();
assert!(
parsed.as_ref().and_then(Inbound::as_ack).is_some(),
"{text}"
);
}
#[test]
fn an_event_is_recognized_by_the_absence_of_one() {
let event = MessageEvent {
code: 2,
numeric: 0.0,
text: "stopped".to_owned(),
payload: None,
synchronous: false,
execution_id: Some(4),
};
let text = serde_json::to_string(&event).unwrap_or_default();
let parsed = Inbound::parse(&text).ok();
assert_eq!(
parsed.as_ref().and_then(Inbound::as_event).map(|e| e.code),
Some(2),
"{text}"
);
}
#[test]
fn the_delay_grows_and_then_stops_growing() {
let backoff = Backoff {
first: Duration::from_secs(1),
longest: Duration::from_secs(30),
attempts: 10,
};
assert!(backoff.delay(0) >= Duration::from_secs(1));
assert!(backoff.delay(1) >= Duration::from_secs(2));
assert!(backoff.delay(3) >= Duration::from_secs(8));
let ceiling = Duration::from_secs(30) + Duration::from_secs(30) / 4;
for attempt in 5..40 {
assert!(
backoff.delay(attempt) <= ceiling,
"attempt {attempt} exceeded the cap"
);
}
}
#[test]
fn the_first_attempt_is_delayed_by_a_random_amount() {
let first = Backoff::first_delay();
assert!(
first <= Duration::from_secs(5),
"{first:?} exceeds the range"
);
}
#[test]
fn the_delay_carries_jitter_above_the_plain_doubling() {
let backoff = Backoff::default();
assert!(backoff.delay(2) >= Duration::from_secs(4));
assert!(backoff.delay(2) <= Duration::from_secs(5));
}
#[test]
fn an_oversized_control_payload_is_refused_before_it_is_sent() {
assert_eq!(super::MAX_CONTROL_PAYLOAD, 125);
}
#[test]
fn a_frame_that_is_neither_shape_is_an_error_rather_than_a_guess() {
assert!(Inbound::parse("not json").is_err());
}
}