use std::time::Duration;
use bamboo_subagent::{AgentRef, AskBody, AskMode, InboxKind, InboxMessage, MsgId, ReplyBody};
use chrono::Utc;
use crate::client::BrokerClient;
use crate::error::{BrokerError, BrokerResult};
pub async fn ask_agent(
endpoint: &str,
me: AgentRef,
token: &str,
target: &str,
question: &str,
mode: AskMode,
timeout: Duration,
) -> BrokerResult<String> {
let mut client = BrokerClient::connect(endpoint, me.clone(), token).await?;
client.subscribe().await?;
ask_over(&mut client, &me, target, question, mode, timeout).await
}
pub async fn ask_over(
client: &mut BrokerClient,
me: &AgentRef,
target: &str,
question: &str,
mode: AskMode,
timeout: Duration,
) -> BrokerResult<String> {
let msg = InboxMessage {
id: MsgId::new(),
from: me.clone(),
kind: InboxKind::Ask,
body: serde_json::to_value(AskBody {
question: question.to_string(),
mode,
})
.expect("AskBody serializes"),
created_at: Utc::now(),
correlation_id: None,
};
let qid = msg.id.clone();
client.deliver(target, msg).await?;
loop {
match tokio::time::timeout(timeout, client.next_message()).await {
Ok(Some(reply)) if reply.correlation_id.as_ref() == Some(&qid) => {
let body: ReplyBody = serde_json::from_value(reply.body)
.map_err(|e| BrokerError::Protocol(format!("bad reply body: {e}")))?;
return Ok(body.answer);
}
Ok(Some(_)) => continue, Ok(None) => {
return Err(BrokerError::Transport(
"connection closed before reply".into(),
))
}
Err(_) => {
let _ = client.cancel(target, &qid).await;
return Err(BrokerError::Transport(format!(
"ask to '{target}' timed out after {timeout:?}"
)));
}
}
}
}
pub async fn request_over(
client: &mut BrokerClient,
me: &AgentRef,
target: &str,
kind: InboxKind,
body: serde_json::Value,
timeout: Duration,
) -> BrokerResult<serde_json::Value> {
let msg = InboxMessage {
id: MsgId::new(),
from: me.clone(),
kind,
body,
created_at: Utc::now(),
correlation_id: None,
};
let qid = msg.id.clone();
client.deliver(target, msg).await?;
loop {
match tokio::time::timeout(timeout, client.next_message()).await {
Ok(Some(reply)) if reply.correlation_id.as_ref() == Some(&qid) => return Ok(reply.body),
Ok(Some(_)) => continue,
Ok(None) => {
return Err(BrokerError::Transport(
"connection closed before reply".into(),
))
}
Err(_) => {
let _ = client.cancel(target, &qid).await;
return Err(BrokerError::Transport(format!(
"request to '{target}' timed out after {timeout:?}"
)));
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::BrokerCore;
use crate::serve::serve_executor;
use crate::server::BrokerServer;
use std::sync::Arc;
use tokio::net::TcpListener;
#[tokio::test]
async fn ask_agent_round_trip_against_echo_executor() {
let dir = tempfile::tempdir().unwrap();
let core = Arc::new(BrokerCore::new(dir.path()));
let server = Arc::new(BrokerServer::new(core, "t"));
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;
});
let endpoint = format!("ws://{addr}");
let ep = endpoint.clone();
tokio::spawn(async move {
let _ = serve_executor(
&ep,
AgentRef {
session_id: "w".into(),
role: None,
},
"t",
Arc::new(bamboo_subagent::EchoExecutor),
)
.await;
});
let answer = ask_agent(
&endpoint,
AgentRef {
session_id: "orch".into(),
role: None,
},
"t",
"w",
"ping pong",
AskMode::Query,
Duration::from_secs(5),
)
.await
.expect("ask returns an answer");
assert_eq!(answer, "echo: ping pong");
}
#[tokio::test]
async fn ask_timeout_cancels_the_worker_run() {
use bamboo_subagent::{ChildExecutor, ChildOutcome, EventSink, RunSpec, SteerInbox};
use std::sync::atomic::{AtomicBool, Ordering};
use tokio_util::sync::CancellationToken;
let dir = tempfile::tempdir().unwrap();
let core = Arc::new(BrokerCore::new(dir.path()));
let server = Arc::new(BrokerServer::new(core, "t"));
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;
});
let endpoint = format!("ws://{addr}");
let was_cancelled = Arc::new(AtomicBool::new(false));
struct ProbeOrPark(Arc<AtomicBool>);
#[async_trait::async_trait]
impl ChildExecutor for ProbeOrPark {
async fn run(
&self,
spec: RunSpec,
_events: EventSink,
_steer: SteerInbox,
cancel: CancellationToken,
) -> ChildOutcome {
if spec.assignment.contains("park") {
cancel.cancelled().await;
self.0.store(true, Ordering::SeqCst);
ChildOutcome::cancelled()
} else {
ChildOutcome::completed("ready")
}
}
}
let ep = endpoint.clone();
let flag = was_cancelled.clone();
tokio::spawn(async move {
let _ = serve_executor(
&ep,
AgentRef {
session_id: "w".into(),
role: None,
},
"t",
Arc::new(ProbeOrPark(flag)),
)
.await;
});
let ready = ask_agent(
&endpoint,
AgentRef {
session_id: "probe".into(),
role: None,
},
"t",
"w",
"ping",
AskMode::Query,
Duration::from_secs(5),
)
.await;
assert_eq!(ready.expect("probe answered"), "ready");
let result = ask_agent(
&endpoint,
AgentRef {
session_id: "orch".into(),
role: None,
},
"t",
"w",
"park on slow work",
AskMode::Query,
Duration::from_millis(300),
)
.await;
assert!(result.is_err(), "ask times out (the worker never replies)");
let mut observed = false;
for _ in 0..100 {
if was_cancelled.load(Ordering::SeqCst) {
observed = true;
break;
}
tokio::time::sleep(Duration::from_millis(20)).await;
}
assert!(
observed,
"ask_agent timeout cancelled the worker's in-flight run"
);
}
}