#![allow(clippy::mutable_key_type)]
use crate::Node;
use crate::message::Message;
use crate::utils::random_string;
use async_trait::async_trait;
use futures_util::Future;
use parking_lot::RwLock;
use std::collections::HashMap;
use std::fmt;
use std::hash::{Hash, Hasher};
use std::marker::Send;
use std::sync::Arc;
use tokio::sync::mpsc::{
Receiver, Sender, UnboundedReceiver, UnboundedSender, channel, unbounded_channel,
};
use tokio::sync::watch;
use tokio::task::JoinHandle;
#[derive(Clone, Debug)]
enum AddrSender {
Unbounded(UnboundedSender<Message>),
Bounded(Sender<Message>),
}
enum AddrReceiver {
Unbounded(UnboundedReceiver<Message>),
Bounded(Receiver<Message>),
}
impl AddrReceiver {
async fn recv(&mut self) -> Option<Message> {
match self {
Self::Unbounded(r) => r.recv().await,
Self::Bounded(r) => r.recv().await,
}
}
}
#[async_trait]
pub trait Actor: Send + Sync + 'static {
async fn handle(&mut self, message: Message, context: &ActorContext);
async fn pre_start(&mut self, _context: &ActorContext) {}
async fn stopping(&mut self, _context: &ActorContext) {}
fn subscribe_to_everything(&self) -> bool {
false
}
fn try_clone_storage(&self) -> Option<Box<dyn Actor>> {
None
}
}
impl dyn Actor {
async fn run(
&mut self,
mut receiver: AddrReceiver,
mut stop_receiver: Receiver<()>,
mut context: ActorContext,
) {
self.pre_start(&context).await;
loop {
tokio::select! {
_v = stop_receiver.recv() => {
context.stop();
break;
},
opt_msg = receiver.recv() => {
let msg = match opt_msg {
Some(msg) => msg,
None => break,
};
self.handle(msg, &context).await
}
}
}
self.stopping(&context).await;
}
}
#[derive(Clone)]
pub struct ActorContext {
pub peer_id: Arc<RwLock<String>>,
pub router: Addr,
stop_signals: Arc<RwLock<HashMap<Addr, Sender<()>>>>,
task_handles: Arc<RwLock<Vec<JoinHandle<()>>>>,
pub addr: Addr,
pub is_stopped: Arc<RwLock<bool>>,
pub shutdown_rx: watch::Receiver<bool>,
pub node: Option<Node>,
}
impl ActorContext {
pub fn new(peer_id: String) -> Self {
Self {
addr: Addr::noop(),
stop_signals: Arc::new(RwLock::new(HashMap::new())),
task_handles: Arc::new(RwLock::new(Vec::new())),
peer_id: Arc::new(RwLock::new(peer_id)),
router: Addr::noop(),
is_stopped: Arc::new(RwLock::new(false)),
shutdown_rx: watch::channel(false).1,
node: None,
}
}
pub fn child_actor_count(&self) -> usize {
self.stop_signals.read().len()
}
fn child_context(&self, addr: Addr, stop_signal: Sender<()>) -> Self {
let mut stop_signals = HashMap::new();
stop_signals.insert(addr.clone(), stop_signal);
Self {
addr,
stop_signals: Arc::new(RwLock::new(stop_signals)),
task_handles: Arc::new(RwLock::new(Vec::new())),
peer_id: self.peer_id.clone(),
router: self.router.clone(),
is_stopped: self.is_stopped.clone(),
shutdown_rx: self.shutdown_rx.clone(),
node: self.node.clone(),
}
}
pub fn start_actor(&self, actor: Box<dyn Actor>) -> Addr {
self.start_actor_or_router(actor, false, None)
}
pub fn start_actor_bounded(&self, actor: Box<dyn Actor>, bound: usize) -> Addr {
self.start_actor_or_router(actor, false, Some(bound))
}
pub fn start_router(&self, actor: Box<dyn Actor>) -> Addr {
self.start_actor_or_router(actor, true, None)
}
pub fn child_task<T>(&self, task: T)
where
T: Future<Output = ()> + Send + 'static,
{
let handle = tokio::spawn(task);
self.task_handles.write().push(handle);
}
pub fn blocking_child_task<F>(&self, task: F)
where
F: FnOnce() + Send + 'static,
{
let handle = tokio::task::spawn_blocking(task);
self.task_handles.write().push(handle);
}
fn start_actor_or_router(
&self,
mut actor: Box<dyn Actor>,
is_router: bool,
bound: Option<usize>,
) -> Addr {
let (addr, receiver) = match bound {
Some(cap) => {
let (sender, receiver) = channel::<Message>(cap);
(Addr::new_bounded(sender), AddrReceiver::Bounded(receiver))
}
None => {
let (sender, receiver) = unbounded_channel::<Message>();
(Addr::new(sender), AddrReceiver::Unbounded(receiver))
}
};
let (stop_sender, stop_receiver) = channel(1);
let mut new_context = self.child_context(addr.clone(), stop_sender.clone());
if is_router {
new_context.router = addr.clone();
}
self.stop_signals.write().insert(addr.clone(), stop_sender);
let stop_signals = self.stop_signals.clone();
let addr_clone = addr.clone();
tokio::spawn(async move {
actor.run(receiver, stop_receiver, new_context).await;
stop_signals.write().remove(&addr_clone);
});
addr
}
pub fn stop(&mut self) {
for handle in self.task_handles.read().iter() {
handle.abort();
}
for signal in self.stop_signals.read().values() {
let _ = signal.try_send(());
}
self.node = None;
*self.is_stopped.write() = true;
}
}
#[derive(Clone, Debug)]
pub struct Addr {
id: String,
sender: AddrSender,
}
impl Addr {
pub fn new(sender: UnboundedSender<Message>) -> Self {
Self {
id: random_string(32),
sender: AddrSender::Unbounded(sender),
}
}
pub fn new_bounded(sender: Sender<Message>) -> Self {
Self {
id: random_string(32),
sender: AddrSender::Bounded(sender),
}
}
#[allow(clippy::result_unit_err)] pub fn send(&self, msg: Message) -> Result<(), ()> {
match &self.sender {
AddrSender::Unbounded(s) => s.send(msg).map_err(|_| ()),
AddrSender::Bounded(s) => s.try_send(msg).map_err(|_| ()),
}
}
pub fn noop() -> Addr {
let (sender, _receiver) = unbounded_channel::<Message>();
Addr::new(sender)
}
}
impl PartialEq for Addr {
fn eq(&self, other: &Addr) -> bool {
self.id == other.id
}
}
impl Eq for Addr {}
impl Hash for Addr {
fn hash<H: Hasher>(&self, state: &mut H) {
self.id.hash(state);
}
}
impl fmt::Display for Addr {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "actor:{}", self.id)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_addr_equality() {
let (s1, _r1) = unbounded_channel::<Message>();
let (s2, _r2) = unbounded_channel::<Message>();
let a1 = Addr::new(s1);
let a2 = Addr::new(s2);
assert_ne!(a1, a2, "different addrs are not equal");
assert_eq!(a1, a1.clone(), "clone is equal");
}
#[test]
fn test_addr_hash() {
let (s1, _r1) = unbounded_channel::<Message>();
let a1 = Addr::new(s1);
let a2 = a1.clone();
let mut set = std::collections::HashSet::new();
set.insert(a1);
assert!(set.contains(&a2), "clone should be found in HashSet");
}
#[test]
fn test_addr_display() {
let (s, _r) = unbounded_channel::<Message>();
let addr = Addr::new(s);
let display = format!("{}", addr);
assert!(display.starts_with("actor:"));
assert_eq!(display.len(), "actor:".len() + 32);
}
#[test]
fn test_addr_noop_sends_silently() {
let addr = Addr::noop();
assert_eq!(addr.id.len(), 32);
}
#[test]
fn test_addr_id_length() {
let (s, _r) = unbounded_channel::<Message>();
let addr = Addr::new(s);
assert_eq!(addr.id.len(), 32);
assert!(
addr.id.chars().all(|c| c.is_ascii_alphanumeric()),
"addr id should be alphanumeric"
);
}
struct TestActor {
received: Arc<RwLock<Vec<Message>>>,
}
#[async_trait]
impl Actor for TestActor {
async fn handle(&mut self, message: Message, _ctx: &ActorContext) {
self.received.write().push(message);
}
}
#[tokio::test]
async fn test_actor_context_new() {
let ctx = ActorContext::new("peer1".to_string());
assert_eq!(*ctx.peer_id.read(), "peer1");
assert_eq!(ctx.child_actor_count(), 0);
assert!(!*ctx.is_stopped.read());
}
#[tokio::test]
async fn test_actor_start_and_send() {
let mut ctx = ActorContext::new("test".to_string());
let received = Arc::new(RwLock::new(Vec::new()));
let actor = TestActor {
received: received.clone(),
};
let _addr = ctx.start_actor(Box::new(actor));
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
assert_eq!(ctx.child_actor_count(), 1);
ctx.stop();
assert!(*ctx.is_stopped.read());
}
}
#[test]
fn test_shutdown_signal_default_false() {
let ctx = ActorContext::new("test-peer".to_string());
assert!(
!*ctx.shutdown_rx.borrow(),
"shutdown_rx should default to false"
);
}
#[test]
fn test_shutdown_signal_propagates_to_child() {
let (tx, rx) = watch::channel(false);
let mut ctx = ActorContext::new("parent".to_string());
ctx.shutdown_rx = rx;
let (stop_tx, _stop_rx) = channel(1);
let child = ctx.child_context(Addr::noop(), stop_tx);
assert!(!*child.shutdown_rx.borrow());
tx.send(true).unwrap();
assert!(
*child.shutdown_rx.borrow(),
"child should see shutdown signal"
);
}
#[test]
fn test_shutdown_signal_isolated_per_node() {
let mut ctx_a = ActorContext::new("node-a".to_string());
let ctx_b = ActorContext::new("node-b".to_string());
assert!(!*ctx_a.shutdown_rx.borrow());
assert!(!*ctx_b.shutdown_rx.borrow());
let (tx_a, rx_a) = watch::channel(false);
ctx_a.shutdown_rx = rx_a;
tx_a.send(true).unwrap();
assert!(*ctx_a.shutdown_rx.borrow());
assert!(
!*ctx_b.shutdown_rx.borrow(),
"unrelated node should not see signal"
);
}