use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, LazyLock};
use std::time::{Duration, Instant};
use parking_lot::Mutex;
use serde::{Deserialize, Serialize};
use crate::remote::link::LinkCommand;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "result", rename_all = "camelCase")]
pub enum AckOutcome {
Done,
Queued,
Refused { reason: String },
Failed { error: String },
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct Envelope {
#[serde(flatten)]
pub command: LinkCommand,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub ack: Option<u64>,
}
impl Envelope {
pub fn parse(json: &str) -> Result<Self, String> {
serde_json::from_str(json).map_err(|e| e.to_string())
}
}
impl From<LinkCommand> for Envelope {
fn from(command: LinkCommand) -> Self {
Self { command, ack: None }
}
}
const REMEMBERED: Duration = Duration::from_secs(60);
pub fn next_id() -> u64 {
let mut bytes = [0u8; 8];
if getrandom::fill(&mut bytes).is_err() {
static FALLBACK: AtomicU64 = AtomicU64::new(1);
return FALLBACK.fetch_add(1, Ordering::Relaxed) | (1 << 63);
}
u64::from_le_bytes(bytes).max(1)
}
type Waiter = (String, crossbeam_channel::Sender<AckOutcome>);
static WAITING: LazyLock<Mutex<HashMap<u64, Waiter>>> = LazyLock::new(Default::default);
pub fn expect(id: u64, device: &str) -> crossbeam_channel::Receiver<AckOutcome> {
let (tx, rx) = crossbeam_channel::bounded(1);
WAITING.lock().insert(id, (device.to_owned(), tx));
rx
}
pub fn resolve(id: u64, from: &str, outcome: AckOutcome) {
let mut waiting = WAITING.lock();
match waiting.get(&id) {
Some((device, _)) if device == from => {}
Some((device, _)) => {
log::warn!("acks: {from} answered a command sent to {device}; ignored");
return;
}
None => return,
}
if let Some((_, tx)) = waiting.remove(&id) {
let _ = tx.try_send(outcome);
}
}
pub fn forget(id: u64) {
WAITING.lock().remove(&id);
}
type Reply = Box<dyn FnOnce(AckOutcome) + Send>;
enum Taken {
Running(Vec<Reply>),
Answered(AckOutcome),
}
static TAKEN: LazyLock<Mutex<HashMap<u64, (Instant, Taken)>>> = LazyLock::new(Default::default);
pub struct Pending {
id: u64,
finished: bool,
}
impl Pending {
pub fn finish(mut self, outcome: AckOutcome) {
self.finished = true;
answer(self.id, outcome);
}
}
impl Drop for Pending {
fn drop(&mut self) {
if !self.finished {
answer(
self.id,
AckOutcome::Failed {
error: "it stopped before finishing".into(),
},
);
}
}
}
impl std::fmt::Debug for Pending {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "Pending({})", self.id)
}
}
fn answer(id: u64, outcome: AckOutcome) {
let replies = {
let mut taken = TAKEN.lock();
match taken.insert(id, (Instant::now(), Taken::Answered(outcome.clone()))) {
Some((_, Taken::Running(replies))) => replies,
_ => Vec::new(),
}
};
for reply in replies {
reply(outcome.clone());
}
}
pub fn accept(id: u64, reply: impl FnOnce(AckOutcome) + Send + 'static) -> Option<Pending> {
let mut taken = TAKEN.lock();
taken.retain(|_, (at, _)| at.elapsed() < REMEMBERED);
match taken.get_mut(&id) {
Some((_, Taken::Running(replies))) => {
replies.push(Box::new(reply));
None
}
Some((_, Taken::Answered(outcome))) => {
let outcome = outcome.clone();
drop(taken);
reply(outcome);
None
}
None => {
taken.insert(id, (Instant::now(), Taken::Running(vec![Box::new(reply)])));
Some(Pending {
id,
finished: false,
})
}
}
}
pub fn take(
envelope: Envelope,
reply: impl FnOnce(u64, AckOutcome) + Send + 'static,
) -> Option<(LinkCommand, Option<Pending>)> {
match envelope.ack {
None => Some((envelope.command, None)),
Some(id) => {
let reply = Arc::new(Mutex::new(Some(reply)));
let pending = accept(id, move |outcome| {
if let Some(reply) = reply.lock().take() {
reply(id, outcome);
}
})?;
Some((envelope.command, Some(pending)))
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn an_envelope_is_the_command_with_an_id_beside_it() {
let pause = Envelope {
command: LinkCommand::Pause,
ack: Some(17),
};
let json = serde_json::to_string(&pause).unwrap();
assert_eq!(json, r#"{"type":"pause","ack":17}"#);
assert_eq!(serde_json::from_str::<Envelope>(&json).unwrap(), pause);
assert_eq!(
serde_json::from_str::<LinkCommand>(&json).unwrap(),
LinkCommand::Pause
);
assert_eq!(
serde_json::from_str::<Envelope>(r#"{"type":"pause"}"#).unwrap(),
Envelope::from(LinkCommand::Pause)
);
}
#[test]
fn ids_are_random_not_counted() {
let ids: Vec<u64> = (0..8).map(|_| next_id()).collect();
assert!(ids.iter().all(|id| *id != 0));
let consecutive = ids.windows(2).filter(|w| w[1] == w[0] + 1).count();
assert_eq!(consecutive, 0, "one id says nothing of the next: {ids:?}");
}
#[test]
fn a_repeat_is_answered_with_the_first_outcome_and_not_acted_on() {
let id = next_id();
let answers = Arc::new(Mutex::new(Vec::new()));
let record = |n: u32| {
let answers = answers.clone();
move |outcome: AckOutcome| answers.lock().push((n, outcome))
};
let first = accept(id, record(1)).expect("acted on");
assert!(accept(id, record(2)).is_none(), "not acted on twice");
assert!(answers.lock().is_empty());
first.finish(AckOutcome::Done);
assert_eq!(
*answers.lock(),
[(1, AckOutcome::Done), (2, AckOutcome::Done)]
);
assert!(accept(id, record(3)).is_none());
assert_eq!(answers.lock().last(), Some(&(3, AckOutcome::Done)));
}
#[test]
fn a_command_dropped_unfinished_is_answered_as_failed() {
let id = next_id();
let answer = Arc::new(Mutex::new(None));
let got = answer.clone();
drop(accept(id, move |o| *got.lock() = Some(o)));
assert!(matches!(*answer.lock(), Some(AckOutcome::Failed { .. })));
}
#[test]
fn a_sender_hears_its_answer_once() {
let id = next_id();
let rx = expect(id, "phone");
resolve(id, "phone", AckOutcome::Queued);
resolve(id, "phone", AckOutcome::Done);
assert_eq!(rx.try_recv(), Ok(AckOutcome::Queued));
assert!(rx.try_recv().is_err());
}
#[test]
fn only_the_device_a_command_went_to_answers_it() {
let id = next_id();
let rx = expect(id, "phone");
resolve(id, "stranger", AckOutcome::Done);
assert!(rx.try_recv().is_err(), "a stranger's answer is not heard");
resolve(id, "phone", AckOutcome::Failed { error: "no".into() });
assert!(matches!(rx.try_recv(), Ok(AckOutcome::Failed { .. })));
}
}