use crate::log::{Event, Logger};
use crate::message::{Message, Request};
use anyhow::{Context as _, bail};
use async_trait::async_trait;
use serde::{Serialize, de::DeserializeOwned};
use std::collections::{HashMap, HashSet};
use std::ops::ControlFlow::{self, Break, Continue};
use tokio::sync::mpsc::{UnboundedReceiver, UnboundedSender};
use tokio::sync::oneshot;
use tokio_util::sync::CancellationToken;
use uuid::Uuid;
pub type ActorId = String;
pub type Builder<B> =
Box<dyn FnOnce(Context<<B as Behavior>::Message, <B as Behavior>::Log>) -> B + Send>;
pub struct ActorInit<B: Behavior> {
pub behavior: Builder<B>,
pub can_send_to: HashSet<ActorId>,
pub can_shut_down: HashSet<ActorId>,
pub has_logger: bool,
}
pub(crate) struct Actor<B: Behavior> {
pub(crate) behavior: B,
pub(crate) ready: Option<oneshot::Sender<()>>,
pub(crate) start: oneshot::Receiver<()>,
pub(crate) mailbox: UnboundedReceiver<Envelope<B::Message>>,
}
impl<B: Behavior> Actor<B> {
pub(crate) async fn run(mut self) -> anyhow::Result<()> {
let shutdown = self.behavior.context().shutdown.mine.clone();
let Some(initialized) = run_unless_stopped(&shutdown, self.behavior.initialize()).await
else {
return Ok(());
};
initialized?;
if let Some(ready) = self.ready.take() {
let _ = ready.send(());
}
let mut start_consumed = false;
loop {
let flow = tokio::select! {
biased;
_ = shutdown.cancelled() => Break(()),
signal = &mut self.start, if !start_consumed => {
start_consumed = true;
self.open(signal).await?
}
envelope = self.mailbox.recv() => {
self.deliver(envelope.expect("mailbox closed")).await?
}
};
if flow.is_break() {
break;
}
}
self.behavior.clean_up().await?;
Ok(())
}
async fn open(
&mut self,
signal: Result<(), oneshot::error::RecvError>,
) -> anyhow::Result<ControlFlow<()>> {
if signal.is_err() {
return Ok(Break(()));
}
let shutdown = self.behavior.context().shutdown.mine.clone();
let Some(opened) = run_unless_stopped(&shutdown, self.behavior.start()).await else {
return Ok(Break(()));
};
opened?;
Ok(Continue(()))
}
async fn deliver(&mut self, envelope: Envelope<B::Message>) -> anyhow::Result<ControlFlow<()>> {
let shutdown = self.behavior.context().shutdown.mine.clone();
match envelope {
Envelope::Statement(message) => {
let Some(received) =
run_unless_stopped(&shutdown, self.behavior.receive(&message)).await
else {
return Ok(Break(()));
};
received?;
}
Envelope::Request(request) => {
let Some(answered) =
run_unless_stopped(&shutdown, self.behavior.answer(request.message())).await
else {
return Ok(Break(()));
};
let _ = request.reply(answered?);
}
}
Ok(Continue(()))
}
}
async fn run_unless_stopped<T>(
shutdown: &CancellationToken,
step: impl Future<Output = T>,
) -> Option<T> {
tokio::select! {
biased;
_ = shutdown.cancelled() => None,
outcome = step => Some(outcome),
}
}
pub struct Context<M: Message, L = M> {
pub id: ActorId,
pub episode: Uuid,
pub(crate) mailboxes: HashMap<ActorId, UnboundedSender<Envelope<M>>>,
pub(crate) shutdown: Shutdown,
pub(crate) log: Option<Logger<L>>,
}
impl<M: Message, L> Context<M, L> {
pub fn log(&self, payload: L) {
if let Some(log) = &self.log {
let _ = log.send(Event::now(self.episode, payload));
}
}
pub fn send(&self, message: M, to: HashSet<ActorId>) -> anyhow::Result<()> {
let senders = to
.iter()
.map(|id| self.mailbox_of(id))
.collect::<anyhow::Result<Vec<_>>>()?;
for sender in senders {
let _ = sender.send(Envelope::Statement(message.clone()));
}
Ok(())
}
pub async fn request(
&self,
message: M,
to: HashSet<ActorId>,
) -> anyhow::Result<HashMap<ActorId, Vec<M>>> {
if to.contains(&self.id) {
bail!(
"{} cannot request from itself: it would wait forever",
self.id
);
}
let senders = to
.iter()
.map(|id| self.mailbox_of(id).map(|sender| (id, sender)))
.collect::<anyhow::Result<Vec<_>>>()?;
let mut pending = Vec::new();
for (id, sender) in senders {
let (request, reply) = Request::new(message.clone());
if sender.send(Envelope::Request(request)).is_ok() {
pending.push((id.clone(), reply));
}
}
let mut replies = HashMap::new();
for (id, reply) in pending {
if let Ok(reply) = reply.await {
replies.insert(id, reply.messages);
}
}
Ok(replies)
}
fn mailbox_of(&self, id: &ActorId) -> anyhow::Result<&UnboundedSender<Envelope<M>>> {
self.mailboxes.get(id).with_context(|| {
format!(
"{} cannot send to {id}: not an actor it may send to",
self.id
)
})
}
pub fn stop(&self, who: &ActorId) -> anyhow::Result<()> {
match self.shutdown.others.get(who) {
Some(token) => {
token.cancel();
Ok(())
}
None => bail!(
"{} cannot stop {who}: not an actor it may shut down",
self.id
),
}
}
pub fn shutdown(&self) {
self.shutdown.mine.cancel();
}
}
#[async_trait]
pub trait Behavior: Send {
type Message: Message;
type Log: Serialize + DeserializeOwned + Send + 'static;
fn context(&self) -> &Context<Self::Message, Self::Log>;
async fn initialize(&mut self) -> anyhow::Result<()> {
Ok(())
}
async fn receive(&mut self, _message: &Self::Message) -> anyhow::Result<()> {
Ok(())
}
async fn answer(&mut self, _message: &Self::Message) -> anyhow::Result<Vec<Self::Message>> {
Ok(vec![])
}
async fn start(&mut self) -> anyhow::Result<()> {
Ok(())
}
async fn clean_up(&mut self) -> anyhow::Result<()> {
Ok(())
}
}
pub(crate) struct Shutdown {
pub(crate) mine: CancellationToken,
pub(crate) others: HashMap<ActorId, CancellationToken>,
}
#[derive(Debug)]
pub(crate) enum Envelope<M: Message> {
Statement(M),
Request(Request<M>),
}
#[cfg(test)]
mod tests {
use super::*;
use serde::{Deserialize, Serialize};
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::Semaphore;
use tokio::sync::mpsc::unbounded_channel;
use tokio::time::{sleep, timeout};
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
struct Note(String);
impl Message for Note {}
fn note(text: &str) -> Note {
Note(text.to_string())
}
struct Echo {
context: Context<Note>,
gate: Option<Arc<Semaphore>>,
slow_start: bool,
broken: bool,
}
#[async_trait]
impl Behavior for Echo {
type Message = Note;
type Log = Note;
fn context(&self) -> &Context<Note> {
&self.context
}
async fn receive(&mut self, _message: &Note) -> anyhow::Result<()> {
self.step().await
}
async fn answer(&mut self, message: &Note) -> anyhow::Result<Vec<Note>> {
self.step().await?;
Ok(vec![message.clone()])
}
async fn start(&mut self) -> anyhow::Result<()> {
if self.slow_start
&& let Some(gate) = &self.gate
{
gate.acquire().await?.forget();
}
Ok(())
}
}
impl Echo {
async fn step(&self) -> anyhow::Result<()> {
if self.broken {
bail!("echo is broken");
}
if let Some(gate) = &self.gate {
gate.acquire().await?.forget();
}
Ok(())
}
}
struct Mute(Context<Note>);
#[async_trait]
impl Behavior for Mute {
type Message = Note;
type Log = Note;
fn context(&self) -> &Context<Note> {
&self.0
}
}
struct Tally {
context: Context<Note>,
seen: usize,
}
#[async_trait]
impl Behavior for Tally {
type Message = Note;
type Log = Note;
fn context(&self) -> &Context<Note> {
&self.context
}
async fn receive(&mut self, _message: &Note) -> anyhow::Result<()> {
self.seen += 1;
Ok(())
}
async fn answer(&mut self, _message: &Note) -> anyhow::Result<Vec<Note>> {
Ok(vec![note("seen"); self.seen])
}
}
struct Rig<B: Behavior = Echo> {
actor: Actor<B>,
sender: UnboundedSender<Envelope<Note>>,
start: oneshot::Sender<()>,
stop: CancellationToken,
}
impl<B: Behavior<Message = Note, Log = Note>> Rig<B> {
fn context(&self) -> &Context<Note> {
self.actor.behavior.context()
}
}
fn rig(
name: &str,
others: HashMap<ActorId, UnboundedSender<Envelope<Note>>>,
can_stop: HashMap<ActorId, CancellationToken>,
) -> Rig {
rig_with(name, others, can_stop, |context| Echo {
context,
gate: None,
slow_start: false,
broken: false,
})
}
fn rig_with<B: Behavior<Message = Note, Log = Note>>(
name: &str,
others: HashMap<ActorId, UnboundedSender<Envelope<Note>>>,
can_stop: HashMap<ActorId, CancellationToken>,
build: impl FnOnce(Context<Note>) -> B,
) -> Rig<B> {
let (sender, mailbox) = unbounded_channel();
let (ready, _) = oneshot::channel();
let (start, started) = oneshot::channel();
let stop = CancellationToken::new();
let mut mailboxes = others;
mailboxes.insert(id(name), sender.clone());
let context = Context {
id: id(name),
episode: Uuid::new_v4(),
mailboxes,
shutdown: Shutdown {
mine: stop.clone(),
others: can_stop,
},
log: None,
};
let actor = Actor {
behavior: build(context),
ready: Some(ready),
start: started,
mailbox,
};
Rig {
actor,
sender,
start,
stop,
}
}
fn id(name: &str) -> ActorId {
name.to_string()
}
#[tokio::test]
async fn send_reaches_each_actor_it_is_addressed_to() {
let (bob, mut bob_mailbox) = unbounded_channel();
let (cat, mut cat_mailbox) = unbounded_channel();
let (dan, mut dan_mailbox) = unbounded_channel();
let ann = rig(
"ann",
HashMap::from([(id("bob"), bob), (id("cat"), cat), (id("dan"), dan)]),
HashMap::new(),
);
ann.context()
.send(note("hello"), HashSet::from([id("bob"), id("cat")]))
.unwrap();
for mailbox in [&mut bob_mailbox, &mut cat_mailbox] {
let heard = mailbox.recv().await.unwrap();
assert!(
matches!(&heard, Envelope::Statement(Note(text)) if text == "hello"),
"{heard:?}"
);
}
assert!(dan_mailbox.try_recv().is_err(), "nothing should reach dan");
}
#[tokio::test]
async fn send_fails_without_sending_if_a_recipient_may_not_be_addressed() {
let (bob, mut bob_mailbox) = unbounded_channel();
let ann = rig("ann", HashMap::from([(id("bob"), bob)]), HashMap::new());
let error = ann
.context()
.send(note("psst"), HashSet::from([id("bob"), id("zed")]))
.unwrap_err();
assert!(error.to_string().contains("zed"), "{error}");
assert!(bob_mailbox.try_recv().is_err(), "nothing should reach bob");
}
#[tokio::test]
async fn send_to_itself_lands_in_the_actors_own_mailbox() {
let mut ann = rig("ann", HashMap::new(), HashMap::new());
ann.context()
.send(note("remember this"), HashSet::from([id("ann")]))
.unwrap();
let heard = ann.actor.mailbox.recv().await.unwrap();
assert!(
matches!(&heard, Envelope::Statement(Note(text)) if text == "remember this"),
"{heard:?}"
);
}
#[tokio::test]
async fn request_collects_a_reply_from_every_recipient() {
let bob = rig("bob", HashMap::new(), HashMap::new());
let cat = rig("cat", HashMap::new(), HashMap::new());
let ann = rig(
"ann",
HashMap::from([
(id("bob"), bob.sender.clone()),
(id("cat"), cat.sender.clone()),
]),
HashMap::new(),
);
let bob_running = tokio::spawn(bob.actor.run());
let cat_running = tokio::spawn(cat.actor.run());
bob.start.send(()).unwrap();
cat.start.send(()).unwrap();
let replies = ann
.context()
.request(note("who's there?"), HashSet::from([id("bob"), id("cat")]))
.await
.unwrap();
let expected = HashMap::from([
(id("bob"), vec![note("who's there?")]),
(id("cat"), vec![note("who's there?")]),
]);
assert_eq!(replies, expected);
bob.stop.cancel();
cat.stop.cancel();
bob_running.await.unwrap().unwrap();
cat_running.await.unwrap().unwrap();
}
#[tokio::test]
async fn a_behavior_keeps_its_state_between_steps() {
let bob = rig_with("bob", HashMap::new(), HashMap::new(), |context| Tally {
context,
seen: 0,
});
let ann = rig(
"ann",
HashMap::from([(id("bob"), bob.sender.clone())]),
HashMap::new(),
);
let running = tokio::spawn(bob.actor.run());
bob.start.send(()).unwrap();
for _ in 0..2 {
ann.context()
.send(note("one more"), HashSet::from([id("bob")]))
.unwrap();
}
let replies = ann
.context()
.request(note("how many?"), HashSet::from([id("bob")]))
.await
.unwrap();
let expected = HashMap::from([(id("bob"), vec![note("seen"), note("seen")])]);
assert_eq!(replies, expected);
bob.stop.cancel();
running.await.unwrap().unwrap();
}
#[tokio::test]
async fn the_default_receive_ignores_the_statement() {
let bob = rig_with("bob", HashMap::new(), HashMap::new(), Mute);
let ann = rig(
"ann",
HashMap::from([(id("bob"), bob.sender.clone())]),
HashMap::new(),
);
let running = tokio::spawn(bob.actor.run());
bob.start.send(()).unwrap();
ann.context()
.send(note("whatever"), HashSet::from([id("bob")]))
.unwrap();
let replies = ann
.context()
.request(note("still there?"), HashSet::from([id("bob")]))
.await
.unwrap();
assert_eq!(replies, HashMap::from([(id("bob"), vec![])]));
bob.stop.cancel();
running.await.unwrap().unwrap();
}
#[tokio::test]
async fn the_default_answer_is_nothing() {
let bob = rig_with("bob", HashMap::new(), HashMap::new(), Mute);
let ann = rig(
"ann",
HashMap::from([(id("bob"), bob.sender.clone())]),
HashMap::new(),
);
let running = tokio::spawn(bob.actor.run());
bob.start.send(()).unwrap();
let replies = ann
.context()
.request(note("anything?"), HashSet::from([id("bob")]))
.await
.unwrap();
assert_eq!(replies, HashMap::from([(id("bob"), vec![])]));
bob.stop.cancel();
running.await.unwrap().unwrap();
}
#[tokio::test]
async fn request_fails_without_sending_if_a_recipient_may_not_be_addressed() {
let (bob, mut bob_mailbox) = unbounded_channel();
let ann = rig("ann", HashMap::from([(id("bob"), bob)]), HashMap::new());
let error = ann
.context()
.request(note("psst"), HashSet::from([id("bob"), id("zed")]))
.await
.unwrap_err();
assert!(error.to_string().contains("zed"), "{error}");
assert!(bob_mailbox.try_recv().is_err(), "nothing should reach bob");
}
#[tokio::test]
async fn request_refuses_to_ask_the_actor_itself() {
let ann = rig("ann", HashMap::new(), HashMap::new());
let error = ann
.context()
.request(note("hello me"), HashSet::from([id("ann")]))
.await
.unwrap_err();
assert!(error.to_string().contains("itself"), "{error}");
}
#[tokio::test]
async fn request_leaves_out_a_recipient_that_has_stopped() {
let (bob, bob_mailbox) = unbounded_channel();
let ann = rig("ann", HashMap::from([(id("bob"), bob)]), HashMap::new());
drop(bob_mailbox);
let replies = ann
.context()
.request(note("anyone?"), HashSet::from([id("bob")]))
.await
.unwrap();
assert!(replies.is_empty(), "{replies:?}");
}
#[tokio::test]
async fn stop_cancels_an_actor_it_may_shut_down_and_refuses_others() {
let bob = rig("bob", HashMap::new(), HashMap::new());
let ann = rig(
"ann",
HashMap::new(),
HashMap::from([(id("bob"), bob.stop.clone())]),
);
ann.context().stop(&id("bob")).unwrap();
assert!(bob.stop.is_cancelled());
let error = ann.context().stop(&id("zed")).unwrap_err();
assert!(error.to_string().contains("zed"), "{error}");
}
#[tokio::test]
async fn log_sends_a_stamped_event_down_the_logger_if_there_is_one() {
let (logger, mut events) = unbounded_channel();
let mut ann = rig("ann", HashMap::new(), HashMap::new());
ann.actor.behavior.context.log = Some(logger);
ann.context().log(note("for the record"));
let event = events.recv().await.unwrap();
assert_eq!(event.payload, note("for the record"));
}
#[tokio::test]
async fn log_does_nothing_without_a_logger() {
let ann = rig("ann", HashMap::new(), HashMap::new());
ann.context().log(note("into the void"));
}
#[tokio::test]
async fn shutdown_cancels_the_actors_own_token() {
let ann = rig("ann", HashMap::new(), HashMap::new());
ann.context().shutdown();
assert!(ann.stop.is_cancelled());
}
fn gated_rig(name: &str) -> (Rig, Arc<Semaphore>) {
let gate = Arc::new(Semaphore::new(0));
let mut rig = rig(name, HashMap::new(), HashMap::new());
rig.actor.behavior.gate = Some(Arc::clone(&gate));
(rig, gate)
}
#[tokio::test(start_paused = true)]
async fn a_kill_stops_an_actor_in_the_middle_of_a_step() {
let (bob, _gate) = gated_rig("bob");
let running = tokio::spawn(bob.actor.run());
bob.start.send(()).unwrap();
bob.sender
.send(Envelope::Statement(note("take your time")))
.unwrap();
sleep(Duration::from_secs(1)).await;
bob.stop.cancel();
let stopped = timeout(Duration::from_secs(5), running).await;
stopped.expect("bob should stop").unwrap().unwrap();
}
#[tokio::test(start_paused = true)]
async fn a_kill_stops_an_actor_in_the_middle_of_starting() {
let (mut bob, _gate) = gated_rig("bob");
bob.actor.behavior.slow_start = true;
let running = tokio::spawn(bob.actor.run());
bob.start.send(()).unwrap();
sleep(Duration::from_secs(1)).await;
bob.stop.cancel();
let stopped = timeout(Duration::from_secs(5), running).await;
stopped.expect("bob should stop").unwrap().unwrap();
}
#[tokio::test(start_paused = true)]
async fn a_behavior_with_no_opening_move_starts_and_waits() {
let bob = rig_with("bob", HashMap::new(), HashMap::new(), Mute);
let running = tokio::spawn(bob.actor.run());
bob.start.send(()).unwrap();
sleep(Duration::from_secs(1)).await;
bob.stop.cancel();
running.await.unwrap().unwrap();
}
#[tokio::test]
async fn an_actor_stops_when_the_episode_is_gone_before_it_starts() {
let bob = rig("bob", HashMap::new(), HashMap::new());
let running = tokio::spawn(bob.actor.run());
drop(bob.start);
running.await.unwrap().unwrap();
}
#[tokio::test]
async fn a_behavior_that_fails_to_receive_a_statement_fails_the_actor() {
let mut bob = rig("bob", HashMap::new(), HashMap::new());
bob.actor.behavior.broken = true;
let running = tokio::spawn(bob.actor.run());
bob.start.send(()).unwrap();
bob.sender.send(Envelope::Statement(note("hello"))).unwrap();
let error = running.await.unwrap().unwrap_err();
assert!(error.to_string().contains("broken"), "{error}");
}
#[tokio::test]
async fn a_behavior_that_fails_to_answer_a_request_fails_the_actor() {
let mut bob = rig("bob", HashMap::new(), HashMap::new());
bob.actor.behavior.broken = true;
let running = tokio::spawn(bob.actor.run());
bob.start.send(()).unwrap();
let (request, _reply) = Request::new(note("well?"));
bob.sender.send(Envelope::Request(request)).unwrap();
let error = running.await.unwrap().unwrap_err();
assert!(error.to_string().contains("broken"), "{error}");
}
#[tokio::test(start_paused = true)]
async fn a_recipient_killed_mid_step_is_left_out_of_the_replies() {
let (bob, _gate) = gated_rig("bob");
let ann = rig(
"ann",
HashMap::from([(id("bob"), bob.sender.clone())]),
HashMap::new(),
);
let bob_running = tokio::spawn(bob.actor.run());
bob.start.send(()).unwrap();
let asking = tokio::spawn(async move {
ann.context()
.request(note("well?"), HashSet::from([id("bob")]))
.await
});
sleep(Duration::from_secs(1)).await;
bob.stop.cancel();
let replies = asking.await.unwrap().unwrap();
assert!(replies.is_empty(), "{replies:?}");
bob_running.await.unwrap().unwrap();
}
#[tokio::test(start_paused = true)]
async fn a_responder_whose_asker_gave_up_carries_on() {
let (bob, gate) = gated_rig("bob");
let ann = rig(
"ann",
HashMap::from([(id("bob"), bob.sender.clone())]),
HashMap::new(),
);
let bob_running = tokio::spawn(bob.actor.run());
bob.start.send(()).unwrap();
let asking = ann
.context()
.request(note("well?"), HashSet::from([id("bob")]));
let gave_up = tokio::time::timeout(Duration::from_secs(1), asking).await;
assert!(gave_up.is_err(), "{gave_up:?}");
gate.add_permits(1);
sleep(Duration::from_secs(1)).await;
bob.stop.cancel();
bob_running.await.unwrap().unwrap();
}
}