use anyhow::Result;
use std::num::NonZero;
use std::sync::Arc;
use crate::events::{DistributedEventFactory, EventHandle};
use crate::observability::VeloMetrics;
use crate::transports::{Transport, VeloBackend};
use velo_ext::{InstanceId, PeerInfo};
use crate::PeerDiscovery;
use crate::messenger::VeloEvents;
use crate::messenger::client::ActiveMessageClient;
use crate::messenger::client::builders::MessageBuilder;
use crate::messenger::handlers::{Handler, HandlerManager};
use crate::messenger::server::ActiveMessageServer;
#[derive(Clone)]
pub struct Messenger {
instance_id: InstanceId,
backend: Arc<VeloBackend>,
client: Arc<ActiveMessageClient>,
server: Arc<ActiveMessageServer>,
handlers: HandlerManager,
discovery: Option<Arc<dyn PeerDiscovery>>,
events: Arc<VeloEvents>,
observability: Option<Arc<VeloMetrics>>,
runtime: tokio::runtime::Handle,
tracker: tokio_util::task::TaskTracker,
large_payload_resolver:
Arc<std::sync::OnceLock<Arc<dyn crate::messenger::large_payload::LargePayloadResolver>>>,
}
pub struct MessengerBuilder {
transports: Vec<Arc<dyn Transport>>,
discovery: Option<Arc<dyn PeerDiscovery>>,
metrics: Option<Arc<VeloMetrics>>,
}
impl MessengerBuilder {
pub fn new() -> Self {
Self {
transports: Vec::new(),
discovery: None,
metrics: None,
}
}
pub fn add_transport(mut self, transport: Arc<dyn Transport>) -> Self {
self.transports.push(transport);
self
}
pub fn discovery(mut self, discovery: Arc<dyn PeerDiscovery>) -> Self {
self.discovery = Some(discovery);
self
}
pub fn metrics(mut self, metrics: Arc<VeloMetrics>) -> Self {
self.metrics = Some(metrics);
self
}
pub async fn build(self) -> Result<Arc<Messenger>> {
Messenger::new(self.transports, self.discovery, self.metrics).await
}
}
impl Default for MessengerBuilder {
fn default() -> Self {
Self::new()
}
}
impl Messenger {
pub fn builder() -> MessengerBuilder {
MessengerBuilder::new()
}
pub(crate) async fn new(
transports: Vec<Arc<dyn Transport>>,
discovery: Option<Arc<dyn PeerDiscovery>>,
metrics: Option<Arc<VeloMetrics>>,
) -> Result<Arc<Self>> {
let (backend, data_streams) = VeloBackend::new(transports, metrics.clone()).await?;
let backend = Arc::new(backend);
let instance_id = backend.instance_id();
let worker_id = instance_id.worker_id();
let response_manager =
crate::messenger::common::responses::ResponseManager::with_observability(
worker_id.as_u64(),
metrics.clone(),
);
let runtime = tokio::runtime::Handle::current();
let tracker = tokio_util::task::TaskTracker::new();
let system_id = NonZero::new(worker_id.as_u64())
.expect("worker_id must be non-zero for distributed events");
let factory = DistributedEventFactory::new(system_id);
let local_base = factory.system().clone();
let events = VeloEvents::new(
instance_id,
local_base,
backend.clone(),
response_manager.clone(),
);
let large_payload_resolver: Arc<
std::sync::OnceLock<Arc<dyn crate::messenger::large_payload::LargePayloadResolver>>,
> = Arc::new(std::sync::OnceLock::new());
let server = ActiveMessageServer::new(
response_manager.clone(),
None,
data_streams,
backend.clone(),
tracker.clone(),
metrics.clone(),
large_payload_resolver.clone(),
)
.await;
let server = Arc::new(server);
struct DefaultErrorHandler {
response_manager: crate::messenger::common::responses::ResponseManager,
}
impl crate::transports::TransportErrorHandler for DefaultErrorHandler {
fn on_error(&self, header: bytes::Bytes, _payload: bytes::Bytes, error: String) {
if let Some(response_id) =
crate::messenger::common::messages::decode_response_id_from_request_header(
&header,
)
{
self.response_manager
.complete_outcome(response_id, Err(error.clone()));
}
tracing::error!("Transport error: {}", error);
}
}
let client = Arc::new(ActiveMessageClient::new(
response_manager.clone(),
backend.clone(),
Arc::new(DefaultErrorHandler {
response_manager: response_manager.clone(),
}),
discovery.clone(),
metrics.clone(),
));
let handlers = HandlerManager::new(server.hub().handlers_arc());
let system = Arc::new(Self {
instance_id,
backend: backend.clone(),
client,
server: server.clone(),
handlers,
discovery,
events: events.clone(),
observability: metrics,
runtime,
tracker,
large_payload_resolver,
});
events.set_messenger(system.clone());
crate::messenger::events::handlers::register_event_handlers(&system.handlers, events)?;
crate::messenger::server::register_system_handlers(&system.handlers)?;
server.hub().set_system(system.clone())?;
Ok(system)
}
pub fn instance_id(&self) -> InstanceId {
self.instance_id
}
pub fn peer_info(&self) -> PeerInfo {
self.backend.peer_info()
}
pub(crate) fn backend(&self) -> &Arc<VeloBackend> {
&self.backend
}
pub(crate) fn discovery(&self) -> Option<Arc<dyn PeerDiscovery>> {
self.discovery.clone()
}
pub(crate) fn observability(&self) -> Option<Arc<VeloMetrics>> {
self.observability.clone()
}
pub fn events(&self) -> &Arc<VeloEvents> {
&self.events
}
pub fn event_manager(&self) -> crate::events::EventManager {
self.events.event_manager()
}
pub fn am_send(
&self,
handler: &str,
) -> Result<crate::messenger::client::builders::AmSendBuilder> {
crate::messenger::client::builders::AmSendBuilder::new(self.client.clone(), handler)
}
pub fn am_send_streaming(
&self,
handler: &str,
) -> Result<crate::messenger::client::builders::AmSendBuilder> {
Ok(
crate::messenger::client::builders::AmSendBuilder::new_unchecked(
self.client.clone(),
handler,
),
)
}
pub fn am_sync(
&self,
handler: &str,
) -> Result<crate::messenger::client::builders::AmSyncBuilder> {
crate::messenger::client::builders::AmSyncBuilder::new(self.client.clone(), handler)
}
pub fn unary(&self, handler: &str) -> Result<crate::messenger::client::builders::UnaryBuilder> {
crate::messenger::client::builders::UnaryBuilder::new(self.client.clone(), handler)
}
pub fn typed_unary<R: serde::de::DeserializeOwned + Send + 'static>(
&self,
handler: &str,
) -> Result<crate::messenger::client::builders::TypedUnaryBuilder<R>> {
crate::messenger::client::builders::TypedUnaryBuilder::new(self.client.clone(), handler)
}
pub fn unary_streaming(
&self,
handler: &str,
) -> crate::messenger::client::builders::UnaryBuilder {
crate::messenger::client::builders::UnaryBuilder::new_unchecked(
self.client.clone(),
handler,
)
}
pub fn typed_unary_streaming<R: serde::de::DeserializeOwned + Send + 'static>(
&self,
handler: &str,
) -> crate::messenger::client::builders::TypedUnaryBuilder<R> {
crate::messenger::client::builders::TypedUnaryBuilder::new_unchecked(
self.client.clone(),
handler,
)
}
pub fn register_handler(&self, handler: Handler) -> Result<()> {
self.handlers.register_handler(handler)
}
pub fn register_streaming_handler(
&self,
handler: crate::messenger::handlers::Handler,
) -> anyhow::Result<()> {
self.handlers.register_internal_handler(handler)
}
pub fn set_large_payload_support(
&self,
stager: Arc<dyn crate::messenger::large_payload::LargePayloadStager>,
resolver: Arc<dyn crate::messenger::large_payload::LargePayloadResolver>,
) {
let _ = self.client.large_payload_stager.set(stager);
let _ = self.large_payload_resolver.set(resolver);
}
pub fn register_peer(&self, peer_info: PeerInfo) -> Result<()> {
let instance_id = peer_info.instance_id();
self.backend.register_peer(peer_info)?;
self.client.register_peer(instance_id);
Ok(())
}
pub async fn discover_and_register_peer(&self, instance_id: InstanceId) -> Result<()> {
tracing::debug!(
target: "crate::messenger::discovery",
%instance_id,
"Discovering peer by instance_id"
);
let discovery = self.discovery.as_ref().ok_or_else(|| {
anyhow::anyhow!(
"No discovery backend configured. Cannot discover instance {}",
instance_id
)
})?;
let peer_info = discovery.discover_by_instance_id(instance_id).await?;
tracing::info!(
target: "crate::messenger::discovery",
%instance_id,
"Discovered peer, registering"
);
self.register_peer(peer_info)
}
pub fn has_event_subscriber(&self, handle: EventHandle, subscriber: InstanceId) -> bool {
self.events.has_subscriber(handle, subscriber)
}
pub async fn available_handlers(&self, instance_id: InstanceId) -> Result<Vec<String>> {
self.client.get_peer_handlers(instance_id).await
}
pub async fn refresh_handlers(&self, instance_id: InstanceId) -> Result<()> {
self.client.refresh_handler_list(instance_id).await
}
pub async fn wait_for_handler(
&self,
instance_id: InstanceId,
handler_name: &str,
) -> Result<()> {
const MAX_ATTEMPTS: u32 = 10;
const DELAY: std::time::Duration = std::time::Duration::from_millis(100);
for _ in 0..MAX_ATTEMPTS {
self.refresh_handlers(instance_id).await?;
let handlers = self.available_handlers(instance_id).await?;
if handlers.contains(&handler_name.to_string()) {
return Ok(());
}
tokio::time::sleep(DELAY).await;
}
anyhow::bail!(
"Timeout waiting for handler '{}' on instance {}",
handler_name,
instance_id
)
}
pub fn list_local_handlers(&self) -> Vec<String> {
self.server.hub().list_handlers()
}
pub fn runtime(&self) -> &tokio::runtime::Handle {
&self.runtime
}
pub fn tracker(&self) -> &tokio_util::task::TaskTracker {
&self.tracker
}
pub(crate) fn message_builder_unchecked(&self, handler: &str) -> MessageBuilder {
MessageBuilder::new_unchecked(self.client.clone(), handler)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::messenger::handlers::Handler;
use crate::transports::{
HealthCheckError, MessageType, SendBackpressure, Transport, TransportAdapter,
TransportError, TransportErrorHandler,
};
use bytes::Bytes;
use futures::future::BoxFuture;
use std::collections::HashMap;
use std::sync::{Arc, Mutex, OnceLock};
use std::time::Duration;
use velo_ext::{PeerInfo, TransportKey, WorkerAddress};
static TEST_TRANSPORT_REGISTRY: OnceLock<Mutex<HashMap<String, TransportAdapter>>> =
OnceLock::new();
fn test_transport_registry() -> &'static Mutex<HashMap<String, TransportAdapter>> {
TEST_TRANSPORT_REGISTRY.get_or_init(|| Mutex::new(HashMap::new()))
}
fn make_test_address(key: &str, endpoint: &str) -> WorkerAddress {
let mut entries = HashMap::<String, Vec<u8>>::new();
entries.insert(key.to_string(), endpoint.as_bytes().to_vec());
WorkerAddress::from_encoded(rmp_serde::to_vec(&entries).unwrap())
}
struct InMemoryTransport {
key: TransportKey,
endpoint: String,
address: WorkerAddress,
peers: Mutex<HashMap<InstanceId, String>>,
}
impl InMemoryTransport {
fn new(endpoint: String) -> Arc<Self> {
let key = TransportKey::from("test");
Arc::new(Self {
key: key.clone(),
address: make_test_address(key.as_str(), &endpoint),
endpoint,
peers: Mutex::new(HashMap::new()),
})
}
}
impl Transport for InMemoryTransport {
fn key(&self) -> TransportKey {
self.key.clone()
}
fn address(&self) -> WorkerAddress {
self.address.clone()
}
fn register(&self, peer_info: PeerInfo) -> Result<(), TransportError> {
let endpoint = peer_info
.worker_address()
.get_entry(self.key.as_str())
.map_err(|_| TransportError::InvalidEndpoint)?
.ok_or(TransportError::NoEndpoint)?;
let endpoint = String::from_utf8(endpoint.to_vec())
.map_err(|_| TransportError::InvalidEndpoint)?;
self.peers
.lock()
.expect("peer map poisoned")
.insert(peer_info.instance_id(), endpoint);
Ok(())
}
fn send_message(
&self,
instance_id: InstanceId,
header: Bytes,
payload: Bytes,
message_type: MessageType,
on_error: Arc<dyn TransportErrorHandler>,
) -> Result<(), SendBackpressure> {
let target_endpoint = match self
.peers
.lock()
.expect("peer map poisoned")
.get(&instance_id)
.cloned()
{
Some(endpoint) => endpoint,
None => {
on_error.on_error(header, payload, "Peer not registered".to_string());
return Ok(());
}
};
let maybe_adapter = test_transport_registry()
.lock()
.expect("transport registry poisoned")
.get(&target_endpoint)
.cloned();
let Some(adapter) = maybe_adapter else {
on_error.on_error(header, payload, "Target transport not started".to_string());
return Ok(());
};
let send_result = match message_type {
MessageType::Message => adapter.message_stream.send((header, payload)),
MessageType::Response | MessageType::ShuttingDown => {
adapter.response_stream.send((header, payload))
}
MessageType::Ack | MessageType::Event => {
adapter.event_stream.send((header, payload))
}
};
if let Err(err) = send_result {
let (header, payload) = err.0;
on_error.on_error(header, payload, "Target receive channel closed".to_string());
}
Ok(())
}
fn start(
&self,
_instance_id: InstanceId,
channels: TransportAdapter,
_rt: tokio::runtime::Handle,
) -> BoxFuture<'_, anyhow::Result<()>> {
let endpoint = self.endpoint.clone();
Box::pin(async move {
test_transport_registry()
.lock()
.expect("transport registry poisoned")
.insert(endpoint, channels);
Ok(())
})
}
fn shutdown(&self) {
test_transport_registry()
.lock()
.expect("transport registry poisoned")
.remove(&self.endpoint);
}
fn check_health(
&self,
instance_id: InstanceId,
_timeout: Duration,
) -> std::pin::Pin<
Box<dyn std::future::Future<Output = Result<(), HealthCheckError>> + Send + '_>,
> {
Box::pin(async move {
if self
.peers
.lock()
.expect("peer map poisoned")
.contains_key(&instance_id)
{
Ok(())
} else {
Err(HealthCheckError::PeerNotRegistered)
}
})
}
}
fn make_transport_pair() -> (Arc<dyn Transport>, Arc<dyn Transport>) {
let a = InMemoryTransport::new(format!("in-memory://{}", InstanceId::new_v4()));
let b = InMemoryTransport::new(format!("in-memory://{}", InstanceId::new_v4()));
(a as Arc<dyn Transport>, b as Arc<dyn Transport>)
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn test_register_streaming_handler_allows_underscore() {
let messenger = Messenger::builder().build().await.unwrap();
let handler = Handler::am_handler("_anchor_test", |_ctx| Ok(())).build();
let result = messenger.register_streaming_handler(handler);
assert!(
result.is_ok(),
"register_streaming_handler should allow underscore-prefixed names"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn test_register_streaming_handler_allows_normal() {
let messenger = Messenger::builder().build().await.unwrap();
let handler = Handler::am_handler("normal_test", |_ctx| Ok(())).build();
let result = messenger.register_streaming_handler(handler);
assert!(
result.is_ok(),
"register_streaming_handler should allow normal handler names"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn test_am_send_streaming_allows_underscore_prefix() {
let messenger = Messenger::builder().build().await.unwrap();
let result = messenger.am_send_streaming("_stream_data");
assert!(
result.is_ok(),
"am_send_streaming should allow underscore-prefixed handler names"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn test_am_send_still_rejects_underscore_prefix() {
let messenger = Messenger::builder().build().await.unwrap();
let result = messenger.am_send("_stream_data");
assert!(
result.is_err(),
"am_send should still reject underscore-prefixed handler names"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn test_am_send_streaming_builder_has_setters() {
let messenger = Messenger::builder().build().await.unwrap();
let builder = messenger.am_send_streaming("_stream_data").unwrap();
let _builder = builder
.raw_payload(bytes::Bytes::from_static(b"test"))
.worker(velo_ext::WorkerId::from_u64(1));
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn test_typed_unary_streaming_allows_underscore_prefix() {
let messenger = Messenger::builder().build().await.unwrap();
let _builder = messenger.typed_unary_streaming::<String>("_anchor_attach");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn test_typed_unary_still_rejects_underscore_prefix() {
let messenger = Messenger::builder().build().await.unwrap();
let result = messenger.typed_unary::<String>("_anchor_attach");
assert!(
result.is_err(),
"typed_unary should still reject underscore-prefixed handler names"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn test_system_and_event_handlers_are_available_immediately_after_startup() {
test_transport_registry()
.lock()
.expect("transport registry poisoned")
.clear();
let (transport_a, transport_b) = make_transport_pair();
let a = Messenger::builder()
.add_transport(transport_a)
.build()
.await
.unwrap();
let b = Messenger::builder()
.add_transport(transport_b)
.build()
.await
.unwrap();
a.register_peer(b.peer_info()).unwrap();
b.register_peer(a.peer_info()).unwrap();
let handlers = a.available_handlers(b.instance_id()).await.unwrap();
assert!(
handlers.iter().any(|handler| handler == "_hello"),
"expected _hello to be available immediately after startup"
);
assert!(
handlers.iter().any(|handler| handler == "_list_handlers"),
"expected _list_handlers to be available immediately after startup"
);
assert!(
handlers.iter().any(|handler| handler == "_event_subscribe"),
"expected _event_subscribe to be available immediately after startup"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn test_await_capacity_chains_through_all_public_builders() {
let messenger = Messenger::builder().build().await.unwrap();
let _ = messenger
.am_send("fan_out")
.unwrap()
.await_capacity()
.raw_payload(bytes::Bytes::from_static(b"x"));
let _ = messenger
.am_sync("fan_out")
.unwrap()
.await_capacity()
.raw_payload(bytes::Bytes::from_static(b"x"));
let _ = messenger
.unary("fan_out")
.unwrap()
.await_capacity()
.raw_payload(bytes::Bytes::from_static(b"x"));
let _ = messenger
.typed_unary::<String>("fan_out")
.unwrap()
.await_capacity()
.raw_payload(bytes::Bytes::from_static(b"x"));
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn test_await_capacity_misuse_fails_without_waiting_for_slot() {
let messenger = Messenger::builder().build().await.unwrap();
let res = tokio::time::timeout(
std::time::Duration::from_millis(50),
messenger
.am_sync("fan_out")
.unwrap()
.await_capacity()
.send(),
)
.await
.expect("missing target must resolve immediately");
assert!(
res.is_err(),
"missing target should produce an immediate error"
);
let res = tokio::time::timeout(
std::time::Duration::from_millis(50),
messenger
.unary("fan_out")
.unwrap()
.await_capacity()
.instance(velo_ext::InstanceId::new_v4())
.worker(velo_ext::WorkerId::from_u64(1))
.send(),
)
.await
.expect("mutually-exclusive target must resolve immediately");
assert!(
res.is_err(),
"mutually-exclusive target should produce an immediate error"
);
}
}