use crate::ws::error::*;
use crate::ws::events::WebsocketEvent;
use crate::ws::worker::{ControlMessage, WorkerLoop};
pub struct WebSocketClient {
config: crate::config::WebSocketConfig,
tls_config: Option<crate::config::TLSConfig>,
callback: Option<crate::ws::MessageCallback>,
control_tx: Option<tokio::sync::mpsc::UnboundedSender<ControlMessage>>,
worker_handle: Option<tokio::task::JoinHandle<WebsocketResult<()>>>,
is_connected: std::sync::Arc<tokio::sync::RwLock<bool>>,
}
impl WebSocketClient {
pub fn new(
config: crate::config::WebSocketConfig,
tls_config: Option<crate::config::TLSConfig>,
) -> Self {
Self {
config,
tls_config,
callback: None,
control_tx: None,
worker_handle: None,
is_connected: std::sync::Arc::new(tokio::sync::RwLock::new(false)),
}
}
pub fn on_message<F>(&mut self, callback: F)
where
F: Fn(WebsocketEvent) + Send + Sync + 'static,
{
self.callback = Some(std::sync::Arc::new(callback));
}
pub async fn start_background(&mut self) -> WebsocketResult<()> {
if self.worker_handle.is_some() {
return Err(WebsocketError::AlreadyConnected);
}
let (control_tx, control_rx) = tokio::sync::mpsc::unbounded_channel();
self.control_tx = Some(control_tx);
let worker_loop = WorkerLoop::new(
self.config.clone(),
self.tls_config.clone(),
self.callback.clone(),
std::sync::Arc::clone(&self.is_connected),
);
let worker_handle = tokio::spawn(async move { worker_loop.run(control_rx).await });
self.worker_handle = Some(worker_handle);
Ok(())
}
pub async fn start_blocking(&mut self) -> WebsocketResult<()> {
let (control_tx, control_rx) = tokio::sync::mpsc::unbounded_channel();
self.control_tx = Some(control_tx);
let worker_loop = WorkerLoop::new(
self.config.clone(),
self.tls_config.clone(),
self.callback.clone(),
std::sync::Arc::clone(&self.is_connected),
);
worker_loop.run(control_rx).await
}
pub async fn stop_background(&mut self) -> WebsocketResult<()> {
if let Some(tx) = &self.control_tx {
let _ = tx.send(ControlMessage::Stop);
}
if let Some(handle) = self.worker_handle.take() {
let _ = tokio::time::timeout(std::time::Duration::from_secs(5), handle).await;
}
self.control_tx = None;
*self.is_connected.write().await = false;
Ok(())
}
pub async fn is_connected(&self) -> bool {
*self.is_connected.read().await
}
pub async fn reconnect(&self) -> WebsocketResult<()> {
if let Some(tx) = &self.control_tx {
tx.send(ControlMessage::Reconnect)
.map_err(|_| WebsocketError::ChannelError)?;
Ok(())
} else {
Err(WebsocketError::NotConnected)
}
}
}
impl Drop for WebSocketClient {
fn drop(&mut self) {
if let Some(tx) = &self.control_tx {
let _ = tx.send(ControlMessage::Stop);
}
}
}
impl std::fmt::Debug for WebSocketClient {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("WebsocketClient")
.field("url", &self.config.url)
.field("is_connected", &self.is_connected)
.field("has_tls_config", &self.tls_config.is_some())
.finish()
}
}