#![cfg(feature = "liminal-transport")]
use std::error::Error;
use std::io::{Read, Write};
use std::net::{SocketAddr, TcpListener, TcpStream};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::time::{Duration, Instant};
use aion_core::{RunId, WorkflowId};
use aion_server::worker::{
ConnectedWorkerRegistry, DispatchRequest, HeartbeatSweeper, HeartbeatTracker,
LiminalConnectionNotifier, LivenessProbe, WorkerDelivery, WorkerHandle, WorkerId,
};
use aion_worker::{ActivityRegistry, RedialTiming, WorkerConfig, serve_with_redial};
use liminal_server::config::{ChannelDef, ServerConfig};
use liminal_server::server::connection::ConnectionSupervisor;
use liminal_server::server::listener::ServerListener;
use serde::{Deserialize, Serialize};
type TestError = Box<dyn Error + Send + Sync>;
type TestResult = Result<(), TestError>;
const HEARTBEAT_WINDOW: Duration = Duration::from_secs(4);
const OBSERVE: Duration = Duration::from_secs(20);
const NAMESPACE: &str = "remote";
const TASK_QUEUE: &str = "gates";
const ACTIVITY_TYPE: &str = "run-check";
#[derive(Debug, Clone, Serialize, Deserialize)]
struct CheckInput {
label: String,
#[serde(default)]
hold_ms: u64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
struct CheckOutput {
ran: String,
}
fn test_error(message: impl std::fmt::Display) -> TestError {
message.to_string().into()
}
#[test]
fn idle_worker_survives_past_the_old_lease_window() -> TestResult {
let harness = Harness::start()?;
let worker = WorkerThread::spawn(vec![harness.address.to_string()])?;
harness.wait_for_worker()?;
std::thread::sleep(HEARTBEAT_WINDOW * 2 + Duration::from_secs(1));
let handle = harness.registered_worker().ok_or_else(|| {
test_error(
"the idle worker was deregistered: its connection lease expired while it was alive \
and connected — nothing kept the lease alive on an idle connection",
)
})?;
let response = push_dispatch(&handle, "idle-then-dispatched")?;
assert_eq!(
response, r#"{"ran":"idle-then-dispatched"}"#,
"the surviving registration must still be a live, dispatchable connection"
);
assert_eq!(
worker.executions(),
1,
"the worker genuinely executed the dispatch"
);
worker.stop();
harness.shutdown()
}
#[test]
fn a_wedged_link_is_detected_and_redialed() -> TestResult {
let harness = Harness::start()?;
let relay = WedgeableRelay::start(harness.address)?;
let worker = WorkerThread::spawn(vec![relay.address.to_string(), harness.address.to_string()])?;
harness.wait_for_worker()?;
let through_relay = harness
.registered_worker()
.ok_or_else(|| test_error("the worker never registered through the relay"))?;
std::thread::sleep(HEARTBEAT_WINDOW);
relay.wedge();
let deadline = Instant::now() + OBSERVE;
let redialed = loop {
if let Some(handle) = harness
.registered_workers()
.into_iter()
.find(|handle| handle.id() != through_relay.id())
{
break handle;
}
if Instant::now() >= deadline {
return Err(test_error(
"the worker never noticed its link was dead: no redial, no re-registration. \
A wedged connection is indistinguishable from an idle one unless something \
is expected to arrive on it",
));
}
std::thread::sleep(Duration::from_millis(50));
};
let response = push_dispatch(&redialed, "after-the-link-died")?;
assert_eq!(response, r#"{"ran":"after-the-link-died"}"#);
assert_eq!(worker.executions(), 1);
worker.stop();
relay.shutdown();
harness.shutdown()
}
#[test]
fn a_long_activity_is_not_mistaken_for_a_dead_link() -> TestResult {
let harness = Harness::start()?;
let worker = WorkerThread::spawn(vec![harness.address.to_string()])?;
harness.wait_for_worker()?;
let handle = harness
.registered_worker()
.ok_or_else(|| test_error("the worker never registered"))?;
let worker_id = handle.id();
std::thread::sleep(HEARTBEAT_WINDOW);
let hold = HEARTBEAT_WINDOW * 2;
let response = push_dispatch_holding(
&handle,
"long-but-alive",
u64::try_from(hold.as_millis()).unwrap_or(u64::MAX),
)?;
assert_eq!(response, r#"{"ran":"long-but-alive"}"#);
let still = harness
.registered_workers()
.into_iter()
.find(|candidate| candidate.id() == worker_id)
.ok_or_else(|| {
test_error(
"the worker tore its healthy connection down after a long activity: the \
dead-man switch counted execution time as silence",
)
})?;
let response = push_dispatch(&still, "still-here")?;
assert_eq!(response, r#"{"ran":"still-here"}"#);
assert_eq!(worker.executions(), 2);
worker.stop();
harness.shutdown()
}
#[test]
fn a_busy_worker_still_answers_the_servers_liveness_ping() -> TestResult {
let harness = Harness::start()?;
let worker = WorkerThread::spawn(vec![harness.address.to_string()])?;
harness.wait_for_worker()?;
let handle = harness
.registered_worker()
.ok_or_else(|| test_error("the worker never registered"))?;
let worker_id = handle.id();
std::thread::sleep(HEARTBEAT_WINDOW);
let hold = HEARTBEAT_WINDOW * 3;
let busy = std::thread::spawn(move || {
push_dispatch_holding(
&handle,
"busy-not-silent",
u64::try_from(hold.as_millis()).unwrap_or(u64::MAX),
)
});
std::thread::sleep(HEARTBEAT_WINDOW / 2);
let observed_at = Instant::now();
std::thread::sleep(HEARTBEAT_WINDOW + Duration::from_millis(250));
assert!(
observed_at.elapsed() > HEARTBEAT_WINDOW,
"the check must sit more than one window past the start of the leg, or a value seeded \
at registration would satisfy it and the assertion would prove nothing"
);
assert!(
harness.dispatch_reachable(worker_id)?,
"a worker busy with a long activity must still be answering pings: the server can only \
keep its dispatch eligibility if the serve loop stayed free to answer while the handler \
ran"
);
let response = busy
.join()
.map_err(|_| test_error("the busy dispatch thread panicked"))??;
assert_eq!(response, r#"{"ran":"busy-not-silent"}"#);
assert_eq!(worker.executions(), 1);
worker.stop();
harness.shutdown()
}
struct Harness {
listener: Option<ServerListener>,
registry: ConnectedWorkerRegistry,
tracker: HeartbeatTracker,
address: SocketAddr,
runtime: Option<tokio::runtime::Runtime>,
shutdown: tokio::sync::watch::Sender<bool>,
}
impl Harness {
fn start() -> Result<Self, TestError> {
let config = ServerConfig {
listen_address: "127.0.0.1:0".parse().map_err(test_error)?,
health_listen_address: reserve_loopback_port()?,
channels: Vec::<ChannelDef>::new(),
routing_rules: Vec::new(),
persistence_path: None,
cluster: None,
auth: None,
drain_timeout_ms: 30_000,
services: liminal_server::config::ServicesConfig::default(),
limits: liminal_server::config::LimitsConfig::default(),
websocket: None,
participant: None,
};
let registry = ConnectedWorkerRegistry::default();
let tracker = HeartbeatTracker::new(HEARTBEAT_WINDOW);
let notifier = Arc::new(
LiminalConnectionNotifier::new(registry.clone())
.with_heartbeat_tracker(tracker.clone()),
);
let supervisor = {
use liminal_server::server::connection::LiminalConnectionServices;
let services =
Arc::new(LiminalConnectionServices::from_config(&config).map_err(test_error)?);
ConnectionSupervisor::with_services_and_notifier(services, notifier.clone())
.map_err(test_error)?
};
if !notifier.bind_supervisor(supervisor.clone()) {
return Err(test_error("notifier supervisor was already bound"));
}
let listener = ServerListener::bind(&config, supervisor).map_err(test_error)?;
let address = listener.local_addr();
let runtime = tokio::runtime::Builder::new_multi_thread()
.enable_all()
.build()
.map_err(test_error)?;
let (shutdown, shutdown_rx) = tokio::sync::watch::channel(false);
let probe = LivenessProbe::new(
notifier,
tracker.clone(),
registry.clone(),
HEARTBEAT_WINDOW,
);
let sweeper = HeartbeatSweeper::new(
tracker.clone(),
registry.clone(),
aion_server::worker::PendingActivities::default(),
aion_server::shutdown::DrainState::default(),
HEARTBEAT_WINDOW,
);
runtime.spawn(probe.run(shutdown_rx.clone()));
runtime.spawn(sweeper.run(shutdown_rx));
Ok(Self {
listener: Some(listener),
registry,
tracker,
address,
runtime: Some(runtime),
shutdown,
})
}
fn dispatch_reachable(&self, worker: WorkerId) -> Result<bool, TestError> {
self.tracker
.is_dispatch_reachable(worker, Instant::now())
.map_err(test_error)
}
fn registered_workers(&self) -> Vec<WorkerHandle> {
self.registry.all_workers().unwrap_or_default()
}
fn registered_worker(&self) -> Option<WorkerHandle> {
self.registered_workers().into_iter().next()
}
fn wait_for_worker(&self) -> TestResult {
let deadline = Instant::now() + OBSERVE;
while Instant::now() < deadline {
if self.registered_worker().is_some() {
return Ok(());
}
std::thread::sleep(Duration::from_millis(10));
}
Err(test_error("the worker never registered"))
}
fn shutdown(mut self) -> TestResult {
self.shutdown.send(true).map_err(test_error)?;
if let Some(listener) = self.listener.take() {
listener.shutdown().map_err(test_error)?;
}
if let Some(runtime) = self.runtime.take() {
runtime.shutdown_timeout(Duration::from_secs(5));
}
Ok(())
}
}
fn push_dispatch(handle: &WorkerHandle, label: &str) -> Result<String, TestError> {
push_dispatch_holding(handle, label, 0)
}
fn push_dispatch_holding(
handle: &WorkerHandle,
label: &str,
hold_ms: u64,
) -> Result<String, TestError> {
let WorkerDelivery::Liminal(delivery) = handle.delivery() else {
return Err(test_error("the registered worker is not liminal-delivered"));
};
let request = DispatchRequest {
activity_type: ACTIVITY_TYPE.to_owned(),
workflow_id: WorkflowId::new_v4(),
ordinal: 0,
run_id: Some(RunId::new_v4()),
completion_token: "dead-man-switch-token".to_owned(),
idempotency_key: "dead-man-switch-key".to_owned(),
input: serde_json::to_vec(&CheckInput {
label: label.to_owned(),
hold_ms,
})
.map_err(test_error)?,
attempt: 1,
labels: std::collections::BTreeMap::new(),
heartbeat_window_ms: 0,
};
let deadline = Instant::now() + OBSERVE;
let response = delivery
.dispatch_held(&request, || Instant::now() < deadline)
.map_err(test_error)?
.ok_or_else(|| {
test_error("the dispatch reply wait was abandoned at the test observation deadline")
})?;
response
.outcome
.map_err(|reason| test_error(format!("the worker failed the dispatch: {reason}")))
}
struct WorkerThread {
stop: Arc<AtomicBool>,
executions: Arc<AtomicUsize>,
handle: Option<std::thread::JoinHandle<()>>,
}
impl WorkerThread {
fn spawn(candidates: Vec<String>) -> Result<Self, TestError> {
let stop = Arc::new(AtomicBool::new(false));
let executions = Arc::new(AtomicUsize::new(0));
let config = WorkerConfig::builder()
.endpoint("unused-direct-address")
.namespace(NAMESPACE)
.task_queue(TASK_QUEUE)
.identity("dead-man-switch-worker")
.max_concurrency(1)
.reconnect_initial_backoff(Duration::from_millis(20))
.reconnect_max_backoff(Duration::from_millis(100))
.reconnect_max_attempts(3)
.build()
.map_err(test_error)?;
let counter = Arc::clone(&executions);
let registry = Arc::new(
ActivityRegistry::new()
.register_activity(ACTIVITY_TYPE, move |input: CheckInput, _context| {
let counter = Arc::clone(&counter);
Box::pin(async move {
counter.fetch_add(1, Ordering::SeqCst);
if input.hold_ms > 0 {
tokio::time::sleep(Duration::from_millis(input.hold_ms)).await;
}
Ok(CheckOutput { ran: input.label })
})
})
.map_err(test_error)?,
);
let thread_stop = Arc::clone(&stop);
let handle = std::thread::spawn(move || {
if let Err(error) = serve_with_redial(
candidates,
&config,
®istry,
RedialTiming::new(Duration::from_millis(20), Duration::from_millis(100)),
&thread_stop,
None,
|| {},
) {
eprintln!("dead-man-switch worker stopped: {error}");
}
});
Ok(Self {
stop,
executions,
handle: Some(handle),
})
}
fn executions(&self) -> usize {
self.executions.load(Ordering::SeqCst)
}
fn stop(mut self) {
self.stop.store(true, Ordering::SeqCst);
if let Some(handle) = self.handle.take() {
let _ = handle.join();
}
}
}
struct WedgeableRelay {
address: SocketAddr,
wedged_from_worker: Arc<AtomicBool>,
wedged_to_worker: Arc<AtomicBool>,
stop: Arc<AtomicBool>,
handle: Option<std::thread::JoinHandle<()>>,
}
impl WedgeableRelay {
fn start(upstream: SocketAddr) -> Result<Self, TestError> {
let listener = TcpListener::bind("127.0.0.1:0").map_err(test_error)?;
let address = listener.local_addr().map_err(test_error)?;
listener.set_nonblocking(true).map_err(test_error)?;
let wedged_from_worker = Arc::new(AtomicBool::new(false));
let wedged_to_worker = Arc::new(AtomicBool::new(false));
let stop = Arc::new(AtomicBool::new(false));
let accept_from_worker = Arc::clone(&wedged_from_worker);
let accept_to_worker = Arc::clone(&wedged_to_worker);
let accept_stop = Arc::clone(&stop);
let handle = std::thread::spawn(move || {
let mut parked: Vec<TcpStream> = Vec::new();
while !accept_stop.load(Ordering::SeqCst) {
match listener.accept() {
Ok((downstream, _)) => {
let Ok(up) = TcpStream::connect(upstream) else {
continue;
};
let Ok(down_read) = downstream.try_clone() else {
continue;
};
let Ok(down_write) = downstream.try_clone() else {
continue;
};
let Ok(up_read) = up.try_clone() else {
continue;
};
let Ok(up_write) = up.try_clone() else {
continue;
};
parked.push(downstream);
parked.push(up);
for (from, to, wedged) in [
(down_read, up_write, &accept_from_worker),
(up_read, down_write, &accept_to_worker),
] {
let pump_wedged = Arc::clone(wedged);
let pump_stop = Arc::clone(&accept_stop);
std::thread::spawn(move || pump(from, to, &pump_wedged, &pump_stop));
}
}
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {
std::thread::sleep(Duration::from_millis(10));
}
Err(_) => break,
}
}
drop(parked);
});
Ok(Self {
address,
wedged_from_worker,
wedged_to_worker,
stop,
handle: Some(handle),
})
}
fn wedge(&self) {
self.wedged_from_worker.store(true, Ordering::SeqCst);
self.wedged_to_worker.store(true, Ordering::SeqCst);
}
fn shutdown(mut self) {
self.stop.store(true, Ordering::SeqCst);
if let Some(handle) = self.handle.take() {
let _ = handle.join();
}
}
}
fn pump(mut from: TcpStream, mut to: TcpStream, wedged: &AtomicBool, stop: &AtomicBool) {
if from
.set_read_timeout(Some(Duration::from_millis(50)))
.is_err()
{
return;
}
let mut buffer = [0_u8; 8192];
while !stop.load(Ordering::SeqCst) {
match from.read(&mut buffer) {
Ok(0) => return,
Ok(read) => {
if wedged.load(Ordering::SeqCst) {
continue;
}
let Some(chunk) = buffer.get(..read) else {
return;
};
if to.write_all(chunk).is_err() || to.flush().is_err() {
return;
}
}
Err(error)
if matches!(
error.kind(),
std::io::ErrorKind::WouldBlock | std::io::ErrorKind::TimedOut
) => {}
Err(_) => return,
}
}
}
fn reserve_loopback_port() -> Result<SocketAddr, TestError> {
let listener = TcpListener::bind("127.0.0.1:0").map_err(test_error)?;
let address = listener.local_addr().map_err(test_error)?;
drop(listener);
Ok(address)
}