use std::collections::HashMap;
use std::io::{BufRead, BufReader, Read, Write};
use std::net::{TcpListener, TcpStream};
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::mpsc::{Receiver, Sender, channel};
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use crate::proto::{HelloError, StreamKind, parse_hello};
const HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(5);
const READ_CHUNK: usize = 8192;
pub type SessionId = u64;
#[derive(Debug, Clone)]
pub enum ServerEvent {
Opened {
id: SessionId,
name: String,
port: u16,
kind: StreamKind,
},
Attached { id: SessionId, reattached: bool },
Bytes { id: SessionId, data: Vec<u8> },
Disconnected { id: SessionId },
Closed { id: SessionId },
}
#[derive(Debug)]
struct Session {
id: SessionId,
port: u16,
live: Arc<AtomicBool>,
generation: Arc<AtomicU64>,
idle_since: Option<Instant>,
shutdown: Arc<AtomicBool>,
}
struct LiveGuard {
sessions: Arc<Mutex<HashMap<String, Session>>>,
name: String,
tx: Sender<ServerEvent>,
id: SessionId,
generation: u64,
}
impl Drop for LiveGuard {
fn drop(&mut self) {
let mut map = self.sessions.lock().unwrap();
let Some(s) = map.get_mut(&self.name) else {
return;
};
if s.generation.load(Ordering::SeqCst) != self.generation {
return;
}
s.live.store(false, Ordering::SeqCst);
s.idle_since = Some(Instant::now());
drop(map);
let _ = self.tx.send(ServerEvent::Disconnected { id: self.id });
}
}
#[derive(Debug)]
pub struct Server {
control_port: u16,
sessions: Arc<Mutex<HashMap<String, Session>>>,
tx: Sender<ServerEvent>,
rx: Receiver<ServerEvent>,
}
impl Server {
pub fn bind(port: u16) -> std::io::Result<Self> {
let control = TcpListener::bind(("127.0.0.1", port))?;
let control_port = control.local_addr()?.port();
let (tx, rx) = channel();
let sessions: Arc<Mutex<HashMap<String, Session>>> = Arc::new(Mutex::new(HashMap::new()));
let s = Arc::clone(&sessions);
let t = tx.clone();
std::thread::spawn(move || {
for stream in control.incoming().flatten() {
let s = Arc::clone(&s);
let t = t.clone();
std::thread::spawn(move || handle_control(stream, &s, &t));
}
});
Ok(Self {
control_port,
sessions,
tx,
rx,
})
}
#[must_use]
pub fn control_port(&self) -> u16 {
self.control_port
}
#[must_use]
pub fn events(&self) -> &Receiver<ServerEvent> {
&self.rx
}
#[must_use]
pub fn live_count(&self) -> usize {
self.sessions
.lock()
.unwrap()
.values()
.filter(|s| s.live.load(Ordering::SeqCst))
.count()
}
pub fn reap(&mut self, ttl: Duration) {
let mut sessions = self.sessions.lock().unwrap();
let mut reaped: Vec<(SessionId, u16, Arc<AtomicBool>)> = Vec::new();
sessions.retain(|_, s| {
let expired = s
.idle_since
.is_some_and(|t| t.elapsed() > ttl && !s.live.load(Ordering::SeqCst));
if expired {
reaped.push((s.id, s.port, Arc::clone(&s.shutdown)));
}
!expired
});
drop(sessions);
for (id, port, shutdown) in reaped {
wake_and_shutdown(id, port, &shutdown);
let _ = self.tx.send(ServerEvent::Closed { id });
}
}
pub fn close_session(&mut self, id: SessionId) {
let removed = {
let mut sessions = self.sessions.lock().unwrap();
let name = sessions
.iter()
.find(|(_, s)| s.id == id)
.map(|(n, _)| n.clone());
name.and_then(|n| sessions.remove(&n))
};
let Some(session) = removed else { return };
wake_and_shutdown(id, session.port, &session.shutdown);
}
}
fn wake_and_shutdown(id: SessionId, port: u16, shutdown: &Arc<AtomicBool>) {
shutdown.store(true, Ordering::SeqCst);
if port != 0 {
const MAX_ATTEMPTS: u32 = 3;
for attempt in 1..=MAX_ATTEMPTS {
match TcpStream::connect(("127.0.0.1", port)) {
Ok(_) => break,
Err(e) if attempt == MAX_ATTEMPTS => {
eprintln!("Failed to wake session {id} on port {port}: {e}");
}
Err(_) => {
std::thread::sleep(Duration::from_millis(10));
}
}
}
}
}
enum AttachOutcome {
SessionGone,
Rejected,
Attached { generation: u64 },
}
static NEXT_ID: AtomicU64 = AtomicU64::new(1);
static NEXT_ANON: AtomicU64 = AtomicU64::new(1);
fn handle_control(
stream: TcpStream,
sessions: &Arc<Mutex<HashMap<String, Session>>>,
tx: &Sender<ServerEvent>,
) {
let _ = stream.set_read_timeout(Some(HANDSHAKE_TIMEOUT));
let mut reader = BufReader::new(match stream.try_clone() {
Ok(s) => s,
Err(_) => return,
});
let mut writer = stream;
let mut line = String::new();
if reader.read_line(&mut line).is_err() {
return;
}
match parse_hello(&line) {
Ok((kind, name)) => match open_or_reuse(&name, kind, sessions, tx) {
Ok(port) => {
let _ = writeln!(writer, "PORT {port}");
}
Err(_) => {
let _ = writeln!(writer, "ERR no port");
}
},
Err(
e @ (HelloError::BadName
| HelloError::MissingVersion
| HelloError::BadVersion
| HelloError::UnsupportedVersion(_)
| HelloError::MissingStreamKind
| HelloError::UnknownStreamKind(_)),
) => {
let _ = writeln!(writer, "{}", e.wire());
}
Err(HelloError::NotHello) => {
let n = NEXT_ANON.fetch_add(1, Ordering::SeqCst);
let name = format!("anon-{n}");
let id = NEXT_ID.fetch_add(1, Ordering::SeqCst);
sessions.lock().unwrap().insert(
name.clone(),
Session {
id,
port: 0,
live: Arc::new(AtomicBool::new(true)),
generation: Arc::new(AtomicU64::new(1)),
idle_since: None,
shutdown: Arc::new(AtomicBool::new(false)),
},
);
let _ = tx.send(ServerEvent::Opened {
id,
name: name.clone(),
port: 0,
kind: StreamKind::Tokens,
});
let _ = tx.send(ServerEvent::Attached {
id,
reattached: false,
});
let _ = tx.send(ServerEvent::Bytes {
id,
data: line.into_bytes(),
});
let _ = writer.set_read_timeout(None);
let _guard = LiveGuard {
sessions: Arc::clone(sessions),
name,
tx: tx.clone(),
id,
generation: 1,
};
pump(reader, id, tx);
}
}
}
fn open_or_reuse(
name: &str,
kind: StreamKind,
sessions: &Arc<Mutex<HashMap<String, Session>>>,
tx: &Sender<ServerEvent>,
) -> std::io::Result<u16> {
{
let mut map = sessions.lock().unwrap();
if let Some(existing) = map.get_mut(name) {
existing.idle_since = None;
let port = existing.port;
return Ok(port);
}
}
let listener = TcpListener::bind(("127.0.0.1", 0))?;
let port = listener.local_addr()?.port();
let id = NEXT_ID.fetch_add(1, Ordering::SeqCst);
let shutdown = Arc::new(AtomicBool::new(false));
sessions.lock().unwrap().insert(
name.to_string(),
Session {
id,
port,
live: Arc::new(AtomicBool::new(false)),
generation: Arc::new(AtomicU64::new(0)),
idle_since: Some(Instant::now()),
shutdown: Arc::clone(&shutdown),
},
);
let _ = tx.send(ServerEvent::Opened {
id,
name: name.to_string(),
port,
kind,
});
let tx = tx.clone();
let sessions = Arc::clone(sessions);
let name = name.to_string();
std::thread::spawn(move || {
for stream in listener.incoming().flatten() {
if shutdown.load(Ordering::SeqCst) {
break;
}
let attach = {
let mut map = sessions.lock().unwrap();
match map.get_mut(&name) {
None => AttachOutcome::SessionGone,
Some(s) if s.live.load(Ordering::SeqCst) => AttachOutcome::Rejected,
Some(s) => {
s.live.store(true, Ordering::SeqCst);
let generation = s.generation.fetch_add(1, Ordering::SeqCst) + 1;
s.idle_since = None;
AttachOutcome::Attached { generation }
}
}
};
match attach {
AttachOutcome::SessionGone => break,
AttachOutcome::Rejected => {
let mut s = stream;
let _ = writeln!(s, "ERR already attached");
}
AttachOutcome::Attached { generation } => {
let reattached = generation > 1;
let _ = tx.send(ServerEvent::Attached { id, reattached });
let tx = tx.clone();
let sessions = Arc::clone(&sessions);
let name = name.clone();
std::thread::spawn(move || {
let _guard = LiveGuard {
sessions,
name,
tx: tx.clone(),
id,
generation,
};
pump(BufReader::new(stream), id, &tx);
});
}
}
}
});
Ok(port)
}
fn pump(mut reader: BufReader<TcpStream>, id: SessionId, tx: &Sender<ServerEvent>) {
let _ = reader.get_ref().set_read_timeout(None);
let mut buf = vec![0u8; READ_CHUNK];
loop {
match reader.read(&mut buf) {
Ok(0) | Err(_) => break,
Ok(n) => {
if tx
.send(ServerEvent::Bytes {
id,
data: buf[..n].to_vec(),
})
.is_err()
{
break;
}
}
}
}
}
#[cfg(test)]
mod guard_generation_tests {
use super::*;
use std::sync::mpsc::channel;
fn make_session(port: u16) -> Session {
Session {
id: 1,
port,
live: Arc::new(AtomicBool::new(true)),
generation: Arc::new(AtomicU64::new(1)),
idle_since: None,
shutdown: Arc::new(AtomicBool::new(false)),
}
}
#[test]
fn outrun_guard_does_not_clobber_a_newer_attachment() {
let sessions: Arc<Mutex<HashMap<String, Session>>> = Arc::new(Mutex::new(HashMap::new()));
let name = "race".to_string();
sessions
.lock()
.unwrap()
.insert(name.clone(), make_session(4242));
let (tx, rx) = channel();
let outrun_guard = LiveGuard {
sessions: Arc::clone(&sessions),
name: name.clone(),
tx,
id: 1,
generation: 1,
};
{
let map = sessions.lock().unwrap();
let s = map.get(&name).unwrap();
s.generation.fetch_add(1, Ordering::SeqCst);
assert!(s.live.load(Ordering::SeqCst));
}
drop(outrun_guard);
let map = sessions.lock().unwrap();
let s = map.get(&name).unwrap();
assert!(
s.live.load(Ordering::SeqCst),
"an outrun guard cleared `live` out from under the new attachment"
);
assert!(
s.idle_since.is_none(),
"an outrun guard marked a currently-attached session idle"
);
drop(map);
assert!(
rx.try_recv().is_err(),
"an outrun guard sent a spurious Disconnected"
);
}
#[test]
fn current_guard_tears_down_normally() {
let sessions: Arc<Mutex<HashMap<String, Session>>> = Arc::new(Mutex::new(HashMap::new()));
let name = "race".to_string();
sessions
.lock()
.unwrap()
.insert(name.clone(), make_session(4242));
let (tx, rx) = channel();
let guard = LiveGuard {
sessions: Arc::clone(&sessions),
name: name.clone(),
tx,
id: 1,
generation: 1,
};
drop(guard);
let map = sessions.lock().unwrap();
let s = map.get(&name).unwrap();
assert!(!s.live.load(Ordering::SeqCst));
assert!(s.idle_since.is_some());
drop(map);
assert!(matches!(
rx.try_recv(),
Ok(ServerEvent::Disconnected { id: 1 })
));
}
}