#[macro_use]
extern crate log;
extern crate uuid;
pub mod protocol;
use std::io::{Read, Write};
use std::net::{Shutdown, TcpStream};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::mpsc::{channel, Receiver, Sender};
use std::sync::{Arc, Mutex};
use std::thread;
use std::thread::JoinHandle;
use uuid::Uuid;
struct TobyReceiver {
tcp_stream: TcpStream,
stop: Arc<AtomicBool>,
to_client_sender: Sender<Vec<u8>>,
id: String,
}
impl TobyReceiver {
fn receive_data(&mut self) {
let mut raw_buff = Vec::new();
let mut curr_size: Option<u64> = None;
let mut done = false;
loop {
trace!("{}: looping toby receiver", self.id);
if self.stop.load(Ordering::Relaxed) {
info!(
"{}: Was told to stop, shutting down inbound message consumer thread",
self.id
);
return;
}
let mut tcpbuf = [0u8; 256];
match self.tcp_stream.read(&mut tcpbuf) {
Ok(bytes) => {
if bytes > 0 {
done = false;
raw_buff.append(&mut tcpbuf[0..bytes].to_vec());
} else {
if done {
info!("{}: read zero bytes from tcp stream indicating client hangup, shutting down everything", self.id);
self.stop.store(true, Ordering::Relaxed);
match self.tcp_stream.shutdown(Shutdown::Both) {
Ok(()) => {}
Err(_) => trace!(
"Got an error while shutting down tcp stream, doing nothing"
),
}
return;
}
done = true;
trace!("{}: Read zero bytes, if this happens again, we will shutdown the thread.", self.id);
continue;
}
}
Err(e) => {
error!("{}: Error waiting for data off of tcp stream, shutting down inbound message consumer thread {}", self.id, e);
return;
}
}
curr_size = compute_curr_size(curr_size, &mut raw_buff);
while curr_size.is_some() && raw_buff.len() >= to_usize(curr_size.unwrap()) {
let parsed_message = raw_buff.drain(0..to_usize(curr_size.unwrap())).collect();
raw_buff.shrink_to_fit();
curr_size = compute_curr_size(None, &mut raw_buff);
match self.to_client_sender.send(parsed_message) {
Ok(_) => {} Err(e) => {
info!("{}: Error sending a complete message to the client, shutting down inbound message consumer thread {}", self.id, e);
return;
}
}
}
}
}
}
pub struct TobyMessenger {
tcp_stream: TcpStream,
stop: Arc<AtomicBool>,
receiver_thread: Option<JoinHandle<()>>,
writer: Mutex<()>, id: String,
}
impl TobyMessenger {
pub fn new(tcp_stream: TcpStream) -> TobyMessenger {
TobyMessenger {
tcp_stream: tcp_stream,
receiver_thread: None,
stop: Arc::new(AtomicBool::new(false)),
writer: Mutex::new(()),
id: Uuid::new_v4().hyphenated().to_string(),
}
}
pub fn id(&self) -> String {
self.id.clone()
}
pub fn send(&self, data: Vec<u8>) -> std::io::Result<()> {
match self.writer.lock() {
Ok(_) => send_actual(&self.tcp_stream, data),
Err(e) => {
error!("Error locking the writer! {}", e);
panic!()
}
}
}
pub fn start(&mut self) -> Result<(Receiver<Vec<u8>>), ()> {
if self.receiver_thread.is_some() {
error!("Calling start on a TobyMessenger that has already started a thread!");
return Err(());
}
let (inbound_sender, inbound_receiver) = channel();
let stop_c = self.stop.clone();
let mut success = true;
let id_c = self.id.clone();
match self.tcp_stream.try_clone() {
Ok(stream) => {
self.receiver_thread = Some(
thread::Builder::new()
.name(format!("toby_rec_{}", self.id).to_string())
.spawn(move || {
let mut rec = TobyReceiver {
tcp_stream: stream,
stop: stop_c,
to_client_sender: inbound_sender,
id: id_c,
};
rec.receive_data();
})
.unwrap(),
);
}
Err(e) => {
error!("{}: Error cloning stream for consumer {}", self.id, e);
success = false;
}
}
if success {
Ok(inbound_receiver)
} else {
Err(())
}
}
pub fn stop_nonblock(&mut self) {
self.stop.store(true, Ordering::Relaxed);
match self.tcp_stream.shutdown(Shutdown::Both) {
Ok(()) => {}
Err(_) => trace!("Got an error while shutting down tcp stream, doing nothing"),
}
}
}
fn compute_curr_size(curr_size: Option<u64>, buf: &mut Vec<u8>) -> Option<u64> {
if curr_size.is_none() {
if buf.len() >= 8 {
let size = Some(bytes_to(&buf[0..8]));
buf.drain(0..8);
return size;
}
None
} else {
curr_size
}
}
fn send_actual(mut stream: &TcpStream, buf: Vec<u8>) -> std::io::Result<()> {
stream.write_all(protocol::encode_tobytcp(buf).as_slice())
}
fn bytes_to(bytes: &[u8]) -> u64 {
let mut ret = 0u64;
let mut i = 0; for byte in bytes {
ret = ret | *byte as u64;
if i < 7 {
ret = ret << 8;
}
i = i + 1;
}
ret
}
fn to_usize(num: u64) -> usize {
num as usize
}
#[cfg(test)]
mod tests {
use std::net::{TcpListener, TcpStream};
use std::thread;
#[test]
fn test_send_data() {
thread::spawn(|| {
let listener = TcpListener::bind("127.0.0.1:8032").unwrap();
for stream in listener.incoming() {
let mut messenger = super::TobyMessenger::new(stream.unwrap());
messenger.start().unwrap();
messenger.send(vec![123, 4, 8]).unwrap();
}
});
let stream = TcpStream::connect("127.0.0.1:8032").unwrap();
let mut messenger = super::TobyMessenger::new(stream);
let receiver = messenger.start().unwrap();
assert_eq!(vec![123, 4, 8], receiver.recv().unwrap());
}
#[test]
fn test_echo_single() {
thread::spawn(|| {
let listener = TcpListener::bind("127.0.0.1:8031").unwrap();
for stream in listener.incoming() {
let mut messenger = super::TobyMessenger::new(stream.unwrap());
let receiver = messenger.start().unwrap();
messenger.send(receiver.recv().unwrap()).unwrap();
}
});
let stream = TcpStream::connect("127.0.0.1:8031").unwrap();
let mut messenger = super::TobyMessenger::new(stream);
let receiver = messenger.start().unwrap();
let data = vec![31, 53, 74, 3, 67, 8, 4];
messenger.send(data.clone()).unwrap();
assert_eq!(data, receiver.recv().unwrap());
}
}