use crate::network::{self, Connection};
use serde::{Serialize, Deserialize};
use std::sync::{Arc, atomic::{AtomicBool, Ordering}};
use std::net::{SocketAddr};
use std::thread::{self, JoinHandle};
use std::fmt::{self};
use std::time::{Duration};
pub type Endpoint = usize;
const NETWORK_SAMPLING_TIMEOUT: u64 = 50;
#[derive(Debug, Clone, Copy)]
pub enum TransportProtocol {
Tcp,
Udp,
}
pub enum NetEvent<InMessage>
where InMessage: for<'b> Deserialize<'b> + Send + 'static {
Message(InMessage, Endpoint),
AddedEndpoint(Endpoint),
RemovedEndpoint(Endpoint),
}
impl fmt::Display for TransportProtocol {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}", format!("{:?}", self).to_lowercase())
}
}
pub struct NetworkManager {
network_event_thread: Option<JoinHandle<()>>,
network_thread_running: Arc<AtomicBool>,
network_controller: network::Controller,
output_buffer: Vec<u8>,
}
impl<'a> NetworkManager {
pub fn new<InMessage, C>(event_callback: C) -> NetworkManager
where InMessage: for<'b> Deserialize<'b> + Send + 'static,
C: Fn(NetEvent<InMessage>) + Send + 'static {
let (network_controller, mut network_receiver) = network::adapter();
let network_thread_running = Arc::new(AtomicBool::new(true));
let running = network_thread_running.clone();
let network_event_thread = thread::spawn(move || {
let timeout = Duration::from_millis(NETWORK_SAMPLING_TIMEOUT);
while running.load(Ordering::Relaxed) {
network_receiver.receive(Some(timeout), |endpoint, event| {
let net_event = match event {
network::Event::Connection => {
NetEvent::AddedEndpoint(endpoint)
},
network::Event::Data(data) => {
let message: InMessage = bincode::deserialize(&data[..]).unwrap();
NetEvent::Message(message, endpoint)
},
network::Event::Disconnection => {
NetEvent::RemovedEndpoint(endpoint)
},
};
event_callback(net_event);
});
}
});
NetworkManager {
network_event_thread: Some(network_event_thread),
network_thread_running,
network_controller,
output_buffer: Vec::new()
}
}
pub fn connect(&mut self, addr: SocketAddr, transport: TransportProtocol) -> Option<Endpoint> {
match transport {
TransportProtocol::Tcp => Connection::new_tcp_stream(addr),
TransportProtocol::Udp => Connection::new_udp_socket(addr),
}
.ok()
.map(|connection| self.network_controller.add_connection(connection))
}
pub fn listen(&mut self, addr: SocketAddr, transport: TransportProtocol) -> Option<Endpoint> {
match transport {
TransportProtocol::Tcp => Connection::new_tcp_listener(addr),
TransportProtocol::Udp => Connection::new_udp_listener(addr),
}
.ok()
.map(|listener| self.network_controller.add_connection(listener))
}
pub fn endpoint_address(&mut self, endpoint: Endpoint) -> Option<SocketAddr> {
self.network_controller.connection_address(endpoint)
}
pub fn remove_endpoint(&mut self, endpoint: Endpoint) -> Option<()> {
self.network_controller.remove_connection(endpoint)
}
pub fn send<OutMessage>(&mut self, endpoint: Endpoint, message: OutMessage) -> Option<()>
where OutMessage: Serialize {
bincode::serialize_into(&mut self.output_buffer, &message).unwrap();
let result = self.network_controller.send(endpoint, &self.output_buffer);
self.output_buffer.clear();
result
}
pub fn send_all<'b, OutMessage>(&mut self, endpoints: impl IntoIterator<Item=&'b Endpoint>, message: OutMessage) -> Result<(), Vec<Endpoint>>
where OutMessage: Serialize {
let mut unrecognized_ids = Vec::new();
bincode::serialize_into(&mut self.output_buffer, &message).unwrap();
for endpoint in endpoints {
if let None = self.network_controller.send(*endpoint, &self.output_buffer) {
unrecognized_ids.push(*endpoint);
}
}
self.output_buffer.clear();
if unrecognized_ids.is_empty() { Ok(()) } else { Err(unrecognized_ids) }
}
}
impl Drop for NetworkManager {
fn drop(&mut self) {
self.network_thread_running.store(false, Ordering::Relaxed);
self.network_event_thread.take().unwrap().join().unwrap();
}
}