use crate::shared::{ClientError, Credentials};
use futures_util::{SinkExt, StreamExt};
use std::borrow::Cow;
use tokio::net::TcpStream;
use tokio_tungstenite::tungstenite::protocol::Message;
use tokio_tungstenite::{connect_async, MaybeTlsStream, WebSocketStream};
use tungstenite::client::IntoClientRequest;
type WsStream = WebSocketStream<MaybeTlsStream<TcpStream>>;
#[derive(Debug)]
pub struct WClient {
current_messages_size: usize,
address: String,
use_tls: bool,
username: String,
password: Option<String>,
ws_connection: Option<WsStream>,
}
impl WClient {
pub fn new(
address: &str,
credentials: Credentials,
use_tls: bool,
) -> Self {
Self {
current_messages_size: 0,
address: address.to_string(),
use_tls,
username: credentials.username,
password: credentials.password,
ws_connection: None
}
}
pub fn update_credentials(&mut self, credentials: Credentials) {
self.username = credentials.username;
self.password = credentials.password;
}
pub fn update_tls(&mut self, use_tls: bool) {
self.use_tls = use_tls;
}
pub fn update_address(&mut self, address: String) {
self.address = address;
}
fn build_url(&self) -> Result<String, ClientError> {
if self.address.starts_with("ws://") || self.address.starts_with("wss://") {
return Ok(self.address.to_string());
}
let scheme = if self.use_tls { "wss" } else { "ws" };
Ok(format!("{scheme}://{}/", self.address))
}
async fn get_ws(&self) -> Result<WsStream, ClientError> {
let url = self.build_url()?;
let (ws, _resp) = connect_async(url.into_client_request().unwrap())
.await
.map_err(|e| ClientError::TlsInitializationError(e.to_string()))?;
Ok(ws)
}
pub async fn prepare(&mut self) -> Result<(), ClientError> {
self.ws_connection = Some(self.get_ws().await?);
Ok(())
}
async fn check_connection(&self) -> Result<(), ClientError> {
if self.ws_connection.is_none() {
Err(ClientError::NoConnectionWRAC)
} else {
Ok(())
}
}
pub async fn register_user(&mut self) -> Result<(), ClientError> {
self.check_connection().await?;
if self.password.is_some() {
let mut ws = self.get_ws().await?;
let payload = format!(
"\x03{}\n{}",
self.username,
self.password.as_deref().unwrap()
);
ws.send(Message::Binary(payload.into()))
.await
.map_err(|e| ClientError::WsSendError(e.to_string()))?;
if let Some(Ok(Message::Binary(buf))) = ws.next().await {
return match buf.first() {
Some(0x01) => Err(ClientError::UsernameAlreadyTaken),
Some(code) => Err(ClientError::UnexpectedResponse(format!("0x{code:02x}"))),
None => Ok(()),
};
}
return Ok(());
}
Err(ClientError::NoPassword)
}
pub async fn fetch_messages_size(&mut self) -> Result<(), ClientError> {
self.check_connection().await?;
let ws = self.ws_connection.as_mut().unwrap();
ws.send(Message::Binary(vec![0x00].into()))
.await
.map_err(|e| ClientError::WsSendError(e.to_string()))?;
if let Some(Ok(msg)) = ws.next().await {
let txt: String = match msg {
Message::Text(t) => t.to_string(),
Message::Binary(b) => String::from_utf8_lossy(&b).into_owned(),
_ => String::new(),
};
self.current_messages_size = txt
.trim()
.parse::<usize>()
.map_err(|_| ClientError::ParseError("Failed to parse messages size".into()))?;
Ok(())
} else {
Err(ClientError::UnexpectedResponse("EOF".into()))
}
}
pub async fn fetch_all_messages(&mut self) -> Result<Vec<Cow<str>>, ClientError> {
self.fetch_messages_size().await?;
let ws = self.ws_connection.as_mut().unwrap();
ws.send(Message::Binary(vec![0x01].into()))
.await
.map_err(|e| ClientError::WsSendError(e.to_string()))?;
let all_msg = ws
.next()
.await
.ok_or_else(|| ClientError::UnexpectedResponse("EOF".into()))?
.map_err(|e| ClientError::WsReadError(e.to_string()))?;
let payload = match all_msg {
Message::Text(t) => t.to_string(),
Message::Binary(b) => String::from_utf8_lossy(&b).into_owned(),
_ => String::new(),
};
Ok(payload
.lines()
.filter(|l| !l.is_empty())
.map(|s| Cow::Owned(s.to_string()))
.collect())
}
pub async fn fetch_new_messages(&mut self) -> Result<Vec<Cow<str>>, ClientError> {
let old_size = self.current_messages_size.clone();
self.fetch_messages_size().await?;
let new_size = self.current_messages_size;
if old_size >= new_size {
return Ok(Vec::new());
}
let ws = self.ws_connection.as_mut().unwrap();
ws.send(Message::Binary(format!("\x00\x02{}", old_size).into()))
.await
.map_err(|e| ClientError::WsSendError(e.to_string()))?;
let diff_msg = ws
.next()
.await
.ok_or_else(|| ClientError::UnexpectedResponse("EOF".into()))?
.map_err(|e| ClientError::WsReadError(e.to_string()))?;
let payload = match diff_msg {
Message::Text(t) => t.to_string(),
Message::Binary(b) => String::from_utf8_lossy(&b).into_owned(),
_ => String::new(),
};
Ok(payload
.lines()
.filter(|l| !l.is_empty())
.map(|s| Cow::Owned(s.to_string()))
.collect())
}
pub async fn send_message(&mut self, message: &str) -> Result<(), ClientError> {
let msg = message.replace("{username}", &self.username);
self.send_custom_message(&msg).await
}
pub async fn send_custom_message(&mut self, message: &str) -> Result<(), ClientError> {
self.check_connection().await?;
let ws = self.ws_connection.as_mut().unwrap();
if self.password.is_some() {
let payload = format!(
"\x02{}\n{}\n{}",
self.username,
self.password.as_deref().unwrap(),
message
);
ws.send(Message::Binary(payload.into()))
.await
.map_err(|e| ClientError::WsSendError(e.to_string()))?;
if let Some(Ok(Message::Binary(buf))) = ws.next().await {
return match buf.first() {
Some(0x01) => Err(ClientError::UserDoesNotExist),
Some(0x02) => Err(ClientError::IncorrectPassword),
Some(code) => Err(ClientError::UnexpectedResponse(format!("0x{code:02x}"))),
None => Ok(()),
};
}
return Ok(());
}
ws.send(Message::Binary(format!("\x01{}", message).into()))
.await
.map_err(|e| ClientError::WsSendError(e.to_string()))?;
Ok(())
}
pub async fn reset(&mut self) {
self.current_messages_size = 0;
self.address.clear();
self.username.clear();
self.password = None;
self.use_tls = false;
if let Some(ws) = &mut self.ws_connection {
let _ = ws.close(None).await;
self.ws_connection = None
}
}
pub fn current_messages_size(&self) -> usize {
self.current_messages_size
}
pub fn tls(&self) -> bool {
self.use_tls
}
pub fn address(&self) -> &str {
&self.address
}
pub fn username(&self) -> &str {
&self.username
}
}