use std::sync::Arc;
use agent_base::{Content, Tool, ToolContext};
use crate::multi_agent::capability::ChildToolCapability;
use crate::multi_agent::child_config::ChildConfig;
use crate::multi_agent::config::{AgentAutonomy, ControlConfig, MultiAgentConfig};
use crate::multi_agent::runtime::MultiAgentRuntime;
use crate::multi_agent::write_gate::WorkspaceWriteGate;
use super::*;
struct StubWrite;
#[async_trait::async_trait]
impl Tool for StubWrite {
fn name(&self) -> &'static str {
"write_file"
}
fn description(&self) -> &'static str {
"fixture"
}
fn schema(&self) -> serde_json::Value {
serde_json::json!({"type": "object", "properties": {"path": {"type": "string"}}})
}
async fn call(
&self,
_args: &serde_json::Value,
_ctx: &ToolContext,
) -> agent_base::AgentResult<Vec<Content>> {
Ok(vec![Content::text("ok")])
}
}
fn runtime_with_gate(gate_enabled: bool) -> Arc<MultiAgentRuntime> {
runtime_with_gate_and_tools(
gate_enabled,
Arc::new(StreamingStub),
vec![Arc::new(StubWrite)],
)
}
fn runtime_with_gate_and_tools(
gate_enabled: bool,
client: Arc<dyn agent_base::llm_trait::LlmProvider>,
tools: Vec<Arc<dyn Tool>>,
) -> Arc<MultiAgentRuntime> {
let config = MultiAgentConfig {
allow_child_write: true,
child_excluded_tools: vec![],
control: ControlConfig {
autonomy: AgentAutonomy::Auto,
child_write_gate: gate_enabled,
..ControlConfig::default()
},
..MultiAgentConfig::enabled()
};
Arc::new(MultiAgentRuntime::new(
config,
client,
tools,
tokio_util::sync::CancellationToken::new(),
None,
agent_base::Language::En,
None,
None,
))
}
#[tokio::test(flavor = "multi_thread")]
async fn gated_tool_wraps_write_file_on_spawned_children() {
let ma = runtime_with_gate(true);
let config = ChildConfig {
system_prompt: Some("p".into()),
..Default::default()
};
let (child_a, _reg, _res) = ma
.build_child_runtime_with_config(&config, true, Some(&ChildToolCapability::Write), "root/a")
.await
.unwrap();
let (child_b, _reg, _res) = ma
.build_child_runtime_with_config(&config, true, Some(&ChildToolCapability::Write), "root/b")
.await
.unwrap();
let _ = (child_a, child_b);
let shared = Arc::new(WorkspaceWriteGate::new());
shared
.try_claim(std::path::Path::new("x.rs"), "root/a")
.unwrap();
let err = shared
.try_claim(std::path::Path::new("x.rs"), "root/b")
.unwrap_err();
assert!(err.contains("root/a"));
}
#[tokio::test(flavor = "multi_thread")]
async fn gate_disabled_still_spawns_write_children() {
let ma = runtime_with_gate(false);
let config = ChildConfig {
system_prompt: Some("p".into()),
..Default::default()
};
let (_child, registered, _res) = ma
.build_child_runtime_with_config(&config, true, Some(&ChildToolCapability::Write), "root/a")
.await
.unwrap();
assert!(registered.contains("write_file"));
}
#[tokio::test(flavor = "multi_thread")]
async fn close_releases_gate_claims() {
let ma = runtime_with_gate(true);
let path = ma
.spawn_child_with_history(
"w",
"p".to_string(),
true,
None,
None,
None, &agent_base::SessionId::new(1),
)
.await
.unwrap();
let path = path.agent_path;
ma.write_gate_for_test()
.try_claim(std::path::Path::new("x.rs"), &path)
.unwrap();
ma.close_agent(&path).unwrap();
poll_until("gate claims released", || {
ma.write_gate_for_test()
.try_claim(std::path::Path::new("x.rs"), "root/probe")
.is_ok()
})
.await;
}
struct SlowWrite(std::time::Duration);
#[async_trait::async_trait]
impl Tool for SlowWrite {
fn name(&self) -> &'static str {
"write_file"
}
fn description(&self) -> &'static str {
"fixture (slow)"
}
fn schema(&self) -> serde_json::Value {
serde_json::json!({"type": "object", "properties": {"path": {"type": "string"}}})
}
async fn call(
&self,
_args: &serde_json::Value,
_ctx: &ToolContext,
) -> agent_base::AgentResult<Vec<Content>> {
tokio::time::sleep(self.0).await;
Ok(vec![Content::text("ok")])
}
}
#[tokio::test(flavor = "multi_thread")]
async fn task_completion_releases_gate_claims() {
let ma = runtime_with_gate_and_tools(
true,
Arc::new(ToolCallOnceStub::new("write_file", "{\"path\":\"x.rs\"}")),
vec![Arc::new(SlowWrite(std::time::Duration::from_millis(500)))],
);
let echo = ma
.spawn_child_with_history(
"wg",
"p".to_string(),
true,
None,
None,
Some(ChildToolCapability::Write),
&agent_base::SessionId::new(1),
)
.await
.unwrap();
assert_eq!(echo.agent_path, "root/wg");
ma.send_task("root/wg", "write it".to_string(), false)
.unwrap();
poll_until("child claims x.rs via GatedTool", || {
ma.write_gate_for_test()
.holder_of(std::path::Path::new("x.rs"))
.as_deref()
== Some("root/wg")
})
.await;
poll_until("child task done", || {
ma.list_agents()
.iter()
.any(|a| a.agent_path == "root/wg" && a.status == "done")
})
.await;
assert!(
ma.write_gate_for_test()
.try_claim(std::path::Path::new("x.rs"), "root/probe")
.is_ok(),
"claims must be released when the task ends, not held until close"
);
}