use std::collections::BTreeSet;
use std::sync::Arc;
use agent_base::{AgentError, SessionId};
use super::child::{ChildGuard, ChildHandle};
use super::child_config::ChildConfig;
use super::preset::ChildPreset;
use super::runtime::MultiAgentRuntime;
#[derive(Default)]
struct ChildFieldSet {
system_prompt: bool,
tool_names: bool,
max_turns: bool,
context_window: bool,
full_permission: bool,
}
pub struct ChildBuilder {
runtime: Arc<MultiAgentRuntime>,
draft: ChildConfig,
explicit: ChildFieldSet,
fork: Option<(String, SessionId)>,
}
impl ChildBuilder {
pub(crate) fn new(runtime: Arc<MultiAgentRuntime>) -> Self {
Self {
runtime,
draft: ChildConfig::default(),
explicit: ChildFieldSet::default(),
fork: None,
}
}
pub fn system_prompt(mut self, prompt: impl Into<String>) -> Self {
self.draft.system_prompt = Some(prompt.into());
self.explicit.system_prompt = true;
self
}
pub fn tool_names(mut self, names: impl IntoIterator<Item = impl Into<String>>) -> Self {
self.draft.tool_names = Some(names.into_iter().map(Into::into).collect());
self.explicit.tool_names = true;
self
}
pub fn add_tool_name(mut self, name: impl Into<String>) -> Self {
self.draft
.tool_names
.get_or_insert_with(BTreeSet::new)
.insert(name.into());
self.explicit.tool_names = true;
self
}
pub fn max_turns(mut self, max: u32) -> Self {
self.draft.max_turns = Some(max);
self.explicit.max_turns = true;
self
}
pub fn context_window(mut self, tokens: usize) -> Self {
self.draft.context_window = Some(tokens);
self.explicit.context_window = true;
self
}
pub fn full_permission(mut self, full: bool) -> Self {
self.draft.full_permission = Some(full);
self.explicit.full_permission = true;
self
}
pub fn preset(mut self, preset: &ChildPreset) -> Self {
let p = &preset.config;
if !self.explicit.system_prompt && p.system_prompt.is_some() {
self.draft.system_prompt = p.system_prompt.clone();
}
if !self.explicit.tool_names && p.tool_names.is_some() {
self.draft.tool_names = p.tool_names.clone();
}
if !self.explicit.max_turns && p.max_turns.is_some() {
self.draft.max_turns = p.max_turns;
}
if !self.explicit.context_window && p.context_window.is_some() {
self.draft.context_window = p.context_window;
}
if !self.explicit.full_permission && p.full_permission.is_some() {
self.draft.full_permission = p.full_permission;
}
self
}
pub fn fork_history(self, mode: impl Into<String>, parent_session: SessionId) -> Self {
Self {
fork: Some((mode.into(), parent_session)),
..self
}
}
pub async fn spawn(self, name: impl Into<String>) -> Result<ChildHandle, AgentError> {
let name = name.into();
if self.draft.system_prompt.as_deref().unwrap_or("").is_empty() {
return Err(AgentError::ConfigError(
"ChildConfig.system_prompt is required (set it directly or use a preset)".into(),
));
}
let parent_messages = match &self.fork {
Some((mode, sid)) => {
self.runtime
.resolve_fork_history(Some(mode.clone()), sid)
.await
}
None => Vec::new(),
};
let spawned = self
.runtime
.spawn_with_config_forked(name, self.draft, parent_messages, None)
.await?;
Ok(ChildHandle::new(
Arc::clone(&self.runtime),
spawned.agent_path().clone(),
spawned.spawned_tools().clone(),
))
}
pub async fn spawn_guarded(self, name: impl Into<String>) -> Result<ChildGuard, AgentError> {
Ok(self.spawn(name).await?.into_guard())
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::time::Duration;
use agent_base::llm_trait::LlmProvider;
use super::*;
use crate::multi_agent::child::ChildOutcome;
use crate::multi_agent::config::MultiAgentConfig;
use crate::multi_agent::preset::tool;
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 full_preset() -> ChildPreset {
ChildPreset::custom(
"full",
"d",
ChildConfig {
system_prompt: Some("preset prompt".into()),
tool_names: Some([tool::READ_FILE.to_string()].into_iter().collect()),
max_turns: Some(32),
context_window: Some(4096),
full_permission: Some(true),
..Default::default()
},
)
}
fn empty_preset() -> ChildPreset {
ChildPreset::custom("empty", "d", ChildConfig::default())
}
fn merge_cells<T>(
field: &str,
user_val: T,
preset_val: T,
set: impl Fn(ChildBuilder, T) -> ChildBuilder,
get: impl Fn(&ChildBuilder) -> Option<T>,
) where
T: PartialEq + Clone + std::fmt::Debug,
{
let ma = runtime();
let full = full_preset();
let empty = empty_preset();
let b = set(ChildBuilder::new(Arc::clone(&ma)), user_val.clone()).preset(&full);
assert_eq!(get(&b), Some(user_val.clone()), "{field}: cell 1");
let b = set(ChildBuilder::new(Arc::clone(&ma)), user_val.clone()).preset(&empty);
assert_eq!(get(&b), Some(user_val.clone()), "{field}: cell 2");
let b = ChildBuilder::new(Arc::clone(&ma)).preset(&full);
assert_eq!(get(&b), Some(preset_val), "{field}: cell 3");
let b = ChildBuilder::new(Arc::clone(&ma)).preset(&empty);
assert_eq!(get(&b), None, "{field}: cell 4");
let b = set(ChildBuilder::new(ma).preset(&full), user_val.clone());
assert_eq!(get(&b), Some(user_val), "{field}: cell 5");
}
#[test]
fn merge_matrix_system_prompt() {
merge_cells(
"system_prompt",
"user prompt".to_string(),
"preset prompt".to_string(),
|b, v| b.system_prompt(v),
|b| b.draft.system_prompt.clone(),
);
}
#[test]
fn merge_matrix_tool_names() {
merge_cells(
"tool_names",
[tool::WRITE_FILE.to_string()].into_iter().collect(),
[tool::READ_FILE.to_string()].into_iter().collect(),
|b, v| b.tool_names(v),
|b| b.draft.tool_names.clone(),
);
}
#[test]
fn merge_matrix_max_turns() {
merge_cells(
"max_turns",
8u32,
32u32,
|b, v| b.max_turns(v),
|b| b.draft.max_turns,
);
}
#[test]
fn merge_matrix_context_window() {
merge_cells(
"context_window",
1024usize,
4096usize,
|b, v| b.context_window(v),
|b| b.draft.context_window,
);
}
#[test]
fn merge_matrix_full_permission() {
merge_cells(
"full_permission",
false,
true,
|b, v| b.full_permission(v),
|b| b.draft.full_permission,
);
}
#[test]
fn add_tool_name_accumulates_and_blocks_preset() {
let b = ChildBuilder::new(runtime())
.add_tool_name("a")
.add_tool_name("b")
.preset(&full_preset());
assert_eq!(
b.draft.tool_names,
Some(["a".to_string(), "b".to_string()].into_iter().collect())
);
}
#[test]
fn add_tool_name_after_preset_appends_to_preset_set() {
let b = ChildBuilder::new(runtime())
.preset(&full_preset())
.add_tool_name("extra");
assert_eq!(
b.draft.tool_names,
Some(
[tool::READ_FILE.to_string(), "extra".to_string()]
.into_iter()
.collect()
)
);
}
#[tokio::test(flavor = "multi_thread")]
async fn spawn_requires_system_prompt() {
let err = runtime()
.child()
.max_turns(4)
.spawn("worker")
.await
.map(|_h| ())
.expect_err("no prompt must fail");
assert!(
matches!(&err, AgentError::ConfigError(s)
if s == "ChildConfig.system_prompt is required (set it directly or use a preset)"),
"got {err:?}"
);
}
#[tokio::test(flavor = "multi_thread")]
async fn spawn_with_preset_prompt_succeeds() {
let ma = runtime();
let handle = ma
.child()
.preset(&ChildPreset::custom(
"p",
"d",
ChildConfig {
system_prompt: Some("prompt".into()),
..Default::default()
},
))
.spawn("w")
.await
.expect("preset prompt spawns");
assert_eq!(handle.agent_path(), "root/w");
assert!(handle.task("go").unwrap());
let outcome = handle.wait(Duration::from_secs(3)).await;
assert!(matches!(outcome, ChildOutcome::Ok { .. }));
assert!(handle.close().unwrap());
}
#[tokio::test(flavor = "multi_thread")]
async fn spawn_guarded_closes_on_drop() {
let ma = runtime();
let guard = ma
.child()
.system_prompt("prompt")
.spawn_guarded("g")
.await
.expect("spawn_guarded");
assert_eq!(guard.handle().agent_path(), "root/g");
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);
}
}