use crate::shared::{ClientError, Credentials};
use native_tls::TlsConnector;
use std::borrow::Cow;
use std::io::{Read, Write};
use std::net::TcpStream;
trait Io: Read + Write {}
impl<T: Read + Write + ?Sized> Io for T {}
type DynStream = Box<dyn Io + Send>;
#[derive(Debug, Clone)]
pub struct Client {
current_messages_size: usize,
address: String,
username: String,
password: Option<String>,
use_tls: bool,
}
impl Client {
pub fn new(address: &str, credentials: Credentials, use_tls: bool) -> Self {
Self {
current_messages_size: 0,
address: address.to_string(),
username: credentials.username,
password: credentials.password,
use_tls,
}
}
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 get_stream(&self) -> Result<DynStream, ClientError> {
let stream = TcpStream::connect(&self.address).map_err(ClientError::ConnectionError)?;
if !self.use_tls {
return Ok(Box::new(stream));
}
let domain = self.address.split(':').next().unwrap_or("localhost");
let connector =
TlsConnector::new().map_err(|e| ClientError::TlsInitializationError(e.to_string()))?;
let tls_stream = connector
.connect(domain, stream)
.map_err(|e| ClientError::TlsInitializationError(e.to_string()))?;
Ok(Box::new(tls_stream))
}
fn remove_nulls(data: &mut Vec<u8>) {
data.retain(|&x| x != 0);
}
pub fn test_connection(&self) -> Result<(), ClientError> {
self.get_stream()?;
Ok(())
}
pub fn register_user(&mut self) -> Result<(), ClientError> {
let mut stream = self.get_stream()?;
if self.password.is_some() {
stream
.write_all(
format!(
"\x03{}\n{}",
self.username,
self.password.as_deref().unwrap()
)
.as_bytes(),
)
.map_err(ClientError::StreamWriteError)?;
let mut buf = [0u8; 2];
let n = stream
.read(&mut buf)
.map_err(ClientError::StreamReadError)?;
if n == 0 {
return Ok(());
}
match buf[0] {
0x01 => Err(ClientError::UsernameAlreadyTaken),
_ => Err(ClientError::UnexpectedResponse(
String::from_utf8_lossy(&buf[..n]).to_string(),
)),
}
} else {
Err(ClientError::NoPassword)
}
}
pub fn fetch_messages_size(&mut self) -> Result<(), ClientError> {
let mut stream = self.get_stream()?;
stream
.write_all(&[0x00])
.map_err(ClientError::StreamWriteError)?;
let mut buf = vec![0u8; 1024];
let n = stream
.read(&mut buf)
.map_err(ClientError::StreamReadError)?;
if n == 0 {
return Err(ClientError::ServerClosedConnection);
}
Self::remove_nulls(&mut buf);
let response = String::from_utf8_lossy(&buf[..n]);
if let Ok(size) = response.parse::<usize>() {
self.current_messages_size = size;
Ok(())
} else {
Err(ClientError::ParseError(
"Failed to parse messages size".to_string(),
))
}
}
pub fn fetch_all_messages(&mut self) -> Result<Vec<Cow<str>>, ClientError> {
let mut stream = self.get_stream()?;
stream
.write_all(&[0x00])
.map_err(ClientError::StreamWriteError)?;
let mut head = vec![0u8; 1024];
let n = stream
.read(&mut head)
.map_err(ClientError::StreamReadError)?;
if n == 0 {
return Err(ClientError::ServerClosedConnection);
}
Self::remove_nulls(&mut head);
let response = String::from_utf8_lossy(&head[..n]);
let size = response
.parse::<usize>()
.map_err(|_| ClientError::ParseError("Failed to parse messages size".to_string()))?;
self.current_messages_size = size;
stream
.write_all(&[0x01])
.map_err(ClientError::StreamWriteError)?;
let mut buffer = vec![0u8; self.current_messages_size];
stream
.read_exact(&mut buffer)
.map_err(ClientError::StreamReadError)?;
Self::remove_nulls(&mut buffer);
let response = String::from_utf8_lossy(&buffer).into_owned();
let vec_messages = response
.lines()
.filter(|l| !l.is_empty())
.map(|s| Cow::Owned(s.to_string()))
.collect();
Ok(vec_messages)
}
pub fn fetch_new_messages(&mut self) -> Result<Vec<Cow<str>>, ClientError> {
let mut stream = self.get_stream()?;
stream
.write_all(&[0x00])
.map_err(ClientError::StreamWriteError)?;
let mut head = vec![0u8; 1024];
let n = stream
.read(&mut head)
.map_err(ClientError::StreamReadError)?;
if n == 0 {
return Err(ClientError::ServerClosedConnection);
}
Self::remove_nulls(&mut head);
let response = String::from_utf8_lossy(&head[..n]);
let size = response
.parse::<usize>()
.map_err(|_| ClientError::ParseError("Failed to parse messages size".to_string()))?;
stream
.write_all(format!("\x02{}", self.current_messages_size).as_bytes())
.map_err(ClientError::StreamWriteError)?;
let mut buffer = vec![0u8; size - self.current_messages_size];
stream
.read_exact(&mut buffer)
.map_err(ClientError::StreamReadError)?;
Self::remove_nulls(&mut buffer);
let response = String::from_utf8_lossy(&buffer).into_owned();
let vec_messages = response
.lines()
.filter(|l| !l.is_empty())
.map(|s| Cow::Owned(s.to_string()))
.collect();
self.current_messages_size = size;
Ok(vec_messages)
}
pub fn send_message(&self, message: &str) -> Result<(), ClientError> {
let message = message.replace("{username}", &self.username);
self.send_custom_message(&message)
}
pub fn send_custom_message(&self, message: &str) -> Result<(), ClientError> {
let mut stream = self.get_stream()?;
if self.password.is_some() {
stream
.write_all(
format!(
"\x02{}\n{}\n{}",
self.username,
self.password.as_deref().unwrap(),
message
)
.as_bytes(),
)
.map_err(ClientError::StreamWriteError)?;
let mut buf = [0u8; 2];
let n = stream
.read(&mut buf)
.map_err(ClientError::StreamReadError)?;
if n == 0 {
return Ok(());
}
return match buf[0] {
0x01 => Err(ClientError::UserDoesNotExist),
0x02 => Err(ClientError::IncorrectPassword),
_ => Err(ClientError::UnexpectedResponse(
String::from_utf8_lossy(&buf[..n]).to_string(),
)),
};
}
stream
.write_all(format!("\x01{}", message).as_bytes())
.map_err(ClientError::StreamWriteError)?;
Ok(())
}
pub fn reset(&mut self) {
self.current_messages_size = 0;
self.address.clear();
self.username.clear();
self.password = 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
}
}