//! RPC helper layer for remote Erlang function calls. //! //! This module provides an abstraction over erl_dist's message passing //! for making remote procedure calls to Erlang nodes. use crate::connection::{CallRef, ConnectionManager}; use crate::error::{RpcError, RpcResult}; use crate::server::FormatterMode; use eetf::{Atom, List, Map, Pid, Term, Tuple}; use erl_dist::message::Message; use std::collections::HashMap; use std::sync::Arc; use std::time::Duration; use tokio::sync::Mutex; use tokio::time::timeout; use crate::connection::NodeConnection; /// Default timeout for RPC calls in milliseconds. const DEFAULT_RPC_TIMEOUT_MS: u64 = 5000; /// Calls `module:function(args)` on the node, through its `rex` server. /// /// `timeout_ms` defaults to [`DEFAULT_RPC_TIMEOUT_MS`]. pub async fn rpc_call( connection_manager: &ConnectionManager, module: &str, function: &str, args: Vec, timeout_ms: Option, ) -> RpcResult { let timeout_duration = Duration::from_millis(timeout_ms.unwrap_or(DEFAULT_RPC_TIMEOUT_MS)); let node = connection_manager.node(); let conn = connection_manager .connection() .await .ok_or_else(|| RpcError::NodeNotConnected { node: node.to_string(), module: module.to_string(), function: function.to_string(), })?; let result = timeout( timeout_duration, execute_rpc( &conn, connection_manager.formatter_mode(), connection_manager.local_node_name(), node, module, function, args, ), ) .await; match result { Ok(inner_result) => inner_result, Err(_) => Err(RpcError::Timeout { node: node.to_string(), module: module.to_string(), function: function.to_string(), timeout_ms: timeout_duration.as_millis() as u64, }), } } /// The pid this server presents to the nodes it talks to. /// /// It addresses no real process: the connection task answers whatever the node /// sends here. The node part must match the name we handshook as, or replies /// are routed to a node that does not exist and every call times out. fn local_pid(local_node_name: &str) -> Pid { Pid::new(local_node_name, 0, 0, 0) } /// Sends one `$gen_call` to `rex` and awaits its reply. async fn execute_rpc( conn: &Arc>, mode: FormatterMode, local_node_name: &str, node: &str, module: &str, function: &str, args: Vec, ) -> RpcResult { let request_sender = conn .lock() .await .request_sender() .map_err(RpcError::Connection)?; // rex, the node's RPC server, expects // {'$gen_call', {From, Tag}, {call, Module, Function, Args, GroupLeader}}. let module_atom = Atom::from(module); let function_atom = Atom::from(function); let args_list = Term::from(List::from(args)); let from_pid = local_pid(local_node_name); let call_tuple = Term::from(Tuple::from(vec![ Term::from(Atom::from("call")), Term::from(module_atom), Term::from(function_atom), args_list, Term::from(Atom::from("user")), // Group leader ])); let call_ref = CallRef::next(); let from_tuple = Term::from(Tuple::from(vec![ Term::from(from_pid.clone()), call_ref.to_term(), ])); let rex_message = Term::from(Tuple::from(vec![ Term::from(Atom::from("$gen_call")), from_tuple, call_tuple, ])); let message = Message::reg_send(from_pid, Atom::from("rex"), rex_message); let result = request_sender .call(message, call_ref) .await .map_err(RpcError::Connection)?; check_badrpc_response(mode, result, node, module, function) } /// Turns a `{badrpc, Reason}` answer into the matching error. fn check_badrpc_response( mode: FormatterMode, term: Term, node: &str, module: &str, function: &str, ) -> RpcResult { if let Term::Tuple(ref tuple) = term { let elements = tuple.elements.as_slice(); if elements.len() == 2 && let Term::Atom(ref atom) = elements[0] && atom.name == "badrpc" { if is_undef(&elements[1]) { return Err(RpcError::Undef { node: node.to_string(), module: module.to_string(), function: function.to_string(), }); } let reason = format_term_for_error(mode, &elements[1]); return Err(RpcError::BadRpc { node: node.to_string(), module: module.to_string(), function: function.to_string(), reason, }); } } Ok(term) } /// Recognises `{'EXIT', {undef, _}}` and `{undef, _}`, which is how a node /// reports being asked for a function it has not loaded. fn is_undef(reason: &Term) -> bool { let Term::Tuple(tuple) = reason else { return false; }; match tuple.elements.as_slice() { [Term::Atom(tag), _] if tag.name == "undef" => true, [Term::Atom(tag), inner] if tag.name == "EXIT" => is_undef(inner), _ => false, } } /// Render a term for an error message, in the server's configured syntax. pub(crate) fn format_term_for_error(mode: FormatterMode, term: &Term) -> String { crate::formatter::get_formatter(mode).format_term(term) } pub(crate) fn atom(name: &str) -> Term { Term::from(Atom::from(name)) } /// Create a tuple term. Only the tests build these; the node sends them. #[cfg(test)] pub(crate) fn tuple(elements: Vec) -> Term { Term::from(Tuple::from(elements)) } /// Create a list term. pub(crate) fn list(elements: Vec) -> Term { Term::from(List::from(elements)) } /// Create a binary term from bytes. Test-only; [`binary_from_str`] covers the /// one case the server itself sends. #[cfg(test)] pub(crate) fn binary(bytes: Vec) -> Term { Term::from(eetf::Binary::from(bytes)) } /// Create a binary term from a string. pub(crate) fn binary_from_str(s: &str) -> Term { Term::from(eetf::Binary::from(s.as_bytes().to_vec())) } /// Create a map term from a vector of key-value pairs. pub(crate) fn map(entries: Vec<(Term, Term)>) -> Term { let map: HashMap = entries.into_iter().collect(); Term::from(Map::from(map)) } pub(crate) fn extract_atom(term: &Term) -> Option<&str> { match term { Term::Atom(a) => Some(&a.name), _ => None, } } /// Extract tuple elements from a term. pub(crate) fn extract_tuple(term: &Term) -> Option<&[Term]> { match term { Term::Tuple(t) => Some(&t.elements), _ => None, } } /// Extract list elements from a term. pub(crate) fn extract_list(term: &Term) -> Option<&[Term]> { match term { Term::List(l) => Some(&l.elements), _ => None, } } /// Extract binary bytes from a term. pub(crate) fn extract_binary(term: &Term) -> Option<&[u8]> { match term { Term::Binary(b) => Some(&b.bytes), _ => None, } } /// Extract map as a reference from a term. pub(crate) fn extract_map(term: &Term) -> Option<&HashMap> { match term { Term::Map(m) => Some(&m.map), _ => None, } } /// Pull `Value` out of an `{ok, Value}` tuple. pub(crate) fn extract_ok_value(term: &Term) -> Option<&Term> { match extract_tuple(term)? { [head, value] if is_atom(head, "ok") => Some(value), _ => None, } } /// Check if a term is a specific atom. pub(crate) fn is_atom(term: &Term, name: &str) -> bool { matches!(term, Term::Atom(a) if a.name == name) } #[cfg(test)] #[allow(clippy::unwrap_used, clippy::expect_used)] mod tests { use super::*; use eetf::FixInteger; #[test] fn atom_helper() { let term = atom("ok"); assert!(matches!(term, Term::Atom(a) if a.name == "ok")); } #[test] fn binary_helper() { let term = binary(vec![1, 2, 3]); if let Term::Binary(b) = term { assert_eq!(b.bytes, vec![1, 2, 3]); } else { panic!("Expected Binary"); } } // ======================================================================== // Extraction helper tests // ======================================================================== #[test] fn extract_atom_success() { let term = Term::from(Atom::from("test")); assert_eq!(extract_atom(&term), Some("test")); } #[test] fn extract_atom_failure() { let term = Term::from(FixInteger::from(42)); assert_eq!(extract_atom(&term), None); } #[test] fn extract_tuple_success() { let term = Term::from(Tuple::from(vec![ Term::from(Atom::from("ok")), Term::from(FixInteger::from(1)), ])); let elements = extract_tuple(&term); assert!(elements.is_some()); assert_eq!(elements.unwrap().len(), 2); } #[test] fn extract_list_success() { let term = Term::from(List::from(vec![Term::from(FixInteger::from(1))])); let elements = extract_list(&term); assert!(elements.is_some()); assert_eq!(elements.unwrap().len(), 1); } #[test] fn extract_binary_success() { let term = Term::from(eetf::Binary::from(b"test".to_vec())); let bytes = extract_binary(&term); assert_eq!(bytes, Some(b"test".as_slice())); } #[test] fn extract_map_success() { let mut map_entries = HashMap::new(); map_entries.insert(Term::from(Atom::from("a")), Term::from(FixInteger::from(1))); let term = Term::from(Map::from(map_entries)); let map = extract_map(&term); assert!(map.is_some()); assert_eq!(map.unwrap().len(), 1); } #[test] fn is_atom_check() { let ok = Term::from(Atom::from("ok")); let error = Term::from(Atom::from("error")); assert!(is_atom(&ok, "ok")); assert!(!is_atom(&ok, "error")); assert!(is_atom(&error, "error")); } #[test] fn extract_ok_value_success() { let ok_tuple = Term::from(Tuple::from(vec![ Term::from(Atom::from("ok")), Term::from(FixInteger::from(42)), ])); let value = extract_ok_value(&ok_tuple); assert!(value.is_some()); assert!(matches!(value.unwrap(), Term::FixInteger(i) if i.value == 42)); } #[test] fn extract_ok_value_failure() { let error_tuple = Term::from(Tuple::from(vec![ Term::from(Atom::from("error")), Term::from(Atom::from("reason")), ])); assert!(extract_ok_value(&error_tuple).is_none()); } // ======================================================================== // badrpc detection tests // ======================================================================== #[test] fn check_badrpc_detects_badrpc() { let badrpc = Term::from(Tuple::from(vec![ Term::from(Atom::from("badrpc")), Term::from(Atom::from("nodedown")), ])); let result = check_badrpc_response(FormatterMode::Erlang, badrpc, "node@host", "mod", "fun"); assert!(result.is_err()); if let Err(RpcError::BadRpc { reason, .. }) = result { assert_eq!(reason, "nodedown"); } else { panic!("Expected BadRpc error"); } } #[test] fn check_badrpc_passes_ok() { let ok = Term::from(Tuple::from(vec![ Term::from(Atom::from("ok")), Term::from(FixInteger::from(42)), ])); let result = check_badrpc_response(FormatterMode::Erlang, ok, "node@host", "mod", "fun"); assert!(result.is_ok()); } #[test] fn check_badrpc_passes_plain_value() { let value = Term::from(FixInteger::from(123)); let result = check_badrpc_response(FormatterMode::Erlang, value, "node@host", "mod", "fun"); assert!(result.is_ok()); } // ======================================================================== // Error formatting tests // ======================================================================== #[test] fn format_term_for_error_atom() { let term = Term::from(Atom::from("test")); assert_eq!(format_term_for_error(FormatterMode::Erlang, &term), "test"); } #[test] fn format_term_for_error_integer() { let term = Term::from(FixInteger::from(42)); assert_eq!(format_term_for_error(FormatterMode::Erlang, &term), "42"); } #[test] fn format_term_for_error_tuple() { let term = Term::from(Tuple::from(vec![ Term::from(Atom::from("error")), Term::from(Atom::from("reason")), ])); assert_eq!( format_term_for_error(FormatterMode::Erlang, &term), "{error, reason}" ); } #[test] fn format_term_for_error_list() { let term = Term::from(List::from(vec![ Term::from(FixInteger::from(1)), Term::from(FixInteger::from(2)), ])); assert_eq!( format_term_for_error(FormatterMode::Erlang, &term), "[1, 2]" ); } #[test] fn format_term_for_error_honours_the_configured_mode() { let term = Term::from(eetf::Binary::from(b"boom".to_vec())); assert_eq!( format_term_for_error(FormatterMode::Erlang, &term), "<<\"boom\">>" ); assert_eq!( format_term_for_error(FormatterMode::Elixir, &term), "\"boom\"" ); } }