use std::{
io::{
Error,
ErrorKind,
Read,
Write,
},
net::{
Shutdown,
SocketAddr,
TcpListener,
TcpStream,
ToSocketAddrs,
},
ops::{
Deref,
DerefMut,
},
sync::{
Arc,
mpsc::{
channel,
Receiver,
Sender,
},
RwLock,
},
thread,
time::Duration,
};
pub struct Client {
connection_state: RwLock<ConnectionState>,
}
impl Client {
pub fn new() -> Arc<Self> {
Arc::new(Self {
connection_state: Default::default(),
})
}
pub fn connect(
self: &Arc<Self>,
entity: &str,
connection_mode: ConnectionMode,
t5: Duration,
t8: Duration,
) -> Result<(SocketAddr, Receiver<Message>), Error> {
let (stream, socket) = match self.connection_state.read().unwrap().deref() {
ConnectionState::NotConnected => {
match connection_mode {
ConnectionMode::Passive => {
let listener = TcpListener::bind(entity)?;
listener.accept()?
},
ConnectionMode::Active => {
let socket = entity.to_socket_addrs()?.next().ok_or(Error::from(ErrorKind::AddrNotAvailable))?;
let stream = TcpStream::connect_timeout(
&socket,
t5,
)?;
(stream, socket)
},
}
},
_ => return Err(Error::from(ErrorKind::AlreadyExists)),
};
stream.set_read_timeout(Some(t8))?;
stream.set_write_timeout(Some(t8))?;
*self.connection_state.write().unwrap().deref_mut() = ConnectionState::Connected(stream);
let (rx_sender, rx_receiver) = channel::<Message>();
let rx_clone: Arc<Client> = self.clone();
thread::spawn(move || {rx_clone.receive(rx_sender.clone())});
Ok((socket, rx_receiver))
}
pub fn disconnect(
self: &Arc<Self>
) -> Result<(), Error> {
match self.connection_state.read().unwrap().deref() {
ConnectionState::NotConnected => return Err(Error::from(ErrorKind::NotConnected)),
ConnectionState::Connected(stream) => {
let _ = stream.shutdown(Shutdown::Both);
},
}
*self.connection_state.write().unwrap().deref_mut() = ConnectionState::NotConnected;
Ok(())
}
}
impl Client {
fn receive(
self: Arc<Self>,
rx_sender: Sender<Message>,
) {
while let ConnectionState::Connected(stream_immutable) = self.connection_state.read().unwrap().deref() {
let res: Result<Option<Message>, Error> = 'rx: {
let mut stream: &TcpStream = stream_immutable;
let mut length_buffer: [u8;4] = [0;4];
let length_bytes: usize = match stream.read(&mut length_buffer) {
Ok(l) => l,
Err(error) => match error.kind() {
ErrorKind::TimedOut => {
break 'rx Ok(None)
},
_ => {
break 'rx Err(error)
},
}
};
if length_bytes != 4 {
break 'rx Err(Error::from(ErrorKind::TimedOut))
}
let length: u32 = u32::from_be_bytes(length_buffer);
if length < 10 {
break 'rx Err(Error::from(ErrorKind::InvalidData))
}
let mut message_buffer: Vec<u8> = vec![0; length as usize];
let message_bytes: usize = match stream.read(&mut message_buffer) {
Ok(message_bytes) => message_bytes,
Err(error) => break 'rx Err(error),
};
if message_bytes != length as usize {
break 'rx Err(Error::from(ErrorKind::TimedOut))
}
match Message::try_from(message_buffer) {
Ok(message) => Ok(Some(message)),
Err(_) => break 'rx Err(Error::from(ErrorKind::InvalidData)),
}
};
match res {
Ok(optional_rx_message) => if let Some(rx_message) = optional_rx_message {
if rx_sender.send(rx_message).is_err() {break}
},
Err(_error) => break,
}
}
}
pub fn transmit(
self: &Arc<Self>,
message: Message,
) -> Result<(), Error> {
match self.connection_state.read().unwrap().deref() {
ConnectionState::Connected(stream_immutable) => 'disconnect: {
let mut stream: &TcpStream = stream_immutable;
let message_buffer: Vec<u8> = (&message).into();
let length: u32 = message_buffer.len() as u32;
let length_buffer: [u8; 4] = length.to_be_bytes();
if stream.write_all(&length_buffer).is_err() {break 'disconnect};
if stream.write_all(&message_buffer).is_err() {break 'disconnect};
return Ok(())
},
ConnectionState::NotConnected => return Err(Error::from(ErrorKind::NotConnected)),
};
self.disconnect()?;
Err(Error::from(ErrorKind::ConnectionAborted))
}
}
#[derive(Debug)]
pub enum ConnectionState {
NotConnected,
Connected(TcpStream)
}
impl Default for ConnectionState {
fn default() -> Self {
ConnectionState::NotConnected
}
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub enum ConnectionMode {
Passive,
Active,
}
impl Default for ConnectionMode {
fn default() -> Self {
ConnectionMode::Passive
}
}
#[derive(Clone, Debug)]
pub struct Message {
pub header: MessageHeader,
pub text: Vec<u8>,
}
impl From<&Message> for Vec<u8> {
fn from(val: &Message) -> Self {
let mut vec: Vec<u8> = vec![];
let header_bytes: [u8;10] = val.header.into();
vec.extend(header_bytes.iter());
vec.extend(&val.text);
vec
}
}
impl TryFrom<Vec<u8>> for Message {
type Error = ();
fn try_from(bytes: Vec<u8>) -> Result<Self, Self::Error> {
if bytes.len() < 10 {return Err(())}
Ok(Self {
header: MessageHeader::from(<[u8;10]>::try_from(&bytes[0..10]).map_err(|_| ())?),
text: bytes[10..].to_vec(),
})
}
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct MessageHeader {
pub session_id : u16,
pub byte_2 : u8,
pub byte_3 : u8,
pub presentation_type : u8,
pub session_type : u8,
pub system : u32,
}
impl From<MessageHeader> for [u8;10] {
fn from(val: MessageHeader) -> Self {
let mut bytes: [u8;10] = [0;10];
let session_id_bytes: [u8;2] = val.session_id.to_be_bytes();
let system_bytes: [u8;4] = val.system.to_be_bytes();
bytes[0] = session_id_bytes[0];
bytes[1] = session_id_bytes[1];
bytes[2] = val.byte_2;
bytes[3] = val.byte_3;
bytes[4] = val.presentation_type;
bytes[5] = val.session_type;
bytes[6] = system_bytes[0];
bytes[7] = system_bytes[1];
bytes[8] = system_bytes[2];
bytes[9] = system_bytes[3];
bytes
}
}
impl From<[u8;10]> for MessageHeader {
fn from(bytes: [u8;10]) -> Self {
Self {
session_id : u16::from_be_bytes(bytes[0..2].try_into().unwrap()),
byte_2 : bytes[2],
byte_3 : bytes[3],
presentation_type : bytes[4],
session_type : bytes[5],
system : u32::from_be_bytes(bytes[6..10].try_into().unwrap()),
}
}
}