use crate::actor::{Actor, ActorId, ActorInit, Behavior, Context, Envelope, Shutdown};
use crate::log::Logger;
use anyhow::{Context as _, bail};
use std::collections::{HashMap, HashSet};
use std::time::Duration;
use tokio::sync::mpsc::{UnboundedReceiver, unbounded_channel};
use tokio::sync::oneshot;
use tokio::task::JoinSet;
use tokio::time::timeout;
use tokio_util::sync::CancellationToken;
use uuid::Uuid;
pub struct Episode<B: Behavior> {
actors: HashMap<ActorId, Actor<B>>,
readies: HashMap<ActorId, oneshot::Receiver<()>>,
starts: HashMap<ActorId, oneshot::Sender<()>>,
}
impl<B: Behavior> Episode<B> {
pub fn new(init: HashMap<ActorId, ActorInit<B>>, logger: Logger<B::Log>) -> Self {
let episode = Uuid::new_v4();
let mut senders = HashMap::new();
let mut shutdowns = HashMap::new();
let staged: Staged<B> = init
.into_iter()
.map(|(id, init)| {
let (tx, rx) = unbounded_channel();
senders.insert(id.clone(), tx);
shutdowns.insert(id.clone(), CancellationToken::new());
(id, (init, rx))
})
.collect();
let mut readies = HashMap::new();
let mut starts = HashMap::new();
let actors = staged
.into_iter()
.map(|(id, (init, mailbox))| {
let (ready, is_ready) = oneshot::channel();
readies.insert(id.clone(), is_ready);
let (start, started) = oneshot::channel();
starts.insert(id.clone(), start);
let mut mailboxes = pick(&senders, &init.can_send_to);
mailboxes.insert(id.clone(), senders[&id].clone());
let context = Context {
id: id.clone(),
episode,
mailboxes,
shutdown: Shutdown {
mine: shutdowns[&id].clone(),
others: pick(&shutdowns, &init.can_shut_down),
},
log: init.has_logger.then(|| logger.clone()),
};
(
id,
Actor {
behavior: (init.behavior)(context),
ready: Some(ready),
start: started,
mailbox,
},
)
})
.collect();
Self {
actors,
readies,
starts,
}
}
pub async fn run(self, patience: Duration) -> anyhow::Result<()>
where
B: 'static,
{
let stops: Vec<_> = self
.actors
.values()
.map(|actor| actor.behavior.context().shutdown.mine.clone())
.collect();
let mut tasks = JoinSet::new();
for (id, actor) in self.actors {
tasks.spawn(async move {
actor
.run()
.await
.with_context(|| format!("actor {id} failed"))
});
}
let episode = async {
if all_ready(self.readies).await {
for start in self.starts.into_values() {
let _ = start.send(());
}
} else {
for stop in &stops {
stop.cancel();
}
}
wait_for_all(&mut tasks, &stops).await
};
match timeout(patience, episode).await {
Ok(outcome) => outcome,
Err(_) => {
for stop in &stops {
stop.cancel();
}
wait_for_all(&mut tasks, &stops).await?;
bail!("ran out of patience after {patience:?}");
}
}
}
}
type Staged<B> = HashMap<
ActorId,
(
ActorInit<B>,
UnboundedReceiver<Envelope<<B as Behavior>::Message>>,
),
>;
async fn all_ready(readies: HashMap<ActorId, oneshot::Receiver<()>>) -> bool {
for ready in readies.into_values() {
if ready.await.is_err() {
return false;
}
}
true
}
async fn wait_for_all(
tasks: &mut JoinSet<anyhow::Result<()>>,
stops: &[CancellationToken],
) -> anyhow::Result<()> {
let mut first_failure = None;
while let Some(outcome) = tasks.join_next().await {
let outcome = outcome.context("an actor panicked").and_then(|ran| ran);
if let Err(failure) = outcome {
for stop in stops {
stop.cancel();
}
first_failure.get_or_insert(failure);
}
}
first_failure.map_or(Ok(()), Err)
}
fn pick<V: Clone>(
directory: &HashMap<ActorId, V>,
allowed: &HashSet<ActorId>,
) -> HashMap<ActorId, V> {
allowed
.iter()
.map(|id| {
let v = directory
.get(id)
.unwrap_or_else(|| panic!("init names unknown actor {id:?}"));
(id.clone(), v.clone())
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::log::Event;
use crate::message::Message;
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use std::sync::Arc;
use tokio::sync::Semaphore;
use tokio::sync::mpsc::UnboundedSender;
#[derive(Debug, Clone, Serialize, Deserialize)]
struct Note;
impl Message for Note {}
#[derive(Default)]
struct Initialization {
gate: Option<Arc<Semaphore>>,
broken: bool,
}
struct Reporter {
context: Context<Note, ActorId>,
initialization: Initialization,
started: UnboundedSender<ActorId>,
fails: bool,
}
#[async_trait]
impl Behavior for Reporter {
type Message = Note;
type Log = ActorId;
fn context(&self) -> &Context<Note, ActorId> {
&self.context
}
async fn initialize(&mut self) -> anyhow::Result<()> {
if self.initialization.broken {
bail!("cannot initialize");
}
if let Some(gate) = &self.initialization.gate {
gate.acquire().await?.forget();
}
Ok(())
}
async fn start(&mut self) -> anyhow::Result<()> {
if self.fails {
bail!("{} refuses to start", self.context.id);
}
self.context.log(self.context.id.clone());
self.started.send(self.context.id.clone())?;
Ok(())
}
}
struct Stage {
episode: Episode<Reporter>,
starts: UnboundedReceiver<ActorId>,
log: UnboundedReceiver<Event<ActorId>>,
}
fn episode_with(
ids: &[&str],
failing: &[&str],
logging: &[&str],
initialization: impl Fn(&str) -> Initialization,
) -> Stage {
let (started, starts) = unbounded_channel();
let (logger, log) = unbounded_channel();
let init = ids
.iter()
.map(|id| {
let started = started.clone();
let fails = failing.contains(id);
let initialization = initialization(id);
let init = ActorInit {
behavior: Box::new(move |context| Reporter {
context,
initialization,
started,
fails,
}),
can_send_to: HashSet::new(),
can_shut_down: HashSet::new(),
has_logger: logging.contains(id),
};
(id.to_string(), init)
})
.collect();
Stage {
episode: Episode::new(init, logger),
starts,
log,
}
}
fn episode_of(ids: &[&str], failing: &[&str]) -> Stage {
episode_with(ids, failing, &[], |_| Initialization::default())
}
fn episode_with_bob_held_up() -> (Stage, Arc<Semaphore>) {
let gate = Arc::new(Semaphore::new(0));
let slow = Arc::clone(&gate);
let stage = episode_with(&["ann", "bob"], &[], &[], move |id| Initialization {
gate: (id == "bob").then(|| Arc::clone(&slow)),
broken: false,
});
(stage, gate)
}
async fn two_starts(starts: &mut UnboundedReceiver<ActorId>) -> HashSet<ActorId> {
let mut started = HashSet::new();
started.insert(starts.recv().await.unwrap());
started.insert(starts.recv().await.unwrap());
started
}
fn ann_and_bob() -> HashSet<ActorId> {
HashSet::from(["ann".to_string(), "bob".to_string()])
}
fn wired(links: &[(&str, &[&str])]) -> Episode<Reporter> {
let (started, _) = unbounded_channel();
let (logger, _) = unbounded_channel();
let init = links
.iter()
.map(|(id, others)| {
let started = started.clone();
let others: HashSet<ActorId> = others.iter().map(|o| o.to_string()).collect();
let init = ActorInit {
behavior: Box::new(move |context| Reporter {
context,
initialization: Initialization::default(),
started,
fails: false,
}),
can_send_to: others.clone(),
can_shut_down: others,
has_logger: false,
};
(id.to_string(), init)
})
.collect();
Episode::new(init, logger)
}
fn stops<B: Behavior>(episode: &Episode<B>) -> Vec<CancellationToken> {
episode
.actors
.values()
.map(|actor| actor.behavior.context().shutdown.mine.clone())
.collect()
}
#[test]
fn new_wires_each_actor_to_itself_and_the_actors_its_init_names() {
let episode = wired(&[("ann", &["bob"]), ("bob", &[])]);
let ann = episode.actors["ann"].behavior.context();
let mut reaches: Vec<_> = ann.mailboxes.keys().cloned().collect();
reaches.sort();
let stops: Vec<_> = ann.shutdown.others.keys().cloned().collect();
assert_eq!(reaches, ["ann", "bob"]);
assert_eq!(stops, ["bob"]);
let bob = episode.actors["bob"].behavior.context();
let reaches: Vec<_> = bob.mailboxes.keys().cloned().collect();
assert_eq!(reaches, ["bob"]);
assert!(bob.shutdown.others.is_empty());
}
#[test]
#[should_panic(expected = "unknown actor \"zed\"")]
fn new_panics_when_an_init_names_an_unknown_actor() {
wired(&[("ann", &["zed"])]);
}
#[tokio::test]
async fn run_starts_every_actor_then_waits_for_them_to_finish() {
let Stage {
episode,
mut starts,
..
} = episode_of(&["ann", "bob"], &[]);
let stops = stops(&episode);
let running = tokio::spawn(episode.run(Duration::from_secs(60)));
assert_eq!(two_starts(&mut starts).await, ann_and_bob());
for stop in stops {
stop.cancel();
}
running.await.unwrap().unwrap();
}
#[tokio::test(start_paused = true)]
async fn run_gives_up_after_patience() {
let Stage {
episode,
starts: _starts,
..
} = episode_of(&["ann", "bob"], &[]);
let error = episode.run(Duration::from_secs(5)).await.unwrap_err();
assert!(error.to_string().contains("patience"), "{error}");
}
#[tokio::test(start_paused = true)]
async fn a_failing_actor_ends_the_episode_with_its_error() {
let Stage {
episode,
starts: _starts,
..
} = episode_of(&["ann", "bob"], &["bob"]);
let error = episode.run(Duration::from_secs(60)).await.unwrap_err();
let text = format!("{error:#}");
assert!(text.contains("actor bob failed"), "{text}");
assert!(text.contains("bob refuses to start"), "{text}");
}
#[tokio::test(start_paused = true)]
async fn no_actor_starts_until_every_actor_has_initialized() {
let (
Stage {
episode,
mut starts,
..
},
gate,
) = episode_with_bob_held_up();
let stops = stops(&episode);
let running = tokio::spawn(episode.run(Duration::from_secs(60)));
tokio::time::sleep(Duration::from_secs(1)).await;
assert!(starts.try_recv().is_err(), "nobody should have started");
gate.add_permits(1);
assert_eq!(two_starts(&mut starts).await, ann_and_bob());
for stop in stops {
stop.cancel();
}
running.await.unwrap().unwrap();
}
#[tokio::test(start_paused = true)]
async fn an_actor_that_fails_to_initialize_fails_the_episode_before_anyone_starts() {
let Stage {
episode,
mut starts,
..
} = episode_with(&["ann", "bob"], &[], &[], |id| Initialization {
gate: None,
broken: id == "bob",
});
let error = episode.run(Duration::from_secs(60)).await.unwrap_err();
let text = format!("{error:#}");
assert!(text.contains("actor bob failed"), "{text}");
assert!(text.contains("cannot initialize"), "{text}");
assert!(starts.try_recv().is_err(), "nobody should have started");
}
#[tokio::test(start_paused = true)]
async fn patience_runs_out_while_an_actor_is_still_initializing() {
let (
Stage {
episode,
mut starts,
..
},
_gate,
) = episode_with_bob_held_up();
let error = episode.run(Duration::from_secs(5)).await.unwrap_err();
assert!(error.to_string().contains("patience"), "{error}");
assert!(starts.try_recv().is_err(), "nobody should have started");
}
#[tokio::test]
async fn only_an_actor_with_the_logger_logs() {
let Stage {
episode,
mut starts,
mut log,
} = episode_with(&["ann", "bob"], &[], &["ann"], |_| {
Initialization::default()
});
let stops = stops(&episode);
let running = tokio::spawn(episode.run(Duration::from_secs(60)));
starts.recv().await.unwrap();
starts.recv().await.unwrap();
for stop in stops {
stop.cancel();
}
running.await.unwrap().unwrap();
let event = log.recv().await.unwrap();
assert_eq!(event.payload, "ann");
assert!(log.try_recv().is_err(), "bob has no logger");
}
}