use std::collections::BTreeSet;
use std::sync::Arc;
use std::time::Duration;
use agent_base::AgentError;
use super::path::AgentPath;
use super::runtime::{MultiAgentRuntime, WaitResult};
#[derive(Clone)]
pub struct ChildHandle {
runtime: Arc<MultiAgentRuntime>,
path: AgentPath,
spawned_tools: BTreeSet<String>,
}
impl ChildHandle {
pub(crate) fn new(
runtime: Arc<MultiAgentRuntime>,
path: AgentPath,
spawned_tools: BTreeSet<String>,
) -> Self {
Self {
runtime,
path,
spawned_tools,
}
}
pub fn agent_path(&self) -> String {
self.path.to_string()
}
pub fn spawned_tools(&self) -> &BTreeSet<String> {
&self.spawned_tools
}
pub fn send(&self, message: impl Into<String>) -> Result<bool, AgentError> {
self.runtime
.send_message(&self.agent_path(), message.into())
.map_err(AgentError::ConfigError)
}
pub fn task(&self, task: impl Into<String>) -> Result<bool, AgentError> {
self.runtime
.send_task(&self.agent_path(), task.into(), false)
.map_err(AgentError::ConfigError)
}
pub async fn wait(&self, timeout: Duration) -> ChildOutcome {
let wr = self
.runtime
.wait_for_result(Some(&self.agent_path()), timeout.as_millis() as u64)
.await;
ChildOutcome::from_wait_result(wr)
}
pub fn close(&self) -> Result<bool, AgentError> {
self.runtime
.close_agent(&self.agent_path())
.map(|r| r.closed)
.map_err(AgentError::ConfigError)
}
pub fn into_guard(self) -> ChildGuard {
ChildGuard {
handle: Some(self),
disarmed: false,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ChildOutcome {
Ok {
text: Option<String>,
has_more: bool,
denied_tools: Vec<String>,
},
Failed { text: Option<String> },
Closed,
Timeout,
}
impl ChildOutcome {
pub(crate) fn from_wait_result(wr: WaitResult) -> Self {
match wr.status.as_str() {
"ok" => Self::Ok {
text: wr.result,
has_more: wr.has_more,
denied_tools: wr.denied_tools,
},
"error" => Self::Failed { text: wr.result },
"closed" => Self::Closed,
_ => Self::Timeout,
}
}
}
pub struct ChildGuard {
handle: Option<ChildHandle>,
disarmed: bool,
}
impl ChildGuard {
pub fn handle(&self) -> &ChildHandle {
self.handle
.as_ref()
.expect("ChildGuard handle is present until into_handle")
}
pub fn into_handle(mut self) -> ChildHandle {
self.disarmed = true;
self.handle
.take()
.expect("handle is present before Drop runs")
}
pub fn disarm(&mut self) {
self.disarmed = true;
}
}
impl Drop for ChildGuard {
fn drop(&mut self) {
if !self.disarmed
&& let Some(handle) = &self.handle
{
let _ = handle.close();
}
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use agent_base::llm_trait::LlmProvider;
use super::*;
use crate::multi_agent::child_config::ChildConfig;
use crate::multi_agent::config::MultiAgentConfig;
use crate::multi_agent::path::AgentPath;
use crate::multi_agent::runtime::WaitResult;
struct EchoLlm;
#[async_trait::async_trait]
impl LlmProvider for EchoLlm {
async fn stream(
&self,
_request: agent_base::llm_trait::ChatRequest,
) -> Result<agent_base::llm_trait::ChatStream, agent_base::llm_trait::LlmError> {
Ok(agent_base::llm_trait::ChatStream::new(Box::pin(
futures_util::stream::iter(vec![
Ok(agent_base::StreamChunk::Text("done".to_string())),
Ok(agent_base::StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
}),
]),
)))
}
async fn chat(
&self,
_request: agent_base::llm_trait::ChatRequest,
) -> Result<agent_base::llm_trait::ChatResponse, agent_base::llm_trait::LlmError> {
unreachable!("unused")
}
fn capabilities(&self) -> agent_base::llm_trait::Capabilities {
agent_base::llm_trait::Capabilities::default()
}
fn info(&self) -> agent_base::llm_trait::ProviderInfo {
agent_base::llm_trait::ProviderInfo {
name: "echo".into(),
model: "echo".into(),
version: None,
}
}
}
fn runtime() -> Arc<MultiAgentRuntime> {
Arc::new(MultiAgentRuntime::new(
MultiAgentConfig::enabled(),
Arc::new(EchoLlm),
vec![],
tokio_util::sync::CancellationToken::new(),
None,
agent_base::Language::En,
None,
None,
))
}
fn wr(status: &str) -> WaitResult {
WaitResult {
status: status.into(),
result: Some("x".into()),
agent_path: Some("root/w".into()),
has_more: false,
denied_tools: vec!["read_file".into()],
}
}
#[test]
fn outcome_maps_all_four_statuses() {
assert_eq!(
ChildOutcome::from_wait_result(wr("ok")),
ChildOutcome::Ok {
text: Some("x".into()),
has_more: false,
denied_tools: vec!["read_file".into()],
}
);
assert_eq!(
ChildOutcome::from_wait_result(wr("error")),
ChildOutcome::Failed {
text: Some("x".into())
}
);
assert_eq!(
ChildOutcome::from_wait_result(wr("closed")),
ChildOutcome::Closed
);
assert_eq!(
ChildOutcome::from_wait_result(wr("timeout")),
ChildOutcome::Timeout
);
}
#[tokio::test(flavor = "multi_thread")]
async fn handle_send_task_wait_close_round_trip() {
let ma = runtime();
let spawned = ma
.spawn_with_config(
"w".to_string(),
ChildConfig {
system_prompt: Some("prompt".into()),
..Default::default()
},
)
.await
.unwrap();
let handle = ChildHandle::new(
Arc::clone(&ma),
spawned.agent_path().clone(),
spawned.spawned_tools().clone(),
);
assert_eq!(handle.agent_path(), "root/w");
assert!(handle.send("heads up").unwrap());
assert!(handle.task("do it").unwrap());
let outcome = handle.wait(Duration::from_secs(3)).await;
assert!(
matches!(outcome, ChildOutcome::Ok { text: Some(ref t), .. } if t == "done"),
"got {outcome:?}"
);
assert!(handle.close().unwrap());
assert!(!handle.close().unwrap(), "idempotent close");
let outcome = handle.wait(Duration::from_millis(500)).await;
assert_eq!(outcome, ChildOutcome::Closed);
}
#[tokio::test(flavor = "multi_thread")]
async fn guard_closes_on_drop_unless_disarmed() {
let ma = runtime();
let spawned = ma
.spawn_with_config(
"g".to_string(),
ChildConfig {
system_prompt: Some("prompt".into()),
..Default::default()
},
)
.await
.unwrap();
let handle = ChildHandle::new(
Arc::clone(&ma),
spawned.agent_path().clone(),
spawned.spawned_tools().clone(),
);
let guard = handle.clone().into_guard();
drop(guard);
for _ in 0..100 {
if ma.registry().lock().unwrap().count() == 0 {
break;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
assert_eq!(ma.registry().lock().unwrap().count(), 0);
let spawned2 = ma
.spawn_with_config(
"h".to_string(),
ChildConfig {
system_prompt: Some("prompt".into()),
..Default::default()
},
)
.await
.unwrap();
let handle2 = ChildHandle::new(
Arc::clone(&ma),
spawned2.agent_path().clone(),
spawned2.spawned_tools().clone(),
);
let mut guard2 = handle2.clone().into_guard();
guard2.disarm();
drop(guard2);
assert!(
ma.registry()
.lock()
.unwrap()
.contains(&AgentPath::root().join("h"))
);
}
}