use std::path::Path;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::Duration;
use bamboo_subagent::{AgentRef, InboxMessage, MsgId};
use futures_util::stream::SplitSink;
use futures_util::{SinkExt, StreamExt};
use tokio::net::TcpStream;
use tokio::sync::mpsc;
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
use tokio_tungstenite::tungstenite::Message;
use tokio_tungstenite::{
connect_async_tls_with_config, Connector, MaybeTlsStream, WebSocketStream,
};
use crate::error::{BrokerError, BrokerResult};
use crate::proto::{BrokerFrame, ClientFrame};
pub fn client_config_trusting_cert(cert_file: &Path) -> BrokerResult<rustls::ClientConfig> {
bamboo_subagent::transport::client_config_trusting_cert(cert_file).map_err(BrokerError::Tls)
}
pub(crate) type WsSink = SplitSink<WebSocketStream<MaybeTlsStream<TcpStream>>, Message>;
const DELIVER_RECEIPT_TIMEOUT: Duration = Duration::from_secs(30);
pub enum ServeEvent {
Message(Option<InboxMessage>),
Cancel(Option<MsgId>),
}
pub struct BrokerClient {
sink: WsSink,
messages: mpsc::UnboundedReceiver<InboxMessage>,
delivered: mpsc::UnboundedReceiver<MsgId>,
errors: mpsc::UnboundedReceiver<(MsgId, String)>,
cancels: mpsc::UnboundedReceiver<MsgId>,
connected: mpsc::UnboundedReceiver<Vec<String>>,
reader_alive: Arc<AtomicBool>,
_supervisor: tokio::task::JoinHandle<()>,
}
impl BrokerClient {
pub async fn connect(endpoint: &str, agent: AgentRef, token: &str) -> BrokerResult<Self> {
Self::connect_with_tls(endpoint, agent, token, None).await
}
pub async fn connect_with_tls(
endpoint: &str,
agent: AgentRef,
token: &str,
tls_config: Option<rustls::ClientConfig>,
) -> BrokerResult<Self> {
let request = endpoint
.into_client_request()
.map_err(|e| BrokerError::Transport(format!("bad endpoint '{endpoint}': {e}")))?;
let connector = tls_config.map(|cfg| Connector::Rustls(Arc::new(cfg)));
let (ws, _resp) = connect_async_tls_with_config(request, None, false, connector)
.await
.map_err(|e| BrokerError::Transport(format!("connect {endpoint}: {e}")))?;
let (mut sink, mut source) = ws.split();
sink.send(Message::text(
ClientFrame::Hello {
agent,
token: token.into(),
}
.to_text(),
))
.await
.map_err(|e| BrokerError::Transport(format!("send hello: {e}")))?;
loop {
match source.next().await {
Some(Ok(Message::Text(t))) => match BrokerFrame::from_text(&t) {
Ok(BrokerFrame::Welcome) => break,
Ok(BrokerFrame::Error { reason, .. }) => return Err(BrokerError::Auth(reason)),
Ok(_) => continue,
Err(e) => return Err(BrokerError::Protocol(format!("bad broker frame: {e}"))),
},
Some(Ok(_)) => continue,
Some(Err(e)) => return Err(BrokerError::Transport(format!("ws: {e}"))),
None => return Err(BrokerError::Transport("closed during handshake".into())),
}
}
let (msg_tx, messages) = mpsc::unbounded_channel();
let (del_tx, delivered) = mpsc::unbounded_channel();
let (err_tx, errors) = mpsc::unbounded_channel();
let (cancel_tx, cancels) = mpsc::unbounded_channel();
let (conn_tx, connected) = mpsc::unbounded_channel();
let reader = tokio::spawn(async move {
while let Some(frame) = source.next().await {
match frame {
Ok(Message::Text(t)) => match BrokerFrame::from_text(&t) {
Ok(BrokerFrame::Message { message }) => {
let _ = msg_tx.send(message);
}
Ok(BrokerFrame::Delivered { id }) => {
let _ = del_tx.send(id);
}
Ok(BrokerFrame::Error {
reason,
id: Some(id),
}) => {
let _ = err_tx.send((id, reason));
}
Ok(BrokerFrame::Error { reason, id: None }) => {
tracing::warn!(
"broker sent an uncorrelated error frame post-handshake: {reason}"
);
}
Ok(BrokerFrame::Cancel { correlation_id }) => {
let _ = cancel_tx.send(correlation_id);
}
Ok(BrokerFrame::Connected { ids }) => {
let _ = conn_tx.send(ids);
}
_ => {}
},
Ok(Message::Close(_)) | Err(_) => break,
_ => {}
}
}
});
let reader_alive = Arc::new(AtomicBool::new(true));
let supervisor = tokio::spawn(reader_supervisor(reader, reader_alive.clone()));
Ok(Self {
sink,
messages,
delivered,
errors,
cancels,
connected,
reader_alive,
_supervisor: supervisor,
})
}
pub async fn deliver(&mut self, to: &str, message: InboxMessage) -> BrokerResult<MsgId> {
self.deliver_with_receipt_timeout(to, message, DELIVER_RECEIPT_TIMEOUT)
.await
}
async fn deliver_with_receipt_timeout(
&mut self,
to: &str,
message: InboxMessage,
receipt_timeout: Duration,
) -> BrokerResult<MsgId> {
let expected = message.id.clone();
self.send(ClientFrame::Deliver {
to: to.into(),
message,
})
.await?;
let deadline = tokio::time::Instant::now() + receipt_timeout;
let correlate = async {
loop {
tokio::select! {
biased;
err = self.errors.recv() => match err {
Some((id, reason)) if id == expected => {
return Err(BrokerError::Rejected(reason));
}
Some(_stale) => continue,
None => {
return Err(BrokerError::Transport(
"connection closed before delivery receipt".into(),
));
}
},
id = self.delivered.recv() => match id {
Some(id) if id == expected => return Ok(id),
Some(_stale) => continue,
None => {
return Err(BrokerError::Transport(
"connection closed before delivery receipt".into(),
));
}
},
}
}
};
match tokio::time::timeout_at(deadline, correlate).await {
Ok(outcome) => outcome,
Err(_) => Err(BrokerError::Transport(
"timed out waiting for delivery receipt from broker".into(),
)),
}
}
pub async fn subscribe(&mut self) -> BrokerResult<()> {
self.send(ClientFrame::Subscribe).await
}
pub async fn list_connected(&mut self, role: &str) -> BrokerResult<Vec<String>> {
self.send(ClientFrame::ListConnected { role: role.into() })
.await?;
match tokio::time::timeout(DELIVER_RECEIPT_TIMEOUT, self.connected.recv()).await {
Ok(Some(ids)) => Ok(ids),
Ok(None) => Err(BrokerError::Transport(
"connection closed before connected-actors reply".into(),
)),
Err(_) => Err(BrokerError::Transport(
"timed out waiting for connected-actors reply from broker".into(),
)),
}
}
pub async fn next_message(&mut self) -> Option<InboxMessage> {
let msg = self.messages.recv().await;
if msg.is_none() && !self.reader_alive.load(Ordering::SeqCst) {
tracing::warn!(
"broker next_message() returned None: reader task exited (connection closed)"
);
}
msg
}
pub async fn next_cancel(&mut self) -> Option<MsgId> {
self.cancels.recv().await
}
pub async fn next_message_or_cancel(&mut self) -> ServeEvent {
tokio::select! {
biased;
cancel = self.cancels.recv() => ServeEvent::Cancel(cancel),
msg = self.messages.recv() => {
if msg.is_none() && !self.reader_alive.load(Ordering::SeqCst) {
tracing::warn!(
"broker next_message() returned None: reader task exited (connection closed)"
);
}
ServeEvent::Message(msg)
}
}
}
pub fn reader_alive(&self) -> bool {
self.reader_alive.load(Ordering::SeqCst)
}
pub async fn ack(&mut self, id: MsgId) -> BrokerResult<()> {
self.send(ClientFrame::Ack { id }).await
}
pub fn into_multiplexed(self, me: AgentRef) -> crate::mux::MultiplexedClient {
crate::mux::MultiplexedClient::spawn(
self.sink,
self.messages,
self.delivered,
self.reader_alive,
me,
)
}
pub async fn cancel(&mut self, to: &str, correlation_id: &MsgId) -> BrokerResult<()> {
self.send(ClientFrame::Cancel {
to: to.into(),
correlation_id: correlation_id.clone(),
})
.await
}
async fn send(&mut self, frame: ClientFrame) -> BrokerResult<()> {
self.sink
.send(Message::text(frame.to_text()))
.await
.map_err(|e| BrokerError::Transport(format!("ws send: {e}")))
}
}
async fn reader_supervisor(reader: tokio::task::JoinHandle<()>, reader_alive: Arc<AtomicBool>) {
let outcome = reader.await;
reader_alive.store(false, Ordering::SeqCst);
match outcome {
Ok(()) => {
tracing::warn!("broker reader task ended; connection closed");
}
Err(err) if err.is_panic() => {
tracing::error!(
"broker reader task panicked: {}",
panic_payload_message(err.into_panic())
);
}
Err(err) => {
tracing::error!("broker reader task ended unexpectedly (cancelled/aborted): {err}");
}
}
}
fn panic_payload_message(payload: Box<dyn std::any::Any + Send>) -> String {
payload
.downcast_ref::<&'static str>()
.map(|s| (*s).to_string())
.or_else(|| payload.downcast_ref::<String>().cloned())
.unwrap_or_else(|| "<non-string panic payload>".to_string())
}
#[cfg(test)]
mod tests {
use super::*;
use bamboo_subagent::{AskBody, AskMode, InboxKind};
use chrono::Utc;
use tokio_tungstenite::accept_async;
fn test_agent(id: &str) -> AgentRef {
AgentRef {
session_id: id.into(),
role: None,
}
}
fn test_ask(from: &str) -> InboxMessage {
InboxMessage {
id: MsgId::new(),
from: test_agent(from),
kind: InboxKind::Ask,
body: serde_json::to_value(AskBody {
question: "ping".into(),
mode: AskMode::Query,
})
.unwrap(),
created_at: Utc::now(),
correlation_id: None,
}
}
async fn broker_that_never_acks() -> String {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("bind");
let addr = listener.local_addr().expect("local_addr");
tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("accept");
let ws = accept_async(stream).await.expect("ws upgrade");
let (mut sink, mut source) = ws.split();
if let Some(Ok(Message::Text(_))) = source.next().await {
let _ = sink
.send(Message::text(BrokerFrame::Welcome.to_text()))
.await;
}
while let Some(Ok(_)) = source.next().await {}
});
format!("ws://{addr}")
}
#[tokio::test]
async fn deliver_times_out_when_broker_never_sends_receipt() {
let endpoint = broker_that_never_acks().await;
let mut client = BrokerClient::connect(&endpoint, test_agent("parent"), "ignored")
.await
.expect("handshake completes");
let started = std::time::Instant::now();
let outcome = tokio::time::timeout(
Duration::from_secs(2),
client.deliver_with_receipt_timeout(
"child",
test_ask("parent"),
Duration::from_millis(50),
),
)
.await;
let result = outcome.expect("deliver() resolved instead of hanging");
assert!(
matches!(result, Err(BrokerError::Transport(ref m)) if m.contains("timed out")),
"expected a timeout transport error, got {result:?}",
);
assert!(
started.elapsed() < Duration::from_secs(1),
"deliver() should fail fast, but took {:?}",
started.elapsed(),
);
}
async fn broker_that_closes_after_handshake() -> String {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("bind");
let addr = listener.local_addr().expect("local_addr");
tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("accept");
let ws = accept_async(stream).await.expect("ws upgrade");
let (mut sink, mut source) = ws.split();
if let Some(Ok(Message::Text(_))) = source.next().await {
let _ = sink
.send(Message::text(BrokerFrame::Welcome.to_text()))
.await;
}
let _ = sink.send(Message::Close(None)).await;
while let Some(Ok(_)) = source.next().await {}
});
format!("ws://{addr}")
}
#[tokio::test]
async fn reader_death_is_surfaced_when_connection_closes() {
let endpoint = broker_that_closes_after_handshake().await;
let mut client = BrokerClient::connect(&endpoint, test_agent("parent"), "ignored")
.await
.expect("handshake completes");
assert!(
client.reader_alive(),
"reader should be alive immediately after connect"
);
let msg = tokio::time::timeout(Duration::from_secs(2), client.next_message())
.await
.expect("next_message() resolved instead of hanging");
assert!(msg.is_none(), "no message expected after the close");
let flagged_dead = tokio::time::timeout(Duration::from_secs(2), async {
while client.reader_alive() {
tokio::task::yield_now().await;
}
})
.await
.is_ok();
assert!(
flagged_dead,
"reader should be marked dead once the connection closed"
);
}
async fn broker_that_echoes_receipts_after_delay() -> String {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("bind");
let addr = listener.local_addr().expect("local_addr");
tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("accept");
let ws = accept_async(stream).await.expect("ws upgrade");
let (mut sink, mut source) = ws.split();
if let Some(Ok(Message::Text(_))) = source.next().await {
let _ = sink
.send(Message::text(BrokerFrame::Welcome.to_text()))
.await;
}
while let Some(Ok(Message::Text(txt))) = source.next().await {
if let Ok(ClientFrame::Deliver { message, .. }) = ClientFrame::from_text(&txt) {
let id = message.id.clone();
tokio::time::sleep(Duration::from_millis(50)).await;
let _ = sink
.send(Message::text(BrokerFrame::Delivered { id }.to_text()))
.await;
}
}
});
format!("ws://{addr}")
}
#[tokio::test]
async fn deliver_skips_a_stale_receipt_from_a_prior_timed_out_deliver() {
let endpoint = broker_that_echoes_receipts_after_delay().await;
let mut client = BrokerClient::connect(&endpoint, test_agent("parent"), "ignored")
.await
.expect("connect");
let msg_a = test_ask("a");
let id_a = msg_a.id.clone();
let res_a = client
.deliver_with_receipt_timeout("target", msg_a, Duration::from_millis(10))
.await;
assert!(res_a.is_err(), "deliver(A) times out before its receipt");
tokio::time::sleep(Duration::from_millis(120)).await;
let msg_b = test_ask("b");
let id_b = msg_b.id.clone();
assert_ne!(id_a, id_b);
let res_b = client
.deliver_with_receipt_timeout("target", msg_b, Duration::from_secs(5))
.await;
assert_eq!(
res_b.expect("deliver(B) succeeds"),
id_b,
"deliver(B) returns its own id, not A's stale receipt"
);
}
async fn broker_that_rejects_every_deliver(reason: &'static str) -> String {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("bind");
let addr = listener.local_addr().expect("local_addr");
tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("accept");
let ws = accept_async(stream).await.expect("ws upgrade");
let (mut sink, mut source) = ws.split();
if let Some(Ok(Message::Text(_))) = source.next().await {
let _ = sink
.send(Message::text(BrokerFrame::Welcome.to_text()))
.await;
}
while let Some(Ok(Message::Text(txt))) = source.next().await {
if let Ok(ClientFrame::Deliver { message, .. }) = ClientFrame::from_text(&txt) {
let _ = sink
.send(Message::text(
BrokerFrame::Error {
reason: reason.to_string(),
id: Some(message.id),
}
.to_text(),
))
.await;
}
}
});
format!("ws://{addr}")
}
#[tokio::test]
async fn deliver_returns_rejected_when_broker_sends_a_correlated_error() {
let endpoint =
broker_that_rejects_every_deliver("mailbox 'child' is full (2 pending messages)").await;
let mut client = BrokerClient::connect(&endpoint, test_agent("parent"), "ignored")
.await
.expect("handshake completes");
let started = std::time::Instant::now();
let outcome = tokio::time::timeout(
Duration::from_secs(2),
client.deliver("child", test_ask("parent")),
)
.await
.expect("deliver() resolves promptly instead of hanging out the receipt timeout");
match outcome {
Err(BrokerError::Rejected(reason)) => {
assert!(
reason.contains("full"),
"rejection reason should be the broker's verbatim message, got: {reason}"
);
}
other => panic!("expected BrokerError::Rejected, got {other:?}"),
}
assert!(
started.elapsed() < Duration::from_secs(1),
"the rejection must reach the caller fast, not via the receipt timeout, took {:?}",
started.elapsed(),
);
}
#[tokio::test]
async fn deliver_skips_a_stale_error_from_a_prior_timed_out_deliver() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("bind");
let addr = listener.local_addr().expect("local_addr");
let endpoint = format!("ws://{addr}");
tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("accept");
let ws = accept_async(stream).await.expect("ws upgrade");
let (mut sink, mut source) = ws.split();
if let Some(Ok(Message::Text(_))) = source.next().await {
let _ = sink
.send(Message::text(BrokerFrame::Welcome.to_text()))
.await;
}
if let Some(Ok(Message::Text(txt))) = source.next().await {
if let Ok(ClientFrame::Deliver { message, .. }) = ClientFrame::from_text(&txt) {
tokio::time::sleep(Duration::from_millis(50)).await;
let _ = sink
.send(Message::text(
BrokerFrame::Error {
reason: "stale rejection for A".into(),
id: Some(message.id),
}
.to_text(),
))
.await;
}
}
if let Some(Ok(Message::Text(txt))) = source.next().await {
if let Ok(ClientFrame::Deliver { message, .. }) = ClientFrame::from_text(&txt) {
let _ = sink
.send(Message::text(
BrokerFrame::Delivered { id: message.id }.to_text(),
))
.await;
}
}
while let Some(Ok(_)) = source.next().await {}
});
let mut client = BrokerClient::connect(&endpoint, test_agent("parent"), "ignored")
.await
.expect("connect");
let msg_a = test_ask("a");
let id_a = msg_a.id.clone();
let res_a = client
.deliver_with_receipt_timeout("target", msg_a, Duration::from_millis(10))
.await;
assert!(res_a.is_err(), "deliver(A) times out before its rejection");
tokio::time::sleep(Duration::from_millis(120)).await;
let msg_b = test_ask("b");
let id_b = msg_b.id.clone();
assert_ne!(id_a, id_b);
let res_b = client
.deliver_with_receipt_timeout("target", msg_b, Duration::from_secs(5))
.await;
assert_eq!(
res_b.expect("deliver(B) succeeds, not misrouted to A's stale rejection"),
id_b,
);
}
}