use futures_util::{SinkExt, StreamExt};
use tokio_tungstenite::{connect_async, tungstenite::protocol::Message, MaybeTlsStream};
use url::Url;
#[async_trait::async_trait]
pub trait WebSocketClientTrait {
async fn connect(&mut self) -> anyhow::Result<()>;
async fn send(&mut self, msg: String) -> anyhow::Result<()>;
async fn receive(&mut self) -> anyhow::Result<Option<String>>;
async fn disconnect(&mut self) -> anyhow::Result<()>;
}
pub struct WebSocketClient {
url: Url,
ws_stream: Option<tokio_tungstenite::WebSocketStream<MaybeTlsStream<tokio::net::TcpStream>>>,
}
impl WebSocketClient {
pub fn new(url: &str) -> anyhow::Result<Self> {
let url = Url::parse(url)?;
Ok(Self {
url,
ws_stream: None,
})
}
}
#[async_trait::async_trait]
impl WebSocketClientTrait for WebSocketClient {
async fn connect(&mut self) -> anyhow::Result<()> {
let (ws_stream, _) = connect_async(self.url.clone()).await?;
self.ws_stream = Some(ws_stream);
Ok(())
}
async fn send(&mut self, msg: String) -> anyhow::Result<()> {
if let Some(ws) = &mut self.ws_stream {
ws.send(Message::Text(msg)).await?;
Ok(())
} else {
Err(anyhow::anyhow!("WebSocket not connected"))
}
}
async fn receive(&mut self) -> anyhow::Result<Option<String>> {
if let Some(ws) = &mut self.ws_stream {
match ws.next().await {
Some(Ok(Message::Text(text))) => Ok(Some(text)),
Some(Ok(Message::Close(_))) | None => Ok(None),
Some(Ok(_)) => self.receive().await, Some(Err(e)) => Err(anyhow::anyhow!(e)),
}
} else {
Err(anyhow::anyhow!("WebSocket not connected"))
}
}
async fn disconnect(&mut self) -> anyhow::Result<()> {
if let Some(ws) = &mut self.ws_stream {
ws.close(None).await?;
self.ws_stream = None;
Ok(())
} else {
Err(anyhow::anyhow!("WebSocket not connected"))
}
}
}