use std::collections::HashMap;
use std::collections::hash_map::Entry;
use std::net::{SocketAddr, TcpStream};
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{Arc, Condvar, Mutex};
use std::thread::JoinHandle;
use beamr::native::native_process::NativeHandlerFactory;
use tungstenite::protocol::{Role, WebSocket};
use super::super::ConnectionSupervisor;
use super::process::WebSocketConnectionProcess;
use super::{AcceptorSettings, HandshakeOutcome, perform_upgrade, pinned_protocol_config};
#[cfg(test)]
#[path = "supervisor_tests.rs"]
mod tests;
#[derive(Clone, Debug)]
pub(super) struct HandshakeSupervisor {
inner: Arc<HandshakeShared>,
}
#[derive(Debug)]
enum WorkerRecord {
Live(JoinHandle<()>),
CompletedBeforeInstall,
}
struct HandshakeCompletion {
shared: Arc<HandshakeShared>,
handshake_id: u64,
}
impl Drop for HandshakeCompletion {
fn drop(&mut self) {
self.shared.complete_worker(self.handshake_id);
}
}
#[derive(Debug)]
struct HandshakeShared {
supervisor: ConnectionSupervisor,
settings: Arc<AcceptorSettings>,
stopping: AtomicBool,
inflight: Mutex<HashMap<u64, TcpStream>>,
workers: Mutex<HashMap<u64, WorkerRecord>>,
completions: Mutex<u64>,
completion_signal: Condvar,
next_id: AtomicU64,
}
impl HandshakeSupervisor {
pub(super) fn new(supervisor: ConnectionSupervisor, settings: AcceptorSettings) -> Self {
Self {
inner: Arc::new(HandshakeShared {
supervisor,
settings: Arc::new(settings),
stopping: AtomicBool::new(false),
inflight: Mutex::new(HashMap::new()),
workers: Mutex::new(HashMap::new()),
completions: Mutex::new(0),
completion_signal: Condvar::new(),
next_id: AtomicU64::new(1),
}),
}
}
pub(super) fn begin(&self, stream: TcpStream, peer_addr: Option<SocketAddr>) {
if self.inner.stopping.load(Ordering::SeqCst) {
if let Err(error) = stream.shutdown(std::net::Shutdown::Both) {
tracing::debug!(?peer_addr, %error, "post-stop handshake socket shutdown failed");
}
return;
}
let handshake_id = self.inner.next_id.fetch_add(1, Ordering::Relaxed);
let guard = match stream.try_clone() {
Ok(guard) => guard,
Err(error) => {
tracing::warn!(?peer_addr, %error, "handshake socket could not be retained");
if let Err(error) = stream.shutdown(std::net::Shutdown::Both) {
tracing::debug!(?peer_addr, %error, "handshake refusal shutdown failed");
}
return;
}
};
match self.inner.inflight.lock() {
Ok(mut inflight) => {
inflight.insert(handshake_id, guard);
}
Err(poisoned) => {
tracing::error!(
?peer_addr,
error = %poisoned,
"handshake registry poisoned; refusing the connection"
);
if let Err(error) = stream.shutdown(std::net::Shutdown::Both) {
tracing::debug!(?peer_addr, %error, "handshake refusal shutdown failed");
}
return;
}
}
let shared = Arc::clone(&self.inner);
let worker = std::thread::spawn(move || {
let completion = HandshakeCompletion {
shared,
handshake_id,
};
completion
.shared
.run_handshake(handshake_id, stream, peer_addr);
});
self.inner.install_handle(handshake_id, worker);
}
pub(super) fn stop(&self) {
self.inner.stopping.store(true, Ordering::SeqCst);
if let Ok(mut inflight) = self.inner.inflight.lock() {
for (handshake_id, guard) in inflight.drain() {
if let Err(error) = guard.shutdown(std::net::Shutdown::Both) {
tracing::debug!(
handshake_id,
%error,
"in-flight handshake socket shutdown failed (already closed)"
);
}
}
}
self.inner.drain_and_join_live();
self.inner.drain_and_join_live();
}
}
impl HandshakeShared {
fn install_handle(&self, handshake_id: u64, worker: JoinHandle<()>) {
if let Ok(mut workers) = self.workers.lock() {
match workers.entry(handshake_id) {
Entry::Occupied(entry) => {
entry.remove();
drop(worker);
}
Entry::Vacant(entry) => {
entry.insert(WorkerRecord::Live(worker));
}
}
}
}
fn complete_worker(&self, handshake_id: u64) {
if let Ok(mut workers) = self.workers.lock() {
match workers.entry(handshake_id) {
Entry::Occupied(entry) => {
entry.remove();
}
Entry::Vacant(entry) => {
entry.insert(WorkerRecord::CompletedBeforeInstall);
}
}
}
if let Ok(mut completions) = self.completions.lock() {
*completions = completions.saturating_add(1);
self.completion_signal.notify_all();
}
}
fn drain_and_join_live(&self) {
let live: Vec<(u64, JoinHandle<()>)> = self.workers.lock().map_or_else(
|_| Vec::new(),
|mut workers| {
let mut handles = Vec::new();
for (handshake_id, record) in workers.drain() {
if let WorkerRecord::Live(handle) = record {
handles.push((handshake_id, handle));
}
}
handles
},
);
for (handshake_id, handle) in live {
if handle.join().is_err() {
tracing::error!(handshake_id, "websocket handshake worker panicked");
}
}
}
fn run_handshake(
&self,
handshake_id: u64,
mut stream: TcpStream,
peer_addr: Option<SocketAddr>,
) {
if let Err(error) = stream.set_nonblocking(false) {
tracing::warn!(?peer_addr, %error, "handshake socket mode change failed");
self.remove_inflight(handshake_id);
return;
}
let outcome = perform_upgrade(&mut stream, &self.settings);
self.remove_inflight(handshake_id);
match outcome {
HandshakeOutcome::Upgraded => {
if self.stopping.load(Ordering::SeqCst) {
if let Err(error) = stream.shutdown(std::net::Shutdown::Both) {
tracing::debug!(?peer_addr, %error, "post-stop upgrade shutdown failed");
}
return;
}
self.spawn_upgraded(stream, peer_addr);
}
HandshakeOutcome::Refused(refusal) => {
tracing::info!(
?peer_addr,
status = %refusal.status(),
reason = %refusal,
"websocket upgrade refused"
);
}
HandshakeOutcome::SocketError(error) => {
tracing::debug!(?peer_addr, %error, "websocket handshake socket failed");
}
}
}
fn spawn_upgraded(&self, stream: TcpStream, peer_addr: Option<SocketAddr>) {
if let Err(error) = stream.set_nonblocking(true) {
tracing::warn!(?peer_addr, %error, "upgraded socket mode change failed");
return;
}
let fd_guard = match stream.try_clone() {
Ok(guard) => guard,
Err(error) => {
tracing::warn!(?peer_addr, %error, "failed to retain connection fd for teardown");
return;
}
};
let socket = WebSocket::from_raw_socket(
stream,
Role::Server,
Some(pinned_protocol_config(self.settings.message_bound)),
);
let holder = Arc::new(Mutex::new(Some(socket)));
let settings = Arc::clone(&self.settings);
let build = move |runtime: Arc<super::super::supervisor::ConnectionRuntime>,
incarnation: Option<liminal_protocol::wire::ConnectionIncarnation>|
-> NativeHandlerFactory {
let holder = Arc::clone(&holder);
let settings = Arc::clone(&settings);
Box::new(move || {
Box::new(WebSocketConnectionProcess::from_holder(
Arc::clone(&runtime),
peer_addr,
&holder,
incarnation,
&settings,
))
})
};
match self
.supervisor
.spawn_transport_connection(peer_addr, fd_guard, &build)
{
Ok(handle) => {
tracing::debug!(
?peer_addr,
connection_pid = handle.pid(),
"websocket connection admitted"
);
}
Err(error) => {
tracing::warn!(?peer_addr, %error, "websocket connection refused at spawn");
}
}
}
fn remove_inflight(&self, handshake_id: u64) {
if let Ok(mut inflight) = self.inflight.lock() {
inflight.remove(&handshake_id);
}
}
}
#[cfg(test)]
impl HandshakeSupervisor {
fn worker_record_count(&self) -> usize {
self.inner.workers.lock().map_or(0, |workers| workers.len())
}
fn wait_for_completions(&self, target: u64) {
let Ok(guard) = self.inner.completions.lock() else {
return;
};
let _held = self
.inner
.completion_signal
.wait_while(guard, |count| *count < target);
}
}