//! Post-handshake distribution session: message channel, ticks, RPC. use crate::error::{Error, Result}; use crate::stream::DistStream; use crate::term; use erl_dist::handshake::{ClientSideHandshake, HandshakeStatus, ServerSideHandshake}; use erl_dist::message::{self, Message, Receiver, Sender}; use erl_dist::node::{Creation, LocalNode, NodeName, PeerNode}; use erl_dist::term::{Atom, FixInteger, List, Mfa, Pid, PidOrAtom, Reference, Term}; use erl_dist::{DistributionFlags, HIGHEST_DISTRIBUTION_PROTOCOL_VERSION}; use futures::channel::{mpsc, oneshot}; use futures::{FutureExt, StreamExt}; use std::collections::HashMap; use std::sync::Arc; use std::time::Duration; use tokio::sync::RwLock; use tracing::{debug, warn}; const TICK_INTERVAL: Duration = Duration::from_secs(15); const DEFAULT_RPC_TIMEOUT: Duration = Duration::from_secs(30); /// How we identify a connected peer. #[derive(Debug, Clone, PartialEq, Eq, Hash)] pub enum PeerId { /// Erlang node name (`foo@host`). Node(String), /// iroh endpoint id (base32). Endpoint(String), } impl std::fmt::Display for PeerId { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { Self::Node(n) => write!(f, "{n}"), Self::Endpoint(e) => write!(f, "iroh:{e}"), } } } /// A live distribution link to one peer. pub struct Session { local: LocalNode, peer: PeerNode, peer_id: PeerId, handle: SessionHandle, _runner: tokio::task::JoinHandle<()>, } /// Cloneable handle used to talk to a session's background task. #[derive(Clone)] pub struct SessionHandle { req_tx: mpsc::Sender, local_name: String, peer_name: String, } enum SessionReq { Rpc { module: Atom, function: Atom, args: List, reply: oneshot::Sender>, }, Send { to: Atom, msg: Term, reply: oneshot::Sender>, }, Raw { msg: Message, reply: oneshot::Sender>, }, } impl Session { /// Client-side: handshake as a connecting node over an established stream. pub async fn connect( stream: DistStream, local_name: &str, cookie: &str, peer_id: PeerId, ) -> Result { let local_node_name: NodeName = local_name .parse() .map_err(|e: erl_dist::node::NodeNameError| Error::InvalidName(e.to_string()))?; let mut local = LocalNode::new(local_node_name, Creation::random()); local.flags |= DistributionFlags::NAME_ME; local.flags |= DistributionFlags::SPAWN; local.flags |= DistributionFlags::DIST_MONITOR; local.flags |= DistributionFlags::DIST_MONITOR_NAME; local.flags |= DistributionFlags::EXPORT_PTR_TAG; local.flags |= DistributionFlags::BIT_BINARIES; local.flags |= DistributionFlags::NEW_FLOATS; local.flags |= DistributionFlags::FUN_TAGS; local.flags |= DistributionFlags::NEW_FUN_TAGS; local.flags |= DistributionFlags::EXTENDED_PIDS_PORTS; local.flags |= DistributionFlags::UTF8_ATOMS; local.flags |= DistributionFlags::MAP_TAGS; local.flags |= DistributionFlags::BIG_CREATION; local.flags |= DistributionFlags::HANDSHAKE_23; local.flags |= DistributionFlags::UNLINK_ID; local.flags |= DistributionFlags::V4_NC; let mut handshake = ClientSideHandshake::new(stream, local.clone(), cookie); let status = handshake .execute_send_name(HIGHEST_DISTRIBUTION_PROTOCOL_VERSION) .await .map_err(Error::handshake)?; match status { HandshakeStatus::Ok | HandshakeStatus::OkSimultaneous => {} HandshakeStatus::Named { name, creation } => { local.name = NodeName::new(&name, local.name.host()) .map_err(|e| Error::InvalidName(e.to_string()))?; local.creation = creation; } HandshakeStatus::Alive => { // continue } other => { return Err(Error::handshake(format!("unexpected status: {other:?}"))); } } let (connection, peer) = handshake .execute_rest(true) .await .map_err(Error::handshake)?; debug!(peer = %peer.name, "handshake complete (client)"); Self::from_connection(connection, local, peer, peer_id) } /// Server-side: accept a handshake on an inbound stream. pub async fn accept( stream: DistStream, local_name: &str, cookie: &str, peer_id: PeerId, ) -> Result { let local_node_name: NodeName = local_name .parse() .map_err(|e: erl_dist::node::NodeNameError| Error::InvalidName(e.to_string()))?; let mut local = LocalNode::new(local_node_name, Creation::random()); local.flags |= DistributionFlags::SPAWN; local.flags |= DistributionFlags::DIST_MONITOR; local.flags |= DistributionFlags::DIST_MONITOR_NAME; local.flags |= DistributionFlags::HANDSHAKE_23; local.flags |= DistributionFlags::UNLINK_ID; local.flags |= DistributionFlags::V4_NC; local.flags |= DistributionFlags::UTF8_ATOMS; local.flags |= DistributionFlags::MAP_TAGS; local.flags |= DistributionFlags::BIG_CREATION; let mut handshake = ServerSideHandshake::new(stream, local.clone(), cookie); let peer_name = handshake .execute_recv_name() .await .map_err(Error::handshake)?; let status = if peer_name.is_none() { // Peer asked for a dynamic name — hand one out. use erl_dist::node::Creation; HandshakeStatus::Named { name: format!("dyn-{}", rand::random::()), creation: Creation::random(), } } else { HandshakeStatus::Ok }; let (connection, peer) = handshake .execute_rest(status) .await .map_err(Error::handshake)?; debug!(peer = %peer.name, "handshake complete (server)"); Self::from_connection(connection, local, peer, peer_id) } fn from_connection( connection: DistStream, local: LocalNode, peer: PeerNode, peer_id: PeerId, ) -> Result { let flags = local.flags & peer.flags; let (msg_tx, msg_rx) = message::channel(connection, flags); let (req_tx, req_rx) = mpsc::channel(256); let handle = SessionHandle { req_tx, local_name: local.name.to_string(), peer_name: peer.name.to_string(), }; let runner_state = Runner { msg_tx, msg_rx: Some(msg_rx), req_rx: Some(req_rx), local: local.clone(), ongoing: HashMap::new(), }; let _runner = tokio::spawn(async move { if let Err(e) = runner_state.run().await { warn!(error = %e, "session runner stopped"); } }); Ok(Self { local, peer, peer_id, handle, _runner, }) } /// Cloneable handle for RPC / send. pub fn handle(&self) -> SessionHandle { self.handle.clone() } /// Local node info after handshake (may have been renamed via NAME_ME). pub fn local_node(&self) -> &LocalNode { &self.local } /// Peer node info. pub fn peer_node(&self) -> &PeerNode { &self.peer } /// Peer identifier used by the bot registry. pub fn peer_id(&self) -> &PeerId { &self.peer_id } } impl SessionHandle { /// Peer Erlang node name. pub fn peer_name(&self) -> &str { &self.peer_name } /// Local Erlang node name. pub fn local_name(&self) -> &str { &self.local_name } /// Remote procedure call: `module:function(args...)`. pub async fn rpc( &self, module: impl Into, function: impl Into, args: List, ) -> Result { self.rpc_timeout(module, function, args, DEFAULT_RPC_TIMEOUT) .await } /// RPC with an explicit timeout. pub async fn rpc_timeout( &self, module: impl Into, function: impl Into, args: List, timeout: Duration, ) -> Result { let (reply, rx) = oneshot::channel(); self.req_tx .clone() .try_send(SessionReq::Rpc { module: module.into(), function: function.into(), args, reply, }) .map_err(|_| Error::Terminated)?; match tokio::time::timeout(timeout, rx).await { Ok(Ok(r)) => r, Ok(Err(_)) => Err(Error::Terminated), Err(_) => Err(Error::Timeout), } } /// Send a message to a registered process name on the peer. pub async fn send(&self, to: impl Into, msg: Term) -> Result<()> { let (reply, rx) = oneshot::channel(); self.req_tx .clone() .try_send(SessionReq::Send { to: to.into(), msg, reply, }) .map_err(|_| Error::Terminated)?; rx.await.map_err(|_| Error::Terminated)? } /// Send a raw distribution message. pub async fn send_raw(&self, msg: Message) -> Result<()> { let (reply, rx) = oneshot::channel(); self.req_tx .clone() .try_send(SessionReq::Raw { msg, reply }) .map_err(|_| Error::Terminated)?; rx.await.map_err(|_| Error::Terminated)? } } struct Runner { msg_tx: Sender, msg_rx: Option>, req_rx: Option>, local: LocalNode, ongoing: HashMap>>, } impl Runner { async fn run(mut self) -> Result<()> { let mut req_rx = self.req_rx.take().expect("req_rx").into_future(); let msg_rx = self.msg_rx.take().expect("msg_rx"); let mut msg_fut = msg_rx.recv_owned().boxed(); let mut tick = tokio::time::interval(TICK_INTERVAL); tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay); loop { tokio::select! { _ = tick.tick() => { if let Err(e) = self.msg_tx.send(Message::Tick).await { return Err(Error::send(e)); } } result = &mut msg_fut => { let (msg, next_rx) = result.map_err(Error::recv)?; self.handle_msg(msg).await?; msg_fut = next_rx.recv_owned().boxed(); } (req, rest) = &mut req_rx => { match req { Some(r) => { self.handle_req(r).await?; req_rx = rest.into_future(); } None => { debug!("session request channel closed"); break; } } } } } Ok(()) } async fn handle_req(&mut self, req: SessionReq) -> Result<()> { match req { SessionReq::Rpc { module, function, args, reply, } => { let req_id = self.make_ref(); let spawn_request = Message::spawn_request( req_id.clone(), self.pid(), self.pid(), Mfa { module: "erpc".into(), function: "execute_call".into(), arity: FixInteger::from(4), }, List::from(vec![Atom::from("monitor").into()]), List::from(vec![ self.make_ref().into(), module.into(), function.into(), args.into(), ]), ); if let Err(e) = self.msg_tx.send(spawn_request).await { let _ = reply.send(Err(Error::send(e))); } else { self.ongoing.insert(req_id, reply); } } SessionReq::Send { to, msg, reply } => { let m = Message::reg_send(self.pid(), to, msg); let r = self.msg_tx.send(m).await.map_err(Error::send); let _ = reply.send(r); } SessionReq::Raw { msg, reply } => { let r = self.msg_tx.send(msg).await.map_err(Error::send); let _ = reply.send(r); } } Ok(()) } async fn handle_msg(&mut self, msg: Message) -> Result<()> { match msg { Message::Tick => Ok(()), Message::SpawnReply(msg) => { if let PidOrAtom::Atom(reason) = msg.result { // Find and fail the matching request if still pending. // The monitor exit will also fire; if spawn failed, clean up. if let Some((_, reply)) = self .ongoing .iter() .find(|(k, _)| k.id == msg.req_id.id) .map(|(k, _)| k.clone()) .and_then(|k| self.ongoing.remove(&k).map(|r| (k, r))) { let _ = reply.send(Err(Error::rpc(format!( "spawn_request failed: {}", reason.name )))); } } Ok(()) } Message::MonitorPExit(msg) => { if let Some(reply) = self.ongoing.remove(&msg.reference) { let _ = reply.send(decode_erpc_result(msg.reason)); } Ok(()) } other => { debug!(?other, "ignored distribution message"); Ok(()) } } } fn node(&self) -> Atom { Atom::from(self.local.name.to_string()) } fn pid(&self) -> Pid { Pid::new(self.node(), 0, 0, self.local.creation.get()) } fn make_ref(&self) -> Reference { term::make_ref(&self.local.name.to_string(), self.local.creation.get()) } } /// Decode `{Ref, return, Value}` / error shapes from `erpc:execute_call`. fn decode_erpc_result(reason: Term) -> Result { // Success shape: {CallerRef, return, Value} OR just the exit reason from monitor // erpc:execute_call exits with {Ref, return, Result} on success. if let Term::Tuple(tup) = &reason { if tup.elements.len() == 3 { if let Term::Atom(kind) = &tup.elements[1] { if kind.name == "return" { return Ok(tup.elements[2].clone()); } if kind.name == "throw" || kind.name == "error" || kind.name == "exit" { return Err(Error::rpc(format!("{reason}"))); } } } } // Some paths just return the value as the exit reason directly. Ok(reason) } /// Registry of active sessions, shared across the bot. #[derive(Clone, Default)] pub struct SessionRegistry { inner: Arc>>, } impl SessionRegistry { /// Insert a session under its peer name and peer_id display key. pub async fn insert(&self, session: &Session) { let mut g = self.inner.write().await; let handle = session.handle(); g.insert(session.peer_node().name.to_string(), handle.clone()); g.insert(session.peer_id().to_string(), handle); } /// Remove by any known key. pub async fn remove(&self, key: &str) { let mut g = self.inner.write().await; g.remove(key); } /// Lookup handle. pub async fn get(&self, key: &str) -> Option { self.inner.read().await.get(key).cloned() } /// All known keys. pub async fn list(&self) -> Vec { self.inner.read().await.keys().cloned().collect() } }