1use crate::actor::{Actor, ActorId, ActorInit, Behavior, Context, Envelope, Shutdown};
6use crate::log::Logger;
7use anyhow::{Context as _, bail};
8use std::collections::{HashMap, HashSet};
9use std::time::Duration;
10use tokio::sync::mpsc::{UnboundedReceiver, unbounded_channel};
11use tokio::sync::oneshot;
12use tokio::task::JoinSet;
13use tokio::time::timeout;
14use tokio_util::sync::CancellationToken;
15use uuid::Uuid;
16
17pub struct Episode<B: Behavior> {
20 actors: HashMap<ActorId, Actor<B>>,
21 readies: HashMap<ActorId, oneshot::Receiver<()>>,
23 starts: HashMap<ActorId, oneshot::Sender<()>>,
25}
26
27impl<B: Behavior> Episode<B> {
28 pub fn new(init: HashMap<ActorId, ActorInit<B>>, logger: Logger<B::Log>) -> Self {
36 let episode = Uuid::new_v4();
41 let mut senders = HashMap::new();
42 let mut shutdowns = HashMap::new();
43 let staged: Staged<B> = init
44 .into_iter()
45 .map(|(id, init)| {
46 let (tx, rx) = unbounded_channel();
47 senders.insert(id.clone(), tx);
48 shutdowns.insert(id.clone(), CancellationToken::new());
49 (id, (init, rx))
50 })
51 .collect();
52 let mut readies = HashMap::new();
56 let mut starts = HashMap::new();
57 let actors = staged
58 .into_iter()
59 .map(|(id, (init, mailbox))| {
60 let (ready, is_ready) = oneshot::channel();
61 readies.insert(id.clone(), is_ready);
62 let (start, started) = oneshot::channel();
63 starts.insert(id.clone(), start);
64 let mut mailboxes = pick(&senders, &init.can_send_to);
65 mailboxes.insert(id.clone(), senders[&id].clone());
66 let context = Context {
67 id: id.clone(),
68 episode,
69 mailboxes,
70 shutdown: Shutdown {
71 mine: shutdowns[&id].clone(),
72 others: pick(&shutdowns, &init.can_shut_down),
73 },
74 log: init.has_logger.then(|| logger.clone()),
75 };
76 (
77 id,
78 Actor {
79 behavior: (init.behavior)(context),
80 ready: Some(ready),
81 start: started,
82 mailbox,
83 },
84 )
85 })
86 .collect();
87 Self {
88 actors,
89 readies,
90 starts,
91 }
92 }
93
94 pub async fn run(self, patience: Duration) -> anyhow::Result<()>
108 where
109 B: 'static,
110 {
111 let stops: Vec<_> = self
112 .actors
113 .values()
114 .map(|actor| actor.behavior.context().shutdown.mine.clone())
115 .collect();
116 let mut tasks = JoinSet::new();
117 for (id, actor) in self.actors {
118 tasks.spawn(async move {
119 actor
120 .run()
121 .await
122 .with_context(|| format!("actor {id} failed"))
123 });
124 }
125 let episode = async {
126 if all_ready(self.readies).await {
127 for start in self.starts.into_values() {
128 let _ = start.send(());
132 }
133 } else {
134 for stop in &stops {
137 stop.cancel();
138 }
139 }
140 wait_for_all(&mut tasks, &stops).await
141 };
142 match timeout(patience, episode).await {
143 Ok(outcome) => outcome,
144 Err(_) => {
145 for stop in &stops {
146 stop.cancel();
147 }
148 wait_for_all(&mut tasks, &stops).await?;
149 bail!("ran out of patience after {patience:?}");
150 }
151 }
152 }
153}
154
155type Staged<B> = HashMap<
158 ActorId,
159 (
160 ActorInit<B>,
161 UnboundedReceiver<Envelope<<B as Behavior>::Message>>,
162 ),
163>;
164
165async fn all_ready(readies: HashMap<ActorId, oneshot::Receiver<()>>) -> bool {
168 for ready in readies.into_values() {
169 if ready.await.is_err() {
170 return false;
171 }
172 }
173 true
174}
175
176async fn wait_for_all(
179 tasks: &mut JoinSet<anyhow::Result<()>>,
180 stops: &[CancellationToken],
181) -> anyhow::Result<()> {
182 let mut first_failure = None;
183 while let Some(outcome) = tasks.join_next().await {
184 let outcome = outcome.context("an actor panicked").and_then(|ran| ran);
185 if let Err(failure) = outcome {
186 for stop in stops {
187 stop.cancel();
188 }
189 first_failure.get_or_insert(failure);
190 }
191 }
192 first_failure.map_or(Ok(()), Err)
193}
194
195fn pick<V: Clone>(
199 directory: &HashMap<ActorId, V>,
200 allowed: &HashSet<ActorId>,
201) -> HashMap<ActorId, V> {
202 allowed
203 .iter()
204 .map(|id| {
205 let v = directory
206 .get(id)
207 .unwrap_or_else(|| panic!("init names unknown actor {id:?}"));
208 (id.clone(), v.clone())
209 })
210 .collect()
211}
212
213#[cfg(test)]
214mod tests {
215 use super::*;
216 use crate::log::Event;
217 use crate::message::Message;
218 use async_trait::async_trait;
219 use serde::{Deserialize, Serialize};
220 use std::sync::Arc;
221 use tokio::sync::Semaphore;
222 use tokio::sync::mpsc::UnboundedSender;
223
224 #[derive(Debug, Clone, Serialize, Deserialize)]
225 struct Note;
226 impl Message for Note {}
227
228 #[derive(Default)]
231 struct Initialization {
232 gate: Option<Arc<Semaphore>>,
234 broken: bool,
236 }
237
238 struct Reporter {
241 context: Context<Note, ActorId>,
242 initialization: Initialization,
243 started: UnboundedSender<ActorId>,
244 fails: bool,
245 }
246 #[async_trait]
247 impl Behavior for Reporter {
248 type Message = Note;
249 type Log = ActorId;
251 fn context(&self) -> &Context<Note, ActorId> {
252 &self.context
253 }
254 async fn initialize(&mut self) -> anyhow::Result<()> {
255 if self.initialization.broken {
256 bail!("cannot initialize");
257 }
258 if let Some(gate) = &self.initialization.gate {
259 gate.acquire().await?.forget();
260 }
261 Ok(())
262 }
263 async fn start(&mut self) -> anyhow::Result<()> {
264 if self.fails {
265 bail!("{} refuses to start", self.context.id);
266 }
267 self.context.log(self.context.id.clone());
268 self.started.send(self.context.id.clone())?;
269 Ok(())
270 }
271 }
272
273 struct Stage {
275 episode: Episode<Reporter>,
276 starts: UnboundedReceiver<ActorId>,
278 log: UnboundedReceiver<Event<ActorId>>,
280 }
281
282 fn episode_with(
287 ids: &[&str],
288 failing: &[&str],
289 logging: &[&str],
290 initialization: impl Fn(&str) -> Initialization,
291 ) -> Stage {
292 let (started, starts) = unbounded_channel();
293 let (logger, log) = unbounded_channel();
294 let init = ids
295 .iter()
296 .map(|id| {
297 let started = started.clone();
298 let fails = failing.contains(id);
299 let initialization = initialization(id);
300 let init = ActorInit {
301 behavior: Box::new(move |context| Reporter {
302 context,
303 initialization,
304 started,
305 fails,
306 }),
307 can_send_to: HashSet::new(),
308 can_shut_down: HashSet::new(),
309 has_logger: logging.contains(id),
310 };
311 (id.to_string(), init)
312 })
313 .collect();
314 Stage {
315 episode: Episode::new(init, logger),
316 starts,
317 log,
318 }
319 }
320
321 fn episode_of(ids: &[&str], failing: &[&str]) -> Stage {
324 episode_with(ids, failing, &[], |_| Initialization::default())
325 }
326
327 fn episode_with_bob_held_up() -> (Stage, Arc<Semaphore>) {
330 let gate = Arc::new(Semaphore::new(0));
331 let slow = Arc::clone(&gate);
332 let stage = episode_with(&["ann", "bob"], &[], &[], move |id| Initialization {
333 gate: (id == "bob").then(|| Arc::clone(&slow)),
334 broken: false,
335 });
336 (stage, gate)
337 }
338
339 async fn two_starts(starts: &mut UnboundedReceiver<ActorId>) -> HashSet<ActorId> {
341 let mut started = HashSet::new();
342 started.insert(starts.recv().await.unwrap());
343 started.insert(starts.recv().await.unwrap());
344 started
345 }
346
347 fn ann_and_bob() -> HashSet<ActorId> {
348 HashSet::from(["ann".to_string(), "bob".to_string()])
349 }
350
351 fn wired(links: &[(&str, &[&str])]) -> Episode<Reporter> {
354 let (started, _) = unbounded_channel();
355 let (logger, _) = unbounded_channel();
356 let init = links
357 .iter()
358 .map(|(id, others)| {
359 let started = started.clone();
360 let others: HashSet<ActorId> = others.iter().map(|o| o.to_string()).collect();
361 let init = ActorInit {
362 behavior: Box::new(move |context| Reporter {
363 context,
364 initialization: Initialization::default(),
365 started,
366 fails: false,
367 }),
368 can_send_to: others.clone(),
369 can_shut_down: others,
370 has_logger: false,
371 };
372 (id.to_string(), init)
373 })
374 .collect();
375 Episode::new(init, logger)
376 }
377
378 fn stops<B: Behavior>(episode: &Episode<B>) -> Vec<CancellationToken> {
379 episode
380 .actors
381 .values()
382 .map(|actor| actor.behavior.context().shutdown.mine.clone())
383 .collect()
384 }
385
386 #[test]
387 fn new_wires_each_actor_to_itself_and_the_actors_its_init_names() {
388 let episode = wired(&[("ann", &["bob"]), ("bob", &[])]);
389
390 let ann = episode.actors["ann"].behavior.context();
391 let mut reaches: Vec<_> = ann.mailboxes.keys().cloned().collect();
392 reaches.sort();
393 let stops: Vec<_> = ann.shutdown.others.keys().cloned().collect();
394 assert_eq!(reaches, ["ann", "bob"]);
395 assert_eq!(stops, ["bob"]);
396 let bob = episode.actors["bob"].behavior.context();
397 let reaches: Vec<_> = bob.mailboxes.keys().cloned().collect();
398 assert_eq!(reaches, ["bob"]);
399 assert!(bob.shutdown.others.is_empty());
400 }
401
402 #[test]
403 #[should_panic(expected = "unknown actor \"zed\"")]
404 fn new_panics_when_an_init_names_an_unknown_actor() {
405 wired(&[("ann", &["zed"])]);
406 }
407
408 #[tokio::test]
409 async fn run_starts_every_actor_then_waits_for_them_to_finish() {
410 let Stage {
411 episode,
412 mut starts,
413 ..
414 } = episode_of(&["ann", "bob"], &[]);
415 let stops = stops(&episode);
416 let running = tokio::spawn(episode.run(Duration::from_secs(60)));
417
418 assert_eq!(two_starts(&mut starts).await, ann_and_bob());
419
420 for stop in stops {
421 stop.cancel();
422 }
423 running.await.unwrap().unwrap();
424 }
425
426 #[tokio::test(start_paused = true)]
427 async fn run_gives_up_after_patience() {
428 let Stage {
430 episode,
431 starts: _starts,
432 ..
433 } = episode_of(&["ann", "bob"], &[]);
434 let error = episode.run(Duration::from_secs(5)).await.unwrap_err();
435 assert!(error.to_string().contains("patience"), "{error}");
436 }
437
438 #[tokio::test(start_paused = true)]
439 async fn a_failing_actor_ends_the_episode_with_its_error() {
440 let Stage {
442 episode,
443 starts: _starts,
444 ..
445 } = episode_of(&["ann", "bob"], &["bob"]);
446 let error = episode.run(Duration::from_secs(60)).await.unwrap_err();
447 let text = format!("{error:#}");
448 assert!(text.contains("actor bob failed"), "{text}");
449 assert!(text.contains("bob refuses to start"), "{text}");
450 }
451
452 #[tokio::test(start_paused = true)]
453 async fn no_actor_starts_until_every_actor_has_initialized() {
454 let (
455 Stage {
456 episode,
457 mut starts,
458 ..
459 },
460 gate,
461 ) = episode_with_bob_held_up();
462 let stops = stops(&episode);
463 let running = tokio::spawn(episode.run(Duration::from_secs(60)));
464
465 tokio::time::sleep(Duration::from_secs(1)).await;
468 assert!(starts.try_recv().is_err(), "nobody should have started");
469
470 gate.add_permits(1);
471 assert_eq!(two_starts(&mut starts).await, ann_and_bob());
472
473 for stop in stops {
474 stop.cancel();
475 }
476 running.await.unwrap().unwrap();
477 }
478
479 #[tokio::test(start_paused = true)]
480 async fn an_actor_that_fails_to_initialize_fails_the_episode_before_anyone_starts() {
481 let Stage {
482 episode,
483 mut starts,
484 ..
485 } = episode_with(&["ann", "bob"], &[], &[], |id| Initialization {
486 gate: None,
487 broken: id == "bob",
488 });
489
490 let error = episode.run(Duration::from_secs(60)).await.unwrap_err();
491
492 let text = format!("{error:#}");
493 assert!(text.contains("actor bob failed"), "{text}");
494 assert!(text.contains("cannot initialize"), "{text}");
495 assert!(starts.try_recv().is_err(), "nobody should have started");
496 }
497
498 #[tokio::test(start_paused = true)]
499 async fn patience_runs_out_while_an_actor_is_still_initializing() {
500 let (
501 Stage {
502 episode,
503 mut starts,
504 ..
505 },
506 _gate,
507 ) = episode_with_bob_held_up();
508
509 let error = episode.run(Duration::from_secs(5)).await.unwrap_err();
510
511 assert!(error.to_string().contains("patience"), "{error}");
512 assert!(starts.try_recv().is_err(), "nobody should have started");
513 }
514
515 #[tokio::test]
516 async fn only_an_actor_with_the_logger_logs() {
517 let Stage {
518 episode,
519 mut starts,
520 mut log,
521 } = episode_with(&["ann", "bob"], &[], &["ann"], |_| {
522 Initialization::default()
523 });
524 let stops = stops(&episode);
525 let running = tokio::spawn(episode.run(Duration::from_secs(60)));
526 starts.recv().await.unwrap();
527 starts.recv().await.unwrap();
528 for stop in stops {
529 stop.cancel();
530 }
531 running.await.unwrap().unwrap();
532
533 let event = log.recv().await.unwrap();
534 assert_eq!(event.payload, "ann");
535 assert!(log.try_recv().is_err(), "bob has no logger");
536 }
537}