use crate::observability::{DispatchFailure, OrderedMetricsHandle};
use crate::transports::VeloBackend;
use bytes::Bytes;
use dashmap::DashMap;
use futures::FutureExt;
use futures::future::BoxFuture;
use std::panic::AssertUnwindSafe;
use std::sync::{Arc, OnceLock};
use std::time::Instant;
use tokio::sync::Semaphore;
use tokio_util::task::TaskTracker;
use tracing::{error, trace, warn};
use velo_ext::WorkerId;
use crate::Messenger;
use crate::messenger::common::events::{EventType, Outcome, encode_event_header};
use crate::messenger::common::messages::ResponseType;
use crate::messenger::common::responses::ResponseId;
use crate::messenger::handlers::{OrderedConfig, OrderingKey, OverflowPolicy};
use crate::messenger::server::lanes::{LaneObserver, LaneRouter, LaneRouterConfig};
#[derive(Clone)]
pub(crate) struct HandlerContext {
pub message_id: ResponseId,
pub payload: Bytes,
pub response_type: ResponseType,
pub headers: Option<std::collections::HashMap<String, String>>,
pub system: Arc<Messenger>,
pub in_flight: Option<Arc<velo_ext::InFlightGuard>>,
}
pub(crate) trait ActiveMessageHandler: Send + Sync {
fn handle(
&self,
ctx: HandlerContext,
) -> std::pin::Pin<Box<dyn std::future::Future<Output = ()> + Send + 'static>>;
fn name(&self) -> &str;
}
pub(crate) trait ActiveMessageDispatcher: Send + Sync {
fn name(&self) -> &str;
fn dispatch(&self, ctx: HandlerContext);
}
pub(crate) struct SpawnedDispatcher<H: ActiveMessageHandler> {
handler: Arc<H>,
task_tracker: TaskTracker,
}
impl<H: ActiveMessageHandler> SpawnedDispatcher<H> {
pub fn new(handler: H, task_tracker: TaskTracker) -> Self {
Self {
handler: Arc::new(handler),
task_tracker,
}
}
}
impl<H: ActiveMessageHandler + 'static> ActiveMessageDispatcher for SpawnedDispatcher<H> {
fn name(&self) -> &str {
self.handler.name()
}
fn dispatch(&self, ctx: HandlerContext) {
let handler = self.handler.clone();
let handler_name = handler.name().to_string();
self.task_tracker.spawn(async move {
trace!(target: "crate::messenger::dispatcher", handler = %handler_name, "Handler task started");
handler.handle(ctx).await;
trace!(target: "crate::messenger::dispatcher", handler = %handler_name, "Handler task completed");
});
}
}
pub(crate) struct InlineDispatcher<H: ActiveMessageHandler> {
handler: Arc<H>,
}
impl<H: ActiveMessageHandler> InlineDispatcher<H> {
pub fn new(handler: H) -> Self {
Self {
handler: Arc::new(handler),
}
}
}
impl<H: ActiveMessageHandler + 'static> ActiveMessageDispatcher for InlineDispatcher<H> {
fn name(&self) -> &str {
self.handler.name()
}
fn dispatch(&self, ctx: HandlerContext) {
let handler = self.handler.clone();
tokio::spawn(async move {
handler.handle(ctx).await;
});
}
}
pub(crate) struct OrderedDispatcher<H: ActiveMessageHandler> {
handler: Arc<H>,
config: OrderedConfig,
bound: OnceLock<BoundRouter>,
limiter: Option<Arc<Semaphore>>,
rendezvous_warning: std::sync::Once,
shed_warning: std::sync::Once,
}
struct BoundRouter {
router: Arc<LaneRouter<LaneKey, LaneItem>>,
metrics: Option<OrderedMetricsHandle>,
}
type LaneKey = Option<WorkerId>;
struct LaneItem {
ctx: HandlerContext,
enqueued_at: Instant,
}
struct OrderedLaneObserver {
metrics: OrderedMetricsHandle,
}
impl LaneObserver for OrderedLaneObserver {
fn lane_created(&self) {
self.metrics.lane_created();
}
fn lane_closed(&self) {
self.metrics.lane_closed();
}
}
impl<H: ActiveMessageHandler + 'static> OrderedDispatcher<H> {
pub fn new(handler: H, config: OrderedConfig) -> Self {
let limiter = config
.max_concurrent
.map(|limit| Arc::new(Semaphore::new(limit)));
Self {
handler: Arc::new(handler),
config,
bound: OnceLock::new(),
limiter,
rendezvous_warning: std::sync::Once::new(),
shed_warning: std::sync::Once::new(),
}
}
fn lane_key(&self, ctx: &HandlerContext) -> LaneKey {
match self.config.key {
OrderingKey::Global => None,
OrderingKey::Sender => Some(WorkerId::from_u64(ctx.message_id.worker_id())),
}
}
fn bind(&self, system: &Arc<Messenger>) -> BoundRouter {
let handler = self.handler.clone();
let handler_name: Arc<str> = Arc::from(self.handler.name());
let limiter = self.limiter.clone();
let metrics = system
.observability()
.as_ref()
.and_then(|m| m.bind_ordered_dispatcher(&handler_name));
let consumer_metrics = metrics.clone();
let consumer = Arc::new(move |item: LaneItem| {
let handler = handler.clone();
let handler_name = Arc::clone(&handler_name);
let limiter = limiter.clone();
let metrics = consumer_metrics.clone();
Box::pin(async move {
if let Some(metrics) = metrics.as_ref() {
metrics.observe_wait(item.enqueued_at.elapsed());
}
let _permit = match limiter.as_ref() {
Some(sem) => sem.clone().acquire_owned().await.ok(),
None => None,
};
let ctx = item.ctx;
let message_id = ctx.message_id;
let response_type = ctx.response_type;
let system = ctx.system.clone();
let outcome = AssertUnwindSafe(async move { handler.handle(ctx).await })
.catch_unwind()
.await;
if outcome.is_err() {
if let Some(metrics) = system.observability().as_ref() {
metrics.record_dispatch_failure(DispatchFailure::OrderedHandlerPanic);
}
error!(
target: "crate::messenger::dispatcher",
handler = %handler_name,
message_id = %message_id,
"Ordered handler panicked; lane preserved"
);
Self::fail_fast(&system, message_id, response_type, "handler panicked");
}
if let Some(metrics) = metrics.as_ref() {
metrics.dequeued();
}
}) as BoxFuture<'static, ()>
});
let observer = metrics
.clone()
.map(|metrics| Arc::new(OrderedLaneObserver { metrics }) as Arc<dyn LaneObserver>);
let router = Arc::new(LaneRouter::new(
consumer,
LaneRouterConfig {
idle_ttl: self.config.idle_lane_ttl,
runtime: system.runtime().clone(),
tracker: system.tracker().clone(),
observer,
},
));
BoundRouter { router, metrics }
}
fn fail_fast(
system: &Arc<Messenger>,
message_id: ResponseId,
response_type: ResponseType,
reason: &'static str,
) {
if matches!(response_type, ResponseType::FireAndForget) {
return;
}
let backend = system.backend().clone();
tokio::spawn(async move {
if let Err(e) = DispatcherHub::send_error_response_static(
&backend,
message_id,
format!("Handler failed: {reason}"),
)
.await
{
error!(
target: "crate::messenger::dispatcher",
"Failed to send error response for ordered handler: {}", e
);
}
});
}
}
impl<H: ActiveMessageHandler + 'static> ActiveMessageDispatcher for OrderedDispatcher<H> {
fn name(&self) -> &str {
self.handler.name()
}
fn dispatch(&self, ctx: HandlerContext) {
let bound = self.bound.get_or_init(|| self.bind(&ctx.system));
let metrics = bound.metrics.as_ref();
if ctx
.headers
.as_ref()
.is_some_and(|h| h.contains_key(crate::messenger::large_payload::RV_HEADER_KEY))
{
self.rendezvous_warning.call_once(|| {
warn!(
target: "crate::messenger::dispatcher",
handler = %self.handler.name(),
"Rendezvous (large-payload) messages resolve out-of-band and are not \
ordered relative to each other, even on an ordered handler"
);
});
}
let capacity = match self.config.overflow {
OverflowPolicy::Reject => self.config.max_queue_depth,
OverflowPolicy::Warn => None,
};
let key = self.lane_key(&ctx);
let message_id = ctx.message_id;
let response_type = ctx.response_type;
let system = ctx.system.clone();
match bound.router.route(
key,
LaneItem {
ctx,
enqueued_at: Instant::now(),
},
capacity,
) {
Ok(depth_before) => {
if let Some(metrics) = metrics {
metrics.enqueued();
}
if let Some(limit) = self.config.max_queue_depth
&& depth_before == limit
{
warn!(
target: "crate::messenger::dispatcher",
handler = %self.handler.name(),
limit,
lane = ?key,
"Ordered lane exceeded max_queue_depth"
);
}
}
Err(_shed) => {
if let Some(m) = system.observability().as_ref() {
m.record_dispatch_failure(DispatchFailure::OrderedLaneShed);
}
self.shed_warning.call_once(|| {
warn!(
target: "crate::messenger::dispatcher",
handler = %self.handler.name(),
limit = self.config.max_queue_depth.unwrap_or_default(),
lane = ?key,
"Shedding messages: ordered lane is at max_queue_depth. \
Logged once per handler; see velo_messenger_dispatch_failures_total \
{{reason=\"ordered_lane_shed\"}} for the rate"
);
});
Self::fail_fast(
&system,
message_id,
response_type,
"ordered lane queue full",
);
}
}
}
}
pub(crate) struct DispatcherHub {
handlers: Arc<DashMap<String, Arc<dyn ActiveMessageDispatcher>>>,
backend: Arc<VeloBackend>,
system: OnceLock<Arc<Messenger>>,
system_ready: tokio::sync::Notify,
}
impl DispatcherHub {
pub fn new(backend: Arc<VeloBackend>) -> Self {
Self {
handlers: Arc::new(DashMap::new()),
backend,
system: OnceLock::new(),
system_ready: tokio::sync::Notify::new(),
}
}
pub fn set_system(&self, system: Arc<Messenger>) -> anyhow::Result<()> {
self.system
.set(system)
.map_err(|_| anyhow::anyhow!("System already initialized"))?;
self.system_ready.notify_waiters();
Ok(())
}
pub(crate) fn system(&self) -> &Arc<Messenger> {
self.system
.get()
.expect("System must be initialized before dispatching messages")
}
pub(crate) async fn wait_for_system(&self) -> &Arc<Messenger> {
if let Some(system) = self.system.get() {
return system;
}
let notified = self.system_ready.notified();
if let Some(system) = self.system.get() {
return system;
}
notified.await;
self.system
.get()
.expect("system must be set after notification")
}
pub(crate) fn handlers_arc(&self) -> Arc<DashMap<String, Arc<dyn ActiveMessageDispatcher>>> {
self.handlers.clone()
}
pub(crate) fn list_handlers(&self) -> Vec<String> {
self.handlers
.iter()
.map(|entry| entry.key().clone())
.collect()
}
pub fn dispatch_message(&self, handler_name: &str, ctx: HandlerContext) {
match self.handlers.get(handler_name) {
Some(dispatcher) => {
dispatcher.dispatch(ctx);
}
None => {
self.handle_unknown_handler(handler_name, ctx);
}
}
}
fn handle_unknown_handler(&self, handler_name: &str, ctx: HandlerContext) {
if let Some(metrics) = ctx.system.observability().as_ref() {
metrics.record_dispatch_failure(DispatchFailure::DispatchUnknownHandler);
}
error!(
target: "crate::messenger::dispatcher",
handler = %handler_name,
message_id = %ctx.message_id,
"No handler registered for message"
);
let backend = self.backend.clone();
let message_id = ctx.message_id;
let handler_name = handler_name.to_string();
match ctx.response_type {
ResponseType::AckNack | ResponseType::Unary => {
let error_message = format!("Handler '{}' not found", handler_name);
tokio::spawn(async move {
if let Err(e) =
Self::send_error_response_static(&backend, message_id, error_message).await
{
error!(
target: "crate::messenger::dispatcher",
"Failed to send error response for unknown handler: {}", e
);
}
});
}
ResponseType::FireAndForget => {
warn!(
target: "crate::messenger::dispatcher",
handler = %handler_name,
"Fire-and-forget message to unknown handler, no response sent"
);
}
}
}
pub(crate) async fn send_error_response(
&self,
response_id: ResponseId,
error_message: String,
) -> anyhow::Result<()> {
Self::send_error_response_static(&self.backend, response_id, error_message).await
}
async fn send_error_response_static(
backend: &VeloBackend,
response_id: ResponseId,
error_message: String,
) -> anyhow::Result<()> {
use crate::transports::MessageType;
let header = encode_event_header(EventType::Ack(response_id, Outcome::Error));
let payload = Bytes::from(error_message.into_bytes());
struct DispatcherErrorHandler;
impl crate::transports::TransportErrorHandler for DispatcherErrorHandler {
fn on_error(&self, _header: Bytes, _payload: Bytes, error: String) {
error!(target: "crate::messenger::dispatcher", "Failed to send error response: {}", error);
}
}
static ERROR_HANDLER: std::sync::OnceLock<
Arc<dyn crate::transports::TransportErrorHandler>,
> = std::sync::OnceLock::new();
let error_handler = ERROR_HANDLER
.get_or_init(|| Arc::new(DispatcherErrorHandler))
.clone();
let outcome = backend.send_message_to_worker(
WorkerId::from_u64(response_id.worker_id()),
header,
payload,
MessageType::Ack,
error_handler,
)?;
if let crate::transports::SendOutcome::Pending(admission) = outcome {
let _ = admission.await;
}
Ok(())
}
}