use std::collections::HashMap;
use std::sync::Arc;
use anyhow::{Result, anyhow};
use bytes::Bytes;
use serde::Serialize;
use serde::de::DeserializeOwned;
use tokio::sync::oneshot;
use super::ActiveMessageClient;
use crate::messenger::common::responses::{ResponseAwaiter, ResponseId, SlotBackpressure};
use crate::messenger::common::{ActiveMessage, MessageMetadata};
use crate::observability::ClientResolution;
use crate::transports::SendOutcome;
use velo_ext::{InstanceId, WorkerId};
mod results;
#[cfg(test)]
mod tests;
use results::{AdmissionReport, Dispatched, SendStage, drive_send_outcome};
pub use results::{Admitted, FireResult, SyncResult, TypedUnaryResult, UnaryResult};
pub struct AmSendBuilder {
inner: MessageBuilder,
}
impl AmSendBuilder {
pub(crate) fn new(client: Arc<ActiveMessageClient>, handler: &str) -> Result<Self> {
Ok(Self {
inner: MessageBuilder::new(client, handler)?,
})
}
pub(crate) fn new_unchecked(client: Arc<ActiveMessageClient>, handler: &str) -> Self {
Self {
inner: MessageBuilder::new_unchecked(client, handler),
}
}
pub fn payload<T: Serialize>(mut self, data: T) -> Result<Self> {
self.inner = self.inner.payload(data)?;
Ok(self)
}
pub fn raw_payload(mut self, data: Bytes) -> Self {
self.inner = self.inner.raw_payload(data);
self
}
pub fn instance(mut self, instance_id: InstanceId) -> Self {
self.inner = self.inner.instance(instance_id);
self
}
pub fn worker(mut self, worker_id: WorkerId) -> Self {
self.inner = self.inner.worker(worker_id);
self
}
pub fn headers(mut self, headers: HashMap<String, String>) -> Self {
self.inner = self.inner.headers(headers);
self
}
pub fn await_capacity(mut self) -> Self {
self.inner = self.inner.await_capacity();
self
}
pub fn send(self) -> FireResult {
self.inner.fire()
}
pub fn send_to(self, target: InstanceId) -> FireResult {
self.inner.instance(target).fire()
}
}
pub struct AmSyncBuilder {
inner: MessageBuilder,
}
impl AmSyncBuilder {
pub(crate) fn new(client: Arc<ActiveMessageClient>, handler: &str) -> Result<Self> {
Ok(Self {
inner: MessageBuilder::new(client, handler)?,
})
}
pub fn payload<T: Serialize>(mut self, data: T) -> Result<Self> {
self.inner = self.inner.payload(data)?;
Ok(self)
}
pub fn raw_payload(mut self, data: Bytes) -> Self {
self.inner = self.inner.raw_payload(data);
self
}
pub fn instance(mut self, instance_id: InstanceId) -> Self {
self.inner = self.inner.instance(instance_id);
self
}
pub fn worker(mut self, worker_id: WorkerId) -> Self {
self.inner = self.inner.worker(worker_id);
self
}
pub fn headers(mut self, headers: HashMap<String, String>) -> Self {
self.inner = self.inner.headers(headers);
self
}
pub fn await_capacity(mut self) -> Self {
self.inner = self.inner.await_capacity();
self
}
pub fn send(self) -> SyncResult {
self.inner.sync()
}
pub fn send_to(self, target: InstanceId) -> SyncResult {
self.inner.instance(target).sync()
}
}
pub struct UnaryBuilder {
inner: MessageBuilder,
}
impl UnaryBuilder {
pub(crate) fn new(client: Arc<ActiveMessageClient>, handler: &str) -> Result<Self> {
Ok(Self {
inner: MessageBuilder::new(client, handler)?,
})
}
pub(crate) fn new_unchecked(client: Arc<ActiveMessageClient>, handler: &str) -> Self {
Self {
inner: MessageBuilder::new_unchecked(client, handler),
}
}
pub fn payload<T: Serialize>(mut self, data: T) -> Result<Self> {
self.inner = self.inner.payload(data)?;
Ok(self)
}
pub fn raw_payload(mut self, data: Bytes) -> Self {
self.inner = self.inner.raw_payload(data);
self
}
pub fn instance(mut self, instance_id: InstanceId) -> Self {
self.inner = self.inner.instance(instance_id);
self
}
pub fn worker(mut self, worker_id: WorkerId) -> Self {
self.inner = self.inner.worker(worker_id);
self
}
pub fn headers(mut self, headers: HashMap<String, String>) -> Self {
self.inner = self.inner.headers(headers);
self
}
pub fn await_capacity(mut self) -> Self {
self.inner = self.inner.await_capacity();
self
}
pub fn send(self) -> UnaryResult {
self.inner.unary()
}
pub fn send_to(self, target: InstanceId) -> UnaryResult {
self.inner.instance(target).unary()
}
}
pub struct TypedUnaryBuilder<R> {
inner: MessageBuilder,
_marker: std::marker::PhantomData<R>,
}
impl<R> TypedUnaryBuilder<R>
where
R: DeserializeOwned + Send + 'static,
{
pub(crate) fn new(client: Arc<ActiveMessageClient>, handler: &str) -> Result<Self> {
Ok(Self {
inner: MessageBuilder::new(client, handler)?,
_marker: std::marker::PhantomData,
})
}
pub(crate) fn new_unchecked(client: Arc<ActiveMessageClient>, handler: &str) -> Self {
Self {
inner: MessageBuilder::new_unchecked(client, handler),
_marker: std::marker::PhantomData,
}
}
pub fn payload<T: Serialize>(mut self, data: T) -> Result<Self> {
self.inner = self.inner.payload(data)?;
Ok(self)
}
pub fn raw_payload(mut self, data: Bytes) -> Self {
self.inner = self.inner.raw_payload(data);
self
}
pub fn instance(mut self, instance_id: InstanceId) -> Self {
self.inner = self.inner.instance(instance_id);
self
}
pub fn worker(mut self, worker_id: WorkerId) -> Self {
self.inner = self.inner.worker(worker_id);
self
}
pub fn headers(mut self, headers: HashMap<String, String>) -> Self {
self.inner = self.inner.headers(headers);
self
}
pub fn await_capacity(mut self) -> Self {
self.inner = self.inner.await_capacity();
self
}
pub fn send(self) -> TypedUnaryResult<R> {
self.inner.typed()
}
pub fn send_to(self, target: InstanceId) -> TypedUnaryResult<R> {
self.inner.instance(target).typed()
}
}
#[derive(Debug)]
enum ResolveError {
UnresolvedPeer,
Other(anyhow::Error),
}
#[derive(Debug, Clone, Copy)]
enum MsgType {
Fire,
Sync,
Unary,
}
#[derive(Debug, Clone, Copy)]
enum SlowPathKind {
Handshake(InstanceId),
Discovery(WorkerId),
}
impl From<ResolveError> for anyhow::Error {
fn from(err: ResolveError) -> Self {
match err {
ResolveError::UnresolvedPeer => anyhow!("Peer not found"),
ResolveError::Other(e) => e,
}
}
}
pub struct MessageBuilder {
client: Arc<ActiveMessageClient>,
handler: String,
payload: Option<Bytes>,
target_instance: Option<InstanceId>,
target_worker: Option<WorkerId>,
headers: Option<HashMap<String, String>>,
await_capacity: bool,
}
fn stage_from_send(send_result: Result<SendOutcome>, awaiter: ResponseAwaiter) -> Dispatched {
match send_result {
Ok(outcome) => Dispatched::issued(outcome, awaiter),
Err(e) => {
tracing::error!(
target: "crate::messenger::client",
error = %e,
"Failed to send message in fast path"
);
Dispatched::failed(format!("Fast-path send failed: {}", e))
}
}
}
fn build_metadata(
response_id: ResponseId,
handler: String,
headers: Option<HashMap<String, String>>,
message_type: MsgType,
) -> MessageMetadata {
match message_type {
MsgType::Fire => MessageMetadata::new_fire(response_id, handler, headers),
MsgType::Sync => MessageMetadata::new_sync(response_id, handler, headers),
MsgType::Unary => MessageMetadata::new_unary(response_id, handler, headers),
}
}
enum Acquisition {
Allocated(ResponseAwaiter),
Deferred(SlotBackpressure),
Exhausted(anyhow::Error),
}
struct SlowPath {
client: Arc<ActiveMessageClient>,
kind: SlowPathKind,
handler: String,
payload: Option<Bytes>,
headers: Option<HashMap<String, String>>,
response_id: ResponseId,
message_type: MsgType,
}
impl SlowPath {
async fn run(self) -> AdmissionReport {
let target = self.resolve().await?;
self.handshake(target).await?;
self.send(target).await
}
async fn resolve(&self) -> std::result::Result<InstanceId, String> {
let worker_id = match self.kind {
SlowPathKind::Handshake(target) => return Ok(target),
SlowPathKind::Discovery(worker_id) => worker_id,
};
match self.client.resolve_peer_via_discovery(worker_id).await {
Ok(instance_id) => Ok(instance_id),
Err(e) => {
self.record(ClientResolution::DiscoveryError);
tracing::error!(
target: "crate::messenger::client",
error = %e,
worker_id = %worker_id,
"Discovery failed"
);
Err(self.fail(format!("Discovery failed: {}", e)))
}
}
}
async fn handshake(&self, target: InstanceId) -> AdmissionReport {
if let Err(e) = self.client.ensure_peer_ready(target, &self.handler).await {
self.record(ClientResolution::HandshakeError);
tracing::error!(
target: "crate::messenger::client",
error = %e,
"Failed to prepare peer in slow path"
);
return Err(self.fail(format!("Handshake failed: {}", e)));
}
Ok(())
}
async fn send(self, target: InstanceId) -> AdmissionReport {
let metadata = build_metadata(
self.response_id,
self.handler,
self.headers,
self.message_type,
);
let message = ActiveMessage {
metadata,
payload: self.payload.unwrap_or_default(),
};
drive_send_outcome(
&self.client,
self.client.send_message(target, message),
self.response_id,
"slow-path",
)
.await
}
fn fail(&self, reason: String) -> String {
let _ = self
.client
.response_manager
.complete_outcome(self.response_id, Err(reason.clone()));
reason
}
fn record(&self, resolution: ClientResolution) {
if let Some(metrics) = self.client.observability.as_ref() {
metrics.record_client_resolution(resolution);
}
}
}
impl MessageBuilder {
pub fn new(client: Arc<ActiveMessageClient>, handler: &str) -> Result<Self> {
validate_handler_name(handler)?;
Ok(Self::new_unchecked(client, handler))
}
pub fn new_unchecked(client: Arc<ActiveMessageClient>, handler: &str) -> Self {
Self {
client,
handler: handler.to_string(),
payload: None,
target_instance: None,
target_worker: None,
headers: None,
await_capacity: false,
}
}
pub fn payload<T: Serialize>(mut self, data: T) -> Result<Self> {
let bytes =
serde_json::to_vec(&data).map_err(|e| anyhow!("failed to serialize payload: {}", e))?;
self.payload = Some(Bytes::from(bytes));
Ok(self)
}
pub fn raw_payload(mut self, data: Bytes) -> Self {
self.payload = Some(data);
self
}
pub fn instance(mut self, instance_id: InstanceId) -> Self {
self.target_instance = Some(instance_id);
self
}
pub fn worker(mut self, worker_id: WorkerId) -> Self {
self.target_worker = Some(worker_id);
self
}
pub fn headers(mut self, headers: HashMap<String, String>) -> Self {
self.headers = Some(headers);
self
}
pub fn await_capacity(mut self) -> Self {
self.await_capacity = true;
self
}
fn resolve_target(&self) -> Result<InstanceId, ResolveError> {
match (self.target_instance, self.target_worker) {
(Some(instance), None) => Ok(instance),
(None, Some(worker)) => self
.client
.backend
.try_translate_worker_id(worker)
.map_err(|_| ResolveError::UnresolvedPeer),
(Some(_), Some(_)) => Err(ResolveError::Other(anyhow!(
"Cannot set both .instance() and .worker() - they are mutually exclusive"
))),
(None, None) => Err(ResolveError::Other(anyhow!(
"Target not set. Call .instance() or .worker() before sending"
))),
}
}
fn create_metadata(&self, response_id: ResponseId, message_type: MsgType) -> MessageMetadata {
build_metadata(
response_id,
self.handler.clone(),
self.headers.clone(),
message_type,
)
}
fn acquire(&self) -> Acquisition {
if self.await_capacity {
match self.client.response_manager.try_register_outcome() {
crate::messenger::common::responses::RegisterOutcome::Allocated(awaiter) => {
Acquisition::Allocated(awaiter)
}
crate::messenger::common::responses::RegisterOutcome::Backpressured(
backpressure,
) => Acquisition::Deferred(backpressure),
}
} else {
match self.client.register_outcome() {
Ok(awaiter) => Acquisition::Allocated(awaiter),
Err(e) => Acquisition::Exhausted(anyhow!("Failed to register outcome: {}", e)),
}
}
}
fn spawn_slow_path(
&self,
kind: SlowPathKind,
response_id: ResponseId,
message_type: MsgType,
) -> oneshot::Receiver<AdmissionReport> {
let (report_tx, report_rx) = oneshot::channel();
let slow_path = SlowPath {
client: self.client.clone(),
kind,
handler: self.handler.clone(),
payload: self.payload.clone(),
headers: self.headers.clone(),
response_id,
message_type,
};
tokio::spawn(async move {
let _ = report_tx.send(slow_path.run().await);
});
report_rx
}
fn dispatch(
self,
target_result: Result<InstanceId, ResolveError>,
awaiter: ResponseAwaiter,
message_type: MsgType,
) -> Dispatched {
let worker_id = self.target_worker;
let response_id = awaiter.response_id();
match target_result {
Ok(target) if self.client.can_send_directly(target, &self.handler) => {
if let Some(metrics) = self.client.observability.as_ref() {
metrics.record_client_resolution(ClientResolution::DirectSuccess);
}
let metadata = self.create_metadata(response_id, message_type);
let message = ActiveMessage {
metadata,
payload: self.payload.unwrap_or_default(),
};
stage_from_send(self.client.send_message(target, message), awaiter)
}
Ok(target) => Dispatched::detached(
self.spawn_slow_path(SlowPathKind::Handshake(target), response_id, message_type),
awaiter,
),
Err(ResolveError::UnresolvedPeer) => {
let Some(worker_id) = worker_id else {
tracing::error!(target: "crate::messenger::client", "UnresolvedPeer but no worker_id set");
return Dispatched::failed("UnresolvedPeer but no worker_id set");
};
Dispatched::detached(
self.spawn_slow_path(
SlowPathKind::Discovery(worker_id),
response_id,
message_type,
),
awaiter,
)
}
Err(ResolveError::Other(_)) => {
unreachable!("ResolveError::Other is short-circuited in make_stage")
}
}
}
fn make_stage(self, message_type: MsgType) -> SendStage {
let target_result = match self.resolve_target() {
Err(ResolveError::Other(e)) => return SendStage::failed(e),
other => other,
};
match self.acquire() {
Acquisition::Allocated(awaiter) => {
SendStage::Dispatched(self.dispatch(target_result, awaiter, message_type))
}
Acquisition::Deferred(backpressure) => {
let deferred = Box::pin(async move {
let awaiter = backpressure.await;
self.dispatch(target_result, awaiter, message_type)
});
SendStage::Acquiring(deferred)
}
Acquisition::Exhausted(e) => SendStage::failed(e),
}
}
pub fn fire(self) -> FireResult {
FireResult {
stage: self.make_stage(MsgType::Fire),
}
}
pub fn sync(self) -> SyncResult {
SyncResult {
stage: self.make_stage(MsgType::Sync),
}
}
pub fn unary(self) -> UnaryResult {
UnaryResult {
stage: self.make_stage(MsgType::Unary),
}
}
pub fn typed<R>(self) -> TypedUnaryResult<R>
where
R: DeserializeOwned + Send + 'static,
{
TypedUnaryResult {
stage: self.make_stage(MsgType::Unary),
_marker: std::marker::PhantomData,
}
}
}
pub(crate) fn validate_handler_name(handler: &str) -> Result<()> {
if handler.starts_with('_') {
anyhow::bail!(
"Cannot directly call system handler '{}'. Use client convenience methods instead: health_check(), ensure_bidirectional_connection(), list_handlers(), await_handler()",
handler
);
}
Ok(())
}