use std::net::{SocketAddr, TcpListener, TcpStream};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::mpsc::{Receiver, Sender, TryRecvError, channel};
use std::sync::Arc;
use std::thread::{self, JoinHandle};
use std::time::{Duration, Instant};
use crate::tensor::{Result, TensorError};
use super::wire::{
CHANNEL_MAGIC_CONTROL, CHANNEL_MAGIC_DATA, CHANNEL_MAGIC_HTTP_GET,
CHANNEL_MAGIC_JOIN, CHANNEL_MAGIC_RENDEZVOUS,
};
const ACCEPT_POLL: Duration = Duration::from_millis(20);
const PEEK_TIMEOUT_SECS: u64 = 5;
pub(crate) struct MuxAccept {
pub rendezvous: Receiver<TcpStream>,
pub data: Receiver<TcpStream>,
pub control: Receiver<TcpStream>,
pub join: Receiver<TcpStream>,
pub status: Receiver<TcpStream>,
}
pub(crate) struct PortMux {
shutdown: Arc<AtomicBool>,
handle: Option<JoinHandle<()>>,
bound_port: u16,
}
impl PortMux {
pub fn start(
listener: TcpListener,
abort: Arc<AtomicBool>,
) -> Result<(Self, MuxAccept)> {
let bound_port = listener
.local_addr()
.map_err(|e| {
TensorError::new(&format!("port_mux: local_addr() failed: {e}"))
})?
.port();
listener.set_nonblocking(true).map_err(|e| {
TensorError::new(&format!("port_mux: set_nonblocking failed: {e}"))
})?;
let (rdv_tx, rdv_rx) = channel();
let (data_tx, data_rx) = channel();
let (ctrl_tx, ctrl_rx) = channel();
let (join_tx, join_rx) = channel();
let (status_tx, status_rx) = channel();
let shutdown = Arc::new(AtomicBool::new(false));
let shutdown_c = Arc::clone(&shutdown);
let handle = thread::Builder::new()
.name(format!("flodl-port-mux:{bound_port}"))
.spawn(move || {
dispatch_loop(
listener, rdv_tx, data_tx, ctrl_tx, join_tx, status_tx,
shutdown_c, abort,
);
})
.map_err(|e| {
TensorError::new(&format!("port_mux: spawn dispatcher failed: {e}"))
})?;
Ok((
PortMux { shutdown, handle: Some(handle), bound_port },
MuxAccept {
rendezvous: rdv_rx,
data: data_rx,
control: ctrl_rx,
join: join_rx,
status: status_rx,
},
))
}
pub fn port(&self) -> u16 {
self.bound_port
}
}
impl Drop for PortMux {
fn drop(&mut self) {
self.shutdown.store(true, Ordering::SeqCst);
if let Some(h) = self.handle.take() {
let _ = h.join();
}
}
}
#[allow(clippy::too_many_arguments)]
fn dispatch_loop(
listener: TcpListener,
rdv_tx: Sender<TcpStream>,
data_tx: Sender<TcpStream>,
ctrl_tx: Sender<TcpStream>,
join_tx: Sender<TcpStream>,
status_tx: Sender<TcpStream>,
shutdown: Arc<AtomicBool>,
abort: Arc<AtomicBool>,
) {
loop {
if shutdown.load(Ordering::SeqCst) || abort.load(Ordering::SeqCst) {
return;
}
let (stream, peer) = match listener.accept() {
Ok(pair) => pair,
Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => {
thread::sleep(ACCEPT_POLL);
continue;
}
Err(e) => {
eprintln!("port_mux: accept failed: {e}");
return;
}
};
dispatch_one(
stream, peer, &rdv_tx, &data_tx, &ctrl_tx, &join_tx, &status_tx,
);
}
}
#[allow(clippy::too_many_arguments)]
fn dispatch_one(
stream: TcpStream,
peer: SocketAddr,
rdv_tx: &Sender<TcpStream>,
data_tx: &Sender<TcpStream>,
ctrl_tx: &Sender<TcpStream>,
join_tx: &Sender<TcpStream>,
status_tx: &Sender<TcpStream>,
) {
crate::distributed::wire::warn_cleartext_public_peer("cluster controller", peer);
if stream.set_nonblocking(false).is_err() {
return;
}
let deadline_secs =
crate::distributed::wire::scaled_deadline_secs(PEEK_TIMEOUT_SECS);
if stream
.set_read_timeout(Some(Duration::from_secs(deadline_secs)))
.is_err()
{
return;
}
let magic = match peek_magic(&stream, Duration::from_secs(deadline_secs)) {
Ok(m) => m,
Err(why) => {
eprintln!(
"port_mux: dropping connection from {peer} ({why}); \
continuing to accept"
);
return;
}
};
if stream.set_read_timeout(None).is_err() {
return;
}
let (tx, name) = match magic {
CHANNEL_MAGIC_RENDEZVOUS => (rdv_tx, "rendezvous"),
CHANNEL_MAGIC_DATA => (data_tx, "data"),
CHANNEL_MAGIC_CONTROL => (ctrl_tx, "control"),
CHANNEL_MAGIC_JOIN => (join_tx, "join"),
CHANNEL_MAGIC_HTTP_GET => (status_tx, "status"),
other => {
eprintln!(
"port_mux: dropping connection from {peer} (unknown channel \
magic 0x{other:08x}); continuing to accept"
);
return;
}
};
if tx.send(stream).is_err() {
eprintln!(
"port_mux: dropping connection from {peer} (no {name}-channel \
subsystem running); continuing to accept"
);
}
}
fn peek_magic(
stream: &TcpStream,
deadline: Duration,
) -> std::result::Result<u32, String> {
let start = Instant::now();
let mut buf = [0u8; 4];
loop {
match stream.peek(&mut buf) {
Ok(0) => return Err("closed before channel magic".to_string()),
Ok(n) if n >= 4 => return Ok(u32::from_le_bytes(buf)),
Ok(_) => {
if start.elapsed() > deadline {
return Err("channel magic incomplete within deadline".to_string());
}
thread::sleep(Duration::from_millis(5));
}
Err(e)
if e.kind() == std::io::ErrorKind::WouldBlock
|| e.kind() == std::io::ErrorKind::TimedOut =>
{
return Err("no channel magic within deadline".to_string());
}
Err(e) => return Err(format!("peek failed: {e}")),
}
}
}
pub(crate) enum StreamSource {
Listener(TcpListener),
Mux(Receiver<TcpStream>),
}
impl StreamSource {
pub fn from_listener(listener: TcpListener, what: &str) -> Result<Self> {
listener.set_nonblocking(true).map_err(|e| {
TensorError::new(&format!("{what}: set_nonblocking failed: {e}"))
})?;
Ok(StreamSource::Listener(listener))
}
pub fn try_accept(&self, what: &str) -> Result<Option<TcpStream>> {
match self {
StreamSource::Listener(listener) => match listener.accept() {
Ok((stream, _peer)) => {
stream.set_nonblocking(false).map_err(|e| {
TensorError::new(&format!(
"{what}: set_nonblocking(false) failed: {e}"
))
})?;
Ok(Some(stream))
}
Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => Ok(None),
Err(e) => Err(TensorError::new(&format!(
"{what}: accept failed: {e}"
))),
},
StreamSource::Mux(rx) => match rx.try_recv() {
Ok(stream) => Ok(Some(stream)),
Err(TryRecvError::Empty) => Ok(None),
Err(TryRecvError::Disconnected) => Err(TensorError::new(&format!(
"{what}: port mux dispatcher exited"
))),
},
}
}
}
#[cfg(test)]
#[path = "port_mux_tests.rs"]
mod tests;