use std::collections::HashMap;
use std::sync::{Arc, Mutex, MutexGuard};
use std::time::Duration;
use aion_proto::ProtoLivenessPing;
use tokio::sync::oneshot;
use super::liveness::PingFailure;
use super::registry::{WorkerId, WorkerMessage, WorkerTaskSender};
use crate::error::ServerError;
#[derive(Clone, Debug)]
pub struct GrpcLivenessTarget {
pub worker_id: WorkerId,
pub sender: WorkerTaskSender,
}
#[derive(Debug)]
struct ArmedPing {
sequence: u64,
answer: oneshot::Sender<u64>,
}
#[derive(Clone, Debug, Default)]
pub struct GrpcLivenessWaiters {
inner: Arc<Mutex<HashMap<WorkerId, ArmedPing>>>,
}
impl GrpcLivenessWaiters {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn arm(
&self,
worker_id: WorkerId,
sequence: u64,
) -> Result<oneshot::Receiver<u64>, ServerError> {
let (answer, wait) = oneshot::channel();
self.waiters()?
.insert(worker_id, ArmedPing { sequence, answer });
Ok(wait)
}
pub fn answer(&self, worker_id: WorkerId, sequence: u64) -> Result<bool, ServerError> {
let mut waiters = self.waiters()?;
let Some(armed) = waiters.get(&worker_id) else {
return Ok(false);
};
if armed.sequence != sequence {
return Ok(false);
}
let Some(armed) = waiters.remove(&worker_id) else {
return Ok(false);
};
drop(waiters);
Ok(armed.answer.send(sequence).is_ok())
}
pub fn disarm(&self, worker_id: WorkerId) -> Result<(), ServerError> {
self.waiters()?.remove(&worker_id);
Ok(())
}
fn waiters(&self) -> Result<MutexGuard<'_, HashMap<WorkerId, ArmedPing>>, ServerError> {
self.inner
.lock()
.map_err(|_| ServerError::lock_poisoned("grpc worker liveness waiters"))
}
}
pub(super) async fn ping_grpc_worker(
waiters: &GrpcLivenessWaiters,
target: &GrpcLivenessTarget,
ping: ProtoLivenessPing,
cadence: Duration,
) -> Result<u64, PingFailure> {
let sequence = ping.liveness_ping;
let permit = tokio::time::timeout(cadence, target.sender.reserve())
.await
.map_err(|_| {
PingFailure::Unaskable(format!(
"the worker's stream delivery channel stayed full for the whole {cadence:?} probe \
cadence; a dispatch queued now would wait behind the same backlog"
))
})?
.map_err(|error| {
PingFailure::Unaskable(format!(
"the worker's stream delivery channel is closed: {error}"
))
})?;
let wait = waiters.arm(target.worker_id, sequence).map_err(|error| {
PingFailure::Unaskable(format!("could not arm the answer waiter: {error}"))
})?;
permit.send(WorkerMessage::LivenessPing(ping));
let answered = tokio::time::timeout(cadence, wait).await;
if let Err(error) = waiters.disarm(target.worker_id) {
tracing::warn!(
%error,
worker_id = target.worker_id.value(),
liveness_ping = sequence,
"could not disarm a spent gRPC liveness waiter; a later answer may be matched \
against a ping this round already judged"
);
}
match answered {
Ok(Ok(echoed)) => Ok(echoed),
Ok(Err(_)) => Err(PingFailure::Unanswered(String::from(
"the answer channel closed before the worker replied",
))),
Err(_) => Err(PingFailure::Unanswered(format!(
"no answer arrived within the {cadence:?} probe cadence"
))),
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use aion_proto::ProtoLivenessPing;
use super::{GrpcLivenessTarget, GrpcLivenessWaiters, ping_grpc_worker};
use crate::worker::registry::{WorkerId, WorkerMessage};
type TestResult = Result<(), Box<dyn std::error::Error>>;
const CADENCE: Duration = Duration::from_millis(200);
fn ping(sequence: u64) -> ProtoLivenessPing {
ProtoLivenessPing {
liveness_ping: sequence,
silence_window_ms: 800,
}
}
#[tokio::test]
async fn an_echoed_sequence_answers_the_armed_ping() -> TestResult {
let waiters = GrpcLivenessWaiters::new();
let (sender, mut stream) = tokio::sync::mpsc::channel(4);
let target = GrpcLivenessTarget {
worker_id: WorkerId::from_value(1),
sender,
};
let answering = {
let waiters = waiters.clone();
tokio::spawn(async move {
let Some(WorkerMessage::LivenessPing(received)) = stream.recv().await else {
return Ok(false);
};
waiters.answer(WorkerId::from_value(1), received.liveness_ping)
})
};
let echoed = ping_grpc_worker(&waiters, &target, ping(9), CADENCE)
.await
.map_err(|failure| format!("ping failed: {failure:?}"))?;
assert_eq!(echoed, 9, "the probe must observe the sequence it sent");
assert!(
answering.await??,
"the answer must MATCH the armed ping, not merely be delivered"
);
Ok(())
}
#[tokio::test]
async fn a_stale_sequence_neither_matches_nor_consumes_the_armed_waiter() -> TestResult {
let waiters = GrpcLivenessWaiters::new();
let worker = WorkerId::from_value(4);
let wait = waiters.arm(worker, 12)?;
assert!(
!waiters.answer(worker, 11)?,
"an answer echoing an earlier sequence is not an answer to this ping"
);
assert!(
waiters.answer(worker, 12)?,
"the armed ping must still be answerable after a stale echo was rejected"
);
assert_eq!(wait.await?, 12);
Ok(())
}
#[tokio::test]
async fn an_answer_with_nothing_armed_is_reported_unmatched() -> TestResult {
let waiters = GrpcLivenessWaiters::new();
assert!(!waiters.answer(WorkerId::from_value(2), 1)?);
Ok(())
}
#[tokio::test]
async fn silence_after_an_admitted_push_is_unanswered_not_unaskable() -> TestResult {
let waiters = GrpcLivenessWaiters::new();
let (sender, _stream) = tokio::sync::mpsc::channel(4);
let target = GrpcLivenessTarget {
worker_id: WorkerId::from_value(3),
sender,
};
let Err(failure) =
ping_grpc_worker(&waiters, &target, ping(1), Duration::from_millis(60)).await
else {
return Err("a worker that never answers must not report success".into());
};
assert!(
matches!(failure, super::PingFailure::Unanswered(_)),
"an admitted push that goes unanswered is evidence about the WORKER: {failure:?}"
);
Ok(())
}
#[tokio::test]
async fn a_channel_that_never_admits_is_unaskable_and_arms_nothing() -> TestResult {
let waiters = GrpcLivenessWaiters::new();
let (sender, _held) = tokio::sync::mpsc::channel(1);
sender.try_send(WorkerMessage::DrainRequest)?;
let worker = WorkerId::from_value(5);
let target = GrpcLivenessTarget {
worker_id: worker,
sender,
};
let Err(failure) =
ping_grpc_worker(&waiters, &target, ping(1), Duration::from_millis(60)).await
else {
return Err("a channel with no free slot cannot carry a ping".into());
};
assert!(
matches!(failure, super::PingFailure::Unaskable(_)),
"a refused admission is evidence about THIS SERVER's reach: {failure:?}"
);
assert!(
!waiters.answer(worker, 1)?,
"a ping that was never sent must leave no armed waiter behind"
);
Ok(())
}
}