use std::collections::HashMap;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::Duration;
use bamboo_subagent::{AgentRef, InboxKind, InboxMessage, MsgId};
use chrono::Utc;
use futures_util::SinkExt;
use serde_json::Value;
use tokio::sync::{mpsc, oneshot, Mutex};
use tokio::task::JoinHandle;
use tokio_tungstenite::tungstenite::Message;
use crate::client::WsSink;
use crate::error::{BrokerError, BrokerResult};
use crate::proto::ClientFrame;
const DELIVER_RECEIPT_TIMEOUT: Duration = Duration::from_secs(30);
type Pending = Arc<Mutex<HashMap<MsgId, oneshot::Sender<InboxMessage>>>>;
pub struct MultiplexedClient {
sink: Mutex<WsSink>,
pending: Pending,
delivered: Mutex<mpsc::UnboundedReceiver<MsgId>>,
reader_alive: Arc<AtomicBool>,
_router: JoinHandle<()>,
me: AgentRef,
}
impl MultiplexedClient {
pub(crate) fn spawn(
sink: WsSink,
mut messages: mpsc::UnboundedReceiver<InboxMessage>,
delivered: mpsc::UnboundedReceiver<MsgId>,
reader_alive: Arc<AtomicBool>,
me: AgentRef,
) -> Self {
let pending: Pending = Arc::new(Mutex::new(HashMap::new()));
let routed = pending.clone();
let router = tokio::spawn(async move {
while let Some(msg) = messages.recv().await {
if let Some(cid) = msg.correlation_id.clone() {
if let Some(tx) = routed.lock().await.remove(&cid) {
let _ = tx.send(msg);
continue;
}
}
tracing::debug!("mcp mux: dropping uncorrelated/late reply");
}
routed.lock().await.clear();
});
Self {
sink: Mutex::new(sink),
pending,
delivered: Mutex::new(delivered),
reader_alive,
_router: router,
me,
}
}
pub fn reader_alive(&self) -> bool {
self.reader_alive.load(Ordering::SeqCst)
}
pub async fn request(
&self,
target: &str,
kind: InboxKind,
body: Value,
timeout: Duration,
) -> BrokerResult<Value> {
let msg = InboxMessage {
id: MsgId::new(),
from: self.me.clone(),
kind,
body,
created_at: Utc::now(),
correlation_id: None,
};
let qid = msg.id.clone();
let (tx, rx) = oneshot::channel();
self.pending.lock().await.insert(qid.clone(), tx);
if let Err(e) = self.deliver(target, msg).await {
self.pending.lock().await.remove(&qid);
return Err(e);
}
match tokio::time::timeout(timeout, rx).await {
Ok(Ok(reply)) => Ok(reply.body),
Ok(Err(_)) => {
self.pending.lock().await.remove(&qid);
Err(BrokerError::Transport(
"connection closed before reply".into(),
))
}
Err(_) => {
self.pending.lock().await.remove(&qid);
let _ = self.cancel(target, &qid).await;
Err(BrokerError::Transport(format!(
"request to '{target}' timed out after {timeout:?}"
)))
}
}
}
async fn deliver(&self, to: &str, message: InboxMessage) -> BrokerResult<MsgId> {
let id = message.id.clone();
let frame = ClientFrame::Deliver {
to: to.into(),
message,
};
{
let mut sink = self.sink.lock().await;
sink.send(Message::text(frame.to_text()))
.await
.map_err(|e| BrokerError::Transport(format!("ws send: {e}")))?;
}
let mut delivered = self.delivered.lock().await;
match tokio::time::timeout(DELIVER_RECEIPT_TIMEOUT, delivered.recv()).await {
Ok(Some(_)) => Ok(id),
Ok(None) => Err(BrokerError::Transport(
"connection closed before delivery receipt".into(),
)),
Err(_) => Err(BrokerError::Transport(
"timed out waiting for delivery receipt from broker".into(),
)),
}
}
async fn cancel(&self, to: &str, correlation_id: &MsgId) -> BrokerResult<()> {
let frame = ClientFrame::Cancel {
to: to.into(),
correlation_id: correlation_id.clone(),
};
let mut sink = self.sink.lock().await;
sink.send(Message::text(frame.to_text()))
.await
.map_err(|e| BrokerError::Transport(format!("ws send: {e}")))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::client::BrokerClient;
use crate::core::BrokerCore;
use crate::server::BrokerServer;
use std::sync::Arc as StdArc;
use tokio::net::TcpListener;
const TOKEN: &str = "t";
async fn start() -> (String, tempfile::TempDir) {
let dir = tempfile::tempdir().unwrap();
let core = StdArc::new(BrokerCore::new(dir.path()));
let server = StdArc::new(BrokerServer::new(core, TOKEN));
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
let _ = server.serve(listener).await;
});
(format!("ws://{addr}"), dir)
}
fn agent(id: &str) -> AgentRef {
AgentRef {
session_id: id.into(),
role: None,
}
}
async fn mux(endpoint: &str, id: &str) -> MultiplexedClient {
let mut c = BrokerClient::connect(endpoint, agent(id), TOKEN)
.await
.unwrap();
c.subscribe().await.unwrap();
c.into_multiplexed(agent(id))
}
#[tokio::test]
async fn concurrent_requests_route_replies_by_correlation_id() {
let (endpoint, _dir) = start().await;
let client = mux(&endpoint, "worker").await;
let mut responder = BrokerClient::connect(&endpoint, agent("orch"), TOKEN)
.await
.unwrap();
responder.subscribe().await.unwrap();
tokio::spawn(async move {
let m1 = responder.next_message().await.unwrap();
let m2 = responder.next_message().await.unwrap();
for m in [m2, m1] {
let reply = InboxMessage {
id: MsgId::new(),
from: agent("orch"),
kind: InboxKind::McpReply,
body: m.body.clone(), created_at: Utc::now(),
correlation_id: Some(m.id.clone()),
};
responder.deliver("worker", reply).await.unwrap();
}
});
let t = Duration::from_secs(5);
let (r1, r2) = tokio::join!(
client.request(
"orch",
InboxKind::McpRequest,
serde_json::json!({ "n": 1 }),
t
),
client.request(
"orch",
InboxKind::McpRequest,
serde_json::json!({ "n": 2 }),
t
),
);
assert_eq!(r1.unwrap(), serde_json::json!({ "n": 1 }));
assert_eq!(r2.unwrap(), serde_json::json!({ "n": 2 }));
}
#[tokio::test]
async fn request_times_out_when_no_reply() {
let (endpoint, _dir) = start().await;
let client = mux(&endpoint, "worker2").await;
let err = client
.request(
"nobody",
InboxKind::McpRequest,
serde_json::json!({}),
Duration::from_millis(200),
)
.await;
assert!(err.is_err(), "request to a non-responder times out");
assert!(
client.pending.lock().await.is_empty(),
"the timed-out waiter was unregistered"
);
}
}