use std::collections::HashMap;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use anyhow::{Result, anyhow};
use bytes::Bytes;
use serde::Serialize;
use serde::de::DeserializeOwned;
use super::ActiveMessageClient;
use crate::messenger::common::{ActiveMessage, MessageMetadata};
use crate::observability::ClientResolution;
use crate::transports::{SendBackpressure, SendOutcome};
use velo_ext::{InstanceId, WorkerId};
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) -> impl Future<Output = Result<()>> {
self.inner.fire()
}
pub fn send_to(self, target: InstanceId) -> impl Future<Output = Result<()>> {
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 {
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,
}
}
}
struct ResponseStage {
bp: Option<SendBackpressure>,
awaiter: Option<crate::messenger::common::responses::ResponseAwaiter>,
immediate_error: Option<anyhow::Error>,
}
enum StageState {
Ready(ResponseStage),
Pending(futures::future::BoxFuture<'static, ResponseStage>),
}
impl StageState {
fn ready(stage: ResponseStage) -> Self {
StageState::Ready(stage)
}
fn error(err: anyhow::Error) -> Self {
StageState::Ready(ResponseStage::error(err))
}
fn poll_raw(&mut self, cx: &mut Context<'_>) -> Poll<Result<Option<Bytes>>> {
loop {
match self {
StageState::Pending(fut) => match fut.as_mut().poll(cx) {
Poll::Ready(stage) => *self = StageState::Ready(stage),
Poll::Pending => return Poll::Pending,
},
StageState::Ready(stage) => return stage.poll_raw(cx),
}
}
}
}
impl ResponseStage {
fn ready(awaiter: crate::messenger::common::responses::ResponseAwaiter) -> Self {
Self {
bp: None,
awaiter: Some(awaiter),
immediate_error: None,
}
}
fn with_bp(
awaiter: crate::messenger::common::responses::ResponseAwaiter,
bp: Option<SendBackpressure>,
) -> Self {
Self {
bp,
awaiter: Some(awaiter),
immediate_error: None,
}
}
fn error(err: anyhow::Error) -> Self {
Self {
bp: None,
awaiter: None,
immediate_error: Some(err),
}
}
fn poll_raw(&mut self, cx: &mut Context<'_>) -> Poll<Result<Option<Bytes>>> {
if let Some(err) = self.immediate_error.take() {
return Poll::Ready(Err(err));
}
if let Some(bp) = self.bp.as_mut() {
match Pin::new(bp).poll(cx) {
Poll::Ready(()) => self.bp = None,
Poll::Pending => return Poll::Pending,
}
}
let awaiter = self
.awaiter
.as_mut()
.expect("ResponseStage polled after completion");
match awaiter.poll_recv(cx) {
Poll::Ready(result) => {
self.awaiter = None;
Poll::Ready(result.map_err(|e| anyhow!(e)))
}
Poll::Pending => Poll::Pending,
}
}
}
pub struct SyncResult {
stage: StageState,
}
impl Future for SyncResult {
type Output = Result<()>;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
self.stage.poll_raw(cx).map(|r| r.map(|_| ()))
}
}
pub struct UnaryResult {
stage: StageState,
}
impl Future for UnaryResult {
type Output = Result<Bytes>;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
self.stage
.poll_raw(cx)
.map(|r| r.map(|b| b.unwrap_or_default()))
}
}
pub struct TypedUnaryResult<R> {
stage: StageState,
_marker: std::marker::PhantomData<R>,
}
impl<R> Unpin for TypedUnaryResult<R> {}
impl<R: DeserializeOwned> Future for TypedUnaryResult<R> {
type Output = Result<R>;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
self.stage.poll_raw(cx).map(|r| match r {
Ok(Some(bytes)) => serde_json::from_slice(&bytes)
.map_err(|e| anyhow!("Failed to deserialize response: {}", e)),
Ok(None) => Err(anyhow!("Expected response data, got empty")),
Err(e) => Err(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,
}
async fn drive_send_outcome(
client: &ActiveMessageClient,
send_result: Result<SendOutcome>,
response_id: crate::messenger::common::responses::ResponseId,
path_description: &'static str,
) {
match send_result {
Ok(SendOutcome::Enqueued) => {}
Ok(SendOutcome::Backpressured(bp)) => bp.await,
Err(e) => {
tracing::error!(
target: "crate::messenger::client",
error = %e,
path = path_description,
"Failed to send message"
);
if let Some(metrics) = client.observability.as_ref() {
metrics.record_client_resolution(ClientResolution::SendError);
}
let _ = client
.response_manager
.complete_outcome(response_id, Err(format!("Send failed: {}", e)));
}
}
}
async fn finish_fire_via_awaiter(
mut awaiter: crate::messenger::common::responses::ResponseAwaiter,
) -> Result<()> {
awaiter
.recv()
.await
.map(|_| ())
.map_err(|e| anyhow!("{}", e))
}
async fn drive_fire_send(
send_result: Result<SendOutcome>,
mut awaiter: crate::messenger::common::responses::ResponseAwaiter,
) -> Result<()> {
use futures::FutureExt;
match send_result {
Ok(SendOutcome::Enqueued) => {}
Ok(SendOutcome::Backpressured(bp)) => bp.await,
Err(e) => return Err(e),
}
match awaiter.recv().now_or_never() {
Some(Err(e)) => Err(anyhow!("Send failed: {}", e)),
_ => Ok(()),
}
}
fn stage_from_send(
client: &ActiveMessageClient,
send_result: Result<SendOutcome>,
response_id: crate::messenger::common::responses::ResponseId,
awaiter: crate::messenger::common::responses::ResponseAwaiter,
) -> ResponseStage {
match send_result {
Ok(SendOutcome::Enqueued) => ResponseStage::ready(awaiter),
Ok(SendOutcome::Backpressured(bp)) => ResponseStage::with_bp(awaiter, Some(bp)),
Err(e) => {
tracing::error!(
target: "crate::messenger::client",
error = %e,
"Failed to send message in fast path"
);
let _ = client
.response_manager
.complete_outcome(response_id, Err(format!("Fast-path send failed: {}", e)));
ResponseStage::ready(awaiter)
}
}
}
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: crate::messenger::common::responses::ResponseId,
message_type: MsgType,
) -> MessageMetadata {
match message_type {
MsgType::Sync => {
MessageMetadata::new_sync(response_id, self.handler.clone(), self.headers.clone())
}
MsgType::Unary => {
MessageMetadata::new_unary(response_id, self.handler.clone(), self.headers.clone())
}
}
}
fn spawn_slow_path(
&self,
kind: SlowPathKind,
response_id: crate::messenger::common::responses::ResponseId,
message_type: MsgType,
) {
let client = self.client.clone();
let handler = self.handler.clone();
let payload = self.payload.clone();
let headers = self.headers.clone();
tokio::spawn(async move {
let target = match kind {
SlowPathKind::Handshake(target) => target,
SlowPathKind::Discovery(worker_id) => {
match client.resolve_peer_via_discovery(worker_id).await {
Ok(instance_id) => instance_id,
Err(e) => {
if let Some(metrics) = client.observability.as_ref() {
metrics.record_client_resolution(ClientResolution::DiscoveryError);
}
tracing::error!(
target: "crate::messenger::client",
error = %e,
worker_id = %worker_id,
"Discovery failed"
);
let _ = client.response_manager.complete_outcome(
response_id,
Err(format!("Discovery failed: {}", e)),
);
return;
}
}
}
};
if let Err(e) = client.ensure_peer_ready(target, &handler).await {
if let Some(metrics) = client.observability.as_ref() {
metrics.record_client_resolution(ClientResolution::HandshakeError);
}
tracing::error!(
target: "crate::messenger::client",
error = %e,
"Failed to prepare peer in slow path"
);
let _ = client
.response_manager
.complete_outcome(response_id, Err(format!("Handshake failed: {}", e)));
return;
}
let metadata = match message_type {
MsgType::Sync => MessageMetadata::new_sync(response_id, handler, headers),
MsgType::Unary => MessageMetadata::new_unary(response_id, handler, headers),
};
let message = ActiveMessage {
metadata,
payload: payload.unwrap_or_default(),
};
drive_send_outcome(
&client,
client.send_message(target, message),
response_id,
"slow-path",
)
.await;
});
}
pub async fn fire(self) -> Result<()> {
let target_result = self.resolve_target();
let worker_id = self.target_worker;
let target_result = match target_result {
Err(ResolveError::Other(e)) => return Err(e),
other => other,
};
let outcome = acquire_awaiter(&self.client, self.await_capacity).await?;
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 response_id = outcome.response_id();
let message = ActiveMessage {
metadata: MessageMetadata::new_fire(response_id, self.handler, self.headers),
payload: self.payload.unwrap_or_default(),
};
drive_fire_send(self.client.send_message(target, message), outcome).await
}
Ok(target) => {
let response_id = outcome.response_id();
self.spawn_fire_slow_path(SlowPathKind::Handshake(target), response_id);
finish_fire_via_awaiter(outcome).await
}
Err(ResolveError::UnresolvedPeer) => {
let Some(worker_id) = worker_id else {
return Err(anyhow!("UnresolvedPeer but no worker_id set"));
};
let response_id = outcome.response_id();
self.spawn_fire_slow_path(SlowPathKind::Discovery(worker_id), response_id);
finish_fire_via_awaiter(outcome).await
}
Err(ResolveError::Other(_)) => unreachable!("Other handled above"),
}
}
fn spawn_fire_slow_path(
&self,
kind: SlowPathKind,
response_id: crate::messenger::common::responses::ResponseId,
) {
let client = self.client.clone();
let handler = self.handler.clone();
let payload = self.payload.clone();
let headers = self.headers.clone();
tokio::spawn(async move {
let target = match kind {
SlowPathKind::Handshake(target) => target,
SlowPathKind::Discovery(worker_id) => {
match client.resolve_peer_via_discovery(worker_id).await {
Ok(t) => t,
Err(e) => {
if let Some(metrics) = client.observability.as_ref() {
metrics.record_client_resolution(ClientResolution::DiscoveryError);
}
tracing::error!(
target: "crate::messenger::client",
error = %e,
worker_id = %worker_id,
"Discovery failed for fire-and-forget"
);
let _ = client.response_manager.complete_outcome(
response_id,
Err(format!("Discovery failed: {}", e)),
);
return;
}
}
}
};
if let Err(e) = client.ensure_peer_ready(target, &handler).await {
if let Some(metrics) = client.observability.as_ref() {
metrics.record_client_resolution(ClientResolution::HandshakeError);
}
tracing::error!(
target: "crate::messenger::client",
error = %e,
"Handshake failed for fire-and-forget"
);
let _ = client
.response_manager
.complete_outcome(response_id, Err(format!("Handshake failed: {}", e)));
return;
}
let message = ActiveMessage {
metadata: MessageMetadata::new_fire(response_id, handler, headers),
payload: payload.unwrap_or_default(),
};
match client.send_message(target, message) {
Ok(SendOutcome::Enqueued) => {
let _ = client
.response_manager
.complete_outcome(response_id, Ok(None));
}
Ok(SendOutcome::Backpressured(bp)) => {
bp.await;
let _ = client
.response_manager
.complete_outcome(response_id, Ok(None));
}
Err(e) => {
tracing::error!(
target: "crate::messenger::client",
error = %e,
"Fire-and-forget send failed (slow path)"
);
let _ = client
.response_manager
.complete_outcome(response_id, Err(format!("Send failed: {}", e)));
}
}
});
}
fn dispatch_with_awaiter(
self,
target_result: Result<InstanceId, ResolveError>,
awaiter: crate::messenger::common::responses::ResponseAwaiter,
message_type: MsgType,
) -> ResponseStage {
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 message = ActiveMessage {
metadata: self.create_metadata(response_id, message_type),
payload: self.payload.unwrap_or_default(),
};
let send_result = self.client.send_message(target, message);
stage_from_send(&self.client, send_result, response_id, awaiter)
}
Ok(target) => {
self.spawn_slow_path(SlowPathKind::Handshake(target), response_id, message_type);
ResponseStage::ready(awaiter)
}
Err(ResolveError::UnresolvedPeer) => {
let Some(worker_id) = worker_id else {
tracing::error!(target: "crate::messenger::client", "UnresolvedPeer but no worker_id set");
return ResponseStage::ready(awaiter);
};
self.spawn_slow_path(
SlowPathKind::Discovery(worker_id),
response_id,
message_type,
);
ResponseStage::ready(awaiter)
}
Err(ResolveError::Other(_)) => {
unreachable!("ResolveError::Other is short-circuited in make_stage_state")
}
}
}
fn make_stage_state(self, message_type: MsgType) -> StageState {
let target_result = match self.resolve_target() {
Err(ResolveError::Other(e)) => return StageState::error(e),
other => other,
};
if self.await_capacity {
let fut = Box::pin(async move {
let awaiter = self.client.response_manager.register_outcome_async().await;
self.dispatch_with_awaiter(target_result, awaiter, message_type)
});
StageState::Pending(fut)
} else {
match self.client.register_outcome() {
Ok(awaiter) => StageState::ready(self.dispatch_with_awaiter(
target_result,
awaiter,
message_type,
)),
Err(e) => StageState::error(anyhow!("Failed to register outcome: {}", e)),
}
}
}
pub fn sync(self) -> SyncResult {
SyncResult {
stage: self.make_stage_state(MsgType::Sync),
}
}
pub fn unary(self) -> UnaryResult {
UnaryResult {
stage: self.make_stage_state(MsgType::Unary),
}
}
pub fn typed<R>(self) -> TypedUnaryResult<R>
where
R: DeserializeOwned + Send + 'static,
{
TypedUnaryResult {
stage: self.make_stage_state(MsgType::Unary),
_marker: std::marker::PhantomData,
}
}
}
async fn acquire_awaiter(
client: &ActiveMessageClient,
await_capacity: bool,
) -> Result<crate::messenger::common::responses::ResponseAwaiter> {
if await_capacity {
Ok(client.response_manager.register_outcome_async().await)
} else {
client
.register_outcome()
.map_err(|e| anyhow!("Failed to register outcome: {}", e))
}
}
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(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::messenger::common::responses::ResponseManager;
use crate::transports::SendBackpressure;
fn make_awaiter() -> (
crate::messenger::common::responses::ResponseAwaiter,
crate::messenger::common::responses::ResponseId,
Arc<ResponseManager>,
) {
let rm = Arc::new(ResponseManager::new(1));
let awaiter = rm.register_outcome().expect("register");
let id = awaiter.response_id();
(awaiter, id, rm)
}
fn ready_bp() -> SendBackpressure {
SendBackpressure::new(Box::pin(async {}))
}
fn pending_bp() -> SendBackpressure {
SendBackpressure::new(Box::pin(futures::future::pending::<()>()))
}
#[tokio::test]
async fn stage_ready_resolves_after_outcome_completes() {
let (awaiter, id, rm) = make_awaiter();
let stage = ResponseStage::ready(awaiter);
let mut result = SyncResult {
stage: StageState::Ready(stage),
};
assert!(rm.complete_outcome(id, Ok(Some(Bytes::from_static(b"ok")))));
let r = tokio::time::timeout(std::time::Duration::from_secs(1), &mut result)
.await
.expect("sync result completes");
assert!(r.is_ok());
}
#[tokio::test]
async fn stage_with_ready_bp_proceeds_to_awaiter() {
let (awaiter, id, rm) = make_awaiter();
let stage = ResponseStage::with_bp(awaiter, Some(ready_bp()));
let mut result = UnaryResult {
stage: StageState::Ready(stage),
};
assert!(rm.complete_outcome(id, Ok(Some(Bytes::from_static(b"hello")))));
let r = tokio::time::timeout(std::time::Duration::from_secs(1), &mut result)
.await
.expect("unary result completes");
assert_eq!(r.unwrap(), Bytes::from_static(b"hello"));
}
#[tokio::test]
async fn stage_pending_bp_blocks_until_resolved() {
let (awaiter, _id, _rm) = make_awaiter();
let stage = ResponseStage::with_bp(awaiter, Some(pending_bp()));
let result = SyncResult {
stage: StageState::Ready(stage),
};
let outcome = tokio::time::timeout(std::time::Duration::from_millis(100), result).await;
assert!(outcome.is_err(), "pending bp should keep result pending");
}
#[tokio::test]
async fn stage_immediate_error_short_circuits() {
let stage = ResponseStage::error(anyhow!("boom"));
let result = SyncResult {
stage: StageState::Ready(stage),
};
let err = result.await.expect_err("immediate_error returns Err");
assert!(err.to_string().contains("boom"));
}
#[tokio::test]
async fn unary_result_empty_response_becomes_empty_bytes() {
let (awaiter, id, rm) = make_awaiter();
let stage = ResponseStage::ready(awaiter);
let mut result = UnaryResult {
stage: StageState::Ready(stage),
};
assert!(rm.complete_outcome(id, Ok(None)));
let r = tokio::time::timeout(std::time::Duration::from_secs(1), &mut result)
.await
.expect("unary resolves")
.unwrap();
assert_eq!(r, Bytes::new());
}
#[tokio::test]
async fn typed_result_deserializes_payload() {
let (awaiter, id, rm) = make_awaiter();
let stage = ResponseStage::ready(awaiter);
let mut result: TypedUnaryResult<i64> = TypedUnaryResult {
stage: StageState::Ready(stage),
_marker: std::marker::PhantomData,
};
assert!(rm.complete_outcome(id, Ok(Some(Bytes::from(b"42".to_vec())))));
let v = tokio::time::timeout(std::time::Duration::from_secs(1), &mut result)
.await
.expect("typed resolves")
.unwrap();
assert_eq!(v, 42);
}
#[tokio::test]
async fn typed_result_empty_response_is_error() {
let (awaiter, id, rm) = make_awaiter();
let stage = ResponseStage::ready(awaiter);
let mut result: TypedUnaryResult<i64> = TypedUnaryResult {
stage: StageState::Ready(stage),
_marker: std::marker::PhantomData,
};
assert!(rm.complete_outcome(id, Ok(None)));
let err = tokio::time::timeout(std::time::Duration::from_secs(1), &mut result)
.await
.expect("typed resolves")
.expect_err("empty response → Err");
assert!(err.to_string().contains("Expected response data"));
}
#[tokio::test]
async fn typed_result_bad_json_is_error() {
let (awaiter, id, rm) = make_awaiter();
let stage = ResponseStage::ready(awaiter);
let mut result: TypedUnaryResult<i64> = TypedUnaryResult {
stage: StageState::Ready(stage),
_marker: std::marker::PhantomData,
};
assert!(rm.complete_outcome(id, Ok(Some(Bytes::from_static(b"not-json")))));
let err = tokio::time::timeout(std::time::Duration::from_secs(1), &mut result)
.await
.expect("typed resolves")
.expect_err("bad json → Err");
assert!(err.to_string().contains("Failed to deserialize"));
}
#[tokio::test]
async fn stage_state_pending_transitions_and_resolves() {
let (awaiter, id, rm) = make_awaiter();
let fut: futures::future::BoxFuture<'static, ResponseStage> =
Box::pin(async move { ResponseStage::ready(awaiter) });
let mut result = SyncResult {
stage: StageState::Pending(fut),
};
assert!(rm.complete_outcome(id, Ok(None)));
let r = tokio::time::timeout(std::time::Duration::from_secs(1), &mut result)
.await
.expect("sync result resolves");
assert!(r.is_ok());
}
#[tokio::test]
async fn stage_state_pending_awaits_inner_future() {
let (awaiter, _id, _rm) = make_awaiter();
let fut: futures::future::BoxFuture<'static, ResponseStage> = Box::pin(async {
futures::future::pending::<()>().await;
ResponseStage::ready(awaiter)
});
let result = SyncResult {
stage: StageState::Pending(fut),
};
let outcome = tokio::time::timeout(std::time::Duration::from_millis(50), result).await;
assert!(
outcome.is_err(),
"pending inner future should keep result pending"
);
}
#[tokio::test]
async fn drive_fire_send_enqueued_is_ok() {
let (awaiter, _id, _rm) = make_awaiter();
assert!(
drive_fire_send(Ok(SendOutcome::Enqueued), awaiter)
.await
.is_ok()
);
}
#[tokio::test]
async fn drive_fire_send_bp_without_error_is_ok() {
let (awaiter, _id, _rm) = make_awaiter();
assert!(
drive_fire_send(Ok(SendOutcome::Backpressured(ready_bp())), awaiter)
.await
.is_ok()
);
}
#[tokio::test]
async fn drive_fire_send_enqueued_with_sync_on_error_surfaces_err() {
let (awaiter, id, rm) = make_awaiter();
assert!(rm.complete_outcome(id, Err("Connection closed immediately".to_string())));
let err = drive_fire_send(Ok(SendOutcome::Enqueued), awaiter)
.await
.expect_err("should surface sync on_error failure");
assert!(err.to_string().contains("Connection closed immediately"));
}
#[tokio::test]
async fn drive_fire_send_bp_with_on_error_surfaces_err() {
let (awaiter, id, rm) = make_awaiter();
assert!(rm.complete_outcome(id, Err("peer disconnected".to_string())));
let err = drive_fire_send(Ok(SendOutcome::Backpressured(ready_bp())), awaiter)
.await
.expect_err("should surface pre-wire failure");
assert!(err.to_string().contains("peer disconnected"));
}
#[tokio::test]
async fn drive_fire_send_sync_err_is_propagated() {
let (awaiter, _id, _rm) = make_awaiter();
let err = drive_fire_send(Err(anyhow!("peer not registered")), awaiter)
.await
.expect_err("sync err propagates");
assert!(err.to_string().contains("peer not registered"));
}
#[tokio::test]
async fn finish_fire_via_awaiter_ok_on_success_completion() {
let (awaiter, id, rm) = make_awaiter();
assert!(rm.complete_outcome(id, Ok(None)));
assert!(finish_fire_via_awaiter(awaiter).await.is_ok());
}
#[tokio::test]
async fn finish_fire_via_awaiter_err_on_failure_completion() {
let (awaiter, id, rm) = make_awaiter();
assert!(rm.complete_outcome(id, Err("Handshake failed: nope".to_string())));
let err = finish_fire_via_awaiter(awaiter)
.await
.expect_err("err completion surfaces");
assert!(err.to_string().contains("Handshake failed"));
}
#[test]
fn validate_handler_name_accepts_public() {
assert!(validate_handler_name("my_handler").is_ok());
}
#[test]
fn validate_handler_name_rejects_system() {
let err = validate_handler_name("_hello").unwrap_err();
assert!(
err.to_string()
.contains("Cannot directly call system handler")
);
}
}