use std::collections::HashMap;
use std::io;
use std::net::{Shutdown, TcpListener, TcpStream};
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use std::thread;
use std::time::{Duration, Instant};
use crate::RpcServer;
struct IdleState {
conn_count: usize,
deadline: Option<Instant>,
}
fn lock<T>(m: &Mutex<T>) -> std::sync::MutexGuard<'_, T> {
m.lock().unwrap_or_else(|e| e.into_inner())
}
fn reap_finished(threads: &mut Vec<thread::JoinHandle<()>>) {
let mut index = 0;
while index < threads.len() {
if threads[index].is_finished() {
let handle = threads.swap_remove(index);
let _ = handle.join();
} else {
index += 1;
}
}
}
fn join_until(threads: &mut Vec<thread::JoinHandle<()>>, deadline: Instant) {
loop {
reap_finished(threads);
if threads.is_empty() || Instant::now() >= deadline {
return;
}
thread::sleep(Duration::from_millis(5));
}
}
pub fn serve_tcp<F: FnOnce(&str, u16)>(
server: Arc<RpcServer>,
host: &str,
port: u16,
idle_timeout: Option<Duration>,
shutdown: Arc<AtomicBool>,
on_bound: F,
) -> io::Result<()> {
let listener = TcpListener::bind((host, port))?;
let bound_port = listener.local_addr()?.port();
listener.set_nonblocking(true).ok();
on_bound(host, bound_port);
let startup_deadline = idle_timeout.map(|t| Instant::now() + t.max(Duration::from_secs(60)));
let state = Arc::new(Mutex::new(IdleState {
conn_count: 0,
deadline: startup_deadline,
}));
let mut threads: Vec<thread::JoinHandle<()>> = Vec::new();
let active = Arc::new(Mutex::new(HashMap::<u64, TcpStream>::new()));
let next_connection_id = AtomicU64::new(1);
loop {
reap_finished(&mut threads);
if shutdown.load(Ordering::Relaxed) {
break;
}
if idle_timeout.is_some() {
let st = lock(&state);
if st.conn_count == 0 {
if let Some(dl) = st.deadline {
if Instant::now() >= dl {
break;
}
}
}
}
match listener.accept() {
Ok((mut conn, _)) => {
conn.set_nonblocking(false).ok();
conn.set_nodelay(true).ok();
let mut reader = match conn.try_clone() {
Ok(reader) => reader,
Err(_) => continue,
};
let interrupter = match conn.try_clone() {
Ok(interrupter) => interrupter,
Err(_) => continue,
};
{
let mut st = lock(&state);
st.conn_count += 1;
st.deadline = None; }
let srv = server.clone();
let state2 = state.clone();
let active2 = active.clone();
let connection_id = next_connection_id.fetch_add(1, Ordering::Relaxed);
lock(&active).insert(connection_id, interrupter);
threads.push(thread::spawn(move || {
srv.serve(&mut reader, &mut conn);
let mut st = lock(&state2);
st.conn_count -= 1;
if st.conn_count == 0 {
if let Some(t) = idle_timeout {
st.deadline = Some(Instant::now() + t);
}
}
drop(st);
lock(&active2).remove(&connection_id);
}));
}
Err(ref e) if e.kind() == io::ErrorKind::WouldBlock => {
thread::sleep(Duration::from_millis(50));
}
Err(_) => break,
}
}
drop(listener);
for connection in lock(&active).values() {
let _ = connection.shutdown(Shutdown::Both);
}
let deadline = Instant::now() + Duration::from_secs(2);
join_until(&mut threads, deadline);
Ok(())
}