mod common;
use std::sync::Arc;
use atman_dsl::parse::parse_file;
use atman_runtime::event::FlowRunId;
use atman_runtime::session::Session;
use atman_runtime::tool::{Tool, ToolArgs, ToolCtx};
use atman_runtime::tools::agent_ctrl::{FlowInterject, FlowRegistry};
use atman_runtime::{Executor, Value, tools};
static HOME_TEST_LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(());
struct HomeGuard(Option<std::ffi::OsString>);
impl HomeGuard {
fn set(home: &std::path::Path) -> Self {
let old = std::env::var_os("HOME");
unsafe {
std::env::set_var("HOME", home);
}
Self(old)
}
}
impl Drop for HomeGuard {
fn drop(&mut self) {
unsafe {
match self.0.take() {
Some(old) => std::env::set_var("HOME", old),
None => std::env::remove_var("HOME"),
}
}
}
}
const SIMPLE_FLOW: &str = r#"flow t(n: Int) -> Int {
return n + 1
}
"#;
#[tokio::test]
async fn root_flow_run_registered_in_flow_registry() {
let file = parse_file(SIMPLE_FLOW).unwrap();
let session = Arc::new(Session::open_ephemeral());
let ex = Executor::with_events(session.sink().clone());
tools::register_tier_zero(&ex.tools);
let out = ex
.run_in_turn(
&file,
"t",
vec![("n".into(), Value::Int(4))],
None,
Some(session.clone()),
)
.await
.unwrap();
assert!(matches!(out, Value::Int(5)));
let root = session.flow_registry.lookup("root");
assert!(root.is_ok(), "root should be in flow_registry");
let root = root.unwrap();
assert_eq!(root.handle, "root");
assert_eq!(session.current_root(), Some("root".to_string()));
}
#[tokio::test]
async fn current_root_cleared_and_reset_across_turns() {
let file = parse_file(SIMPLE_FLOW).unwrap();
let session = Arc::new(Session::open_ephemeral());
let ex = Executor::with_events(session.sink().clone());
tools::register_tier_zero(&ex.tools);
ex.run_in_turn(
&file,
"t",
vec![("n".into(), Value::Int(1))],
None,
Some(session.clone()),
)
.await
.unwrap();
assert_eq!(session.current_root(), Some("root".to_string()));
ex.run_in_turn(
&file,
"t",
vec![("n".into(), Value::Int(2))],
None,
Some(session.clone()),
)
.await
.unwrap();
assert_eq!(session.current_root(), Some("root".to_string()));
assert!(session.flow_registry.lookup("root").is_ok());
}
#[tokio::test]
async fn no_session_no_root_registration() {
let file = parse_file(SIMPLE_FLOW).unwrap();
let ex = Executor::new();
tools::register_tier_zero(&ex.tools);
let out = ex
.run(&file, "t", vec![("n".into(), Value::Int(4))])
.await
.unwrap();
assert!(matches!(out, Value::Int(5)));
}
#[tokio::test]
async fn flow_interject_delivers_to_target_entry_channel() {
let registry = Arc::new(FlowRegistry::new());
let entry = registry.create_entry("sub_1".into(), "g".into(), "m".into(), FlowRunId::now());
let ctx = ToolCtx::new().with_flow_registry(registry);
let args = ToolArgs {
positional: vec![Value::Str("sub_1".into()), Value::Str("wake up".into())],
named: vec![],
};
FlowInterject.call(args, &ctx).await.unwrap();
let pending = entry.pending_injections.lock().unwrap();
assert_eq!(pending.len(), 1);
assert_eq!(pending[0].text, "wake up");
}
#[tokio::test]
async fn flow_interject_unknown_handle_errors() {
let registry = Arc::new(FlowRegistry::new());
let ctx = ToolCtx::new().with_flow_registry(registry);
let args = ToolArgs {
positional: vec![Value::Str("nope".into()), Value::Str("x".into())],
named: vec![],
};
let err = FlowInterject.call(args, &ctx).await.unwrap_err();
assert!(err.to_string().contains("not found"));
}
#[tokio::test]
async fn flow_interject_cancels_running_subagent_llm() {
let _home_lock = HOME_TEST_LOCK.lock().await;
let _registry =
common::ModelRegistryGuard::acquire(common::config([common::model_for_provider(
"mock", "mock", 100_000, None,
)]))
.await;
use atman_runtime::providers::mock::MockProvider;
use atman_runtime::tool::{Tool, ToolRegistry};
use atman_runtime::tools::agent_ctrl::{
AgentSpawn, FlowInterject, FlowRegistry, FlowRunStatus,
};
use std::io::Write;
let tmp = tempfile::tempdir().unwrap();
let commands_dir = tmp.path().join(".config").join("atman").join("commands");
std::fs::create_dir_all(&commands_dir).unwrap();
let flow_src = r#"flow describe() -> string { return "test" }
flow test_flow(goal: string) -> string {
reply = llm.call(
model: "mock",
prompt: goal,
context: "session",
)
return text_concat(reply)
}
"#;
let mut f = std::fs::File::create(commands_dir.join("test_interject.at")).unwrap();
f.write_all(flow_src.as_bytes()).unwrap();
let _home = HomeGuard::set(tmp.path());
let registry = Arc::new(FlowRegistry::new());
let providers = atman_runtime::provider::ProviderRegistry::new();
providers.register(Arc::new(
MockProvider::new("mock")
.with_fallback(atman_runtime::Value::Str("ok".into()))
.with_chunk_delay(std::time::Duration::from_secs(1)),
));
let tools = ToolRegistry::new();
atman_runtime::tools::register_tier_zero(&tools);
let (stream_tx, _) = tokio::sync::broadcast::channel::<atman_runtime::stream::StreamFrame>(256);
let ctx = ToolCtx::new()
.with_registry(Arc::new(tools))
.with_providers(Arc::new(providers))
.with_flow_registry(registry.clone())
.with_stream_tx(stream_tx);
let spawn_args = ToolArgs {
positional: vec![],
named: vec![
(
"flow".into(),
Value::Str("test_interject.at@test_flow".into()),
),
(
"arguments".into(),
Value::Struct(vec![("goal".into(), Value::Str("test goal".into()))]),
),
],
};
let result = AgentSpawn.call(spawn_args, &ctx).await.unwrap();
let handle = match result {
Value::Struct(fields) => fields
.iter()
.find(|(k, _)| k == "handle")
.and_then(|(_, v)| {
if let Value::Str(s) = v {
Some(s.clone())
} else {
None
}
})
.expect("expected handle"),
_ => panic!("expected struct with handle"),
};
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
let interject_args = ToolArgs {
positional: vec![Value::Str(handle.clone()), Value::Str("stop now".into())],
named: vec![("level".into(), Value::Str("l4_hard_stop".into()))],
};
let _ = FlowInterject.call(interject_args, &ctx).await.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
let entry = registry.lookup(&handle).unwrap();
let status = entry.status.lock().unwrap().clone();
let is_err = matches!(status, FlowRunStatus::Err { .. });
assert!(
is_err,
"sub-agent should be err after interjection, got: {:?}",
status
);
}
#[tokio::test]
async fn l1_nudge_text_appears_in_entry_messages() {
let _home_lock = HOME_TEST_LOCK.lock().await;
let _registry =
common::ModelRegistryGuard::acquire(common::config([common::model_for_provider(
"mock", "mock", 100_000, None,
)]))
.await;
use atman_runtime::Value;
use atman_runtime::providers::mock::MockProvider;
use atman_runtime::tool::{Tool, ToolRegistry};
use atman_runtime::tools::agent_ctrl::{AgentSpawn, FlowRegistry};
use std::io::Write;
let tmp = tempfile::tempdir().unwrap();
let commands_dir = tmp.path().join(".config").join("atman").join("commands");
std::fs::create_dir_all(&commands_dir).unwrap();
let flow_src = r#"flow describe() -> string { return "test" }
flow test_flow(goal: string) -> string {
session.push(message.user(goal))
reply = llm.call(model: "mock", context: "session")
return text_concat(reply)
}
"#;
let mut f = std::fs::File::create(commands_dir.join("test_l1.at")).unwrap();
f.write_all(flow_src.as_bytes()).unwrap();
let _home = HomeGuard::set(tmp.path());
let registry = Arc::new(FlowRegistry::new());
let providers = atman_runtime::provider::ProviderRegistry::new();
providers.register(Arc::new(
MockProvider::new("mock")
.with_fallback(Value::Str("done".into()))
.with_chunk_delay(std::time::Duration::from_millis(50)),
));
let tools = ToolRegistry::new();
atman_runtime::tools::register_tier_zero(&tools);
let (stream_tx, _) = tokio::sync::broadcast::channel::<atman_runtime::stream::StreamFrame>(256);
let ctx = ToolCtx::new()
.with_registry(Arc::new(tools))
.with_providers(Arc::new(providers))
.with_flow_registry(registry.clone())
.with_stream_tx(stream_tx);
let spawn_args = ToolArgs {
positional: vec![],
named: vec![
("flow".into(), Value::Str("test_l1.at@test_flow".into())),
(
"arguments".into(),
Value::Struct(vec![("goal".into(), Value::Str("hello".into()))]),
),
],
};
let result = AgentSpawn.call(spawn_args, &ctx).await.unwrap();
let handle = match result {
Value::Struct(fields) => fields
.iter()
.find(|(k, _)| k == "handle")
.and_then(|(_, v)| {
if let Value::Str(s) = v {
Some(s.clone())
} else {
None
}
})
.expect("expected handle"),
_ => panic!("expected struct"),
};
let inj = atman_runtime::injection::Injection::new_pending(
atman_runtime::event::TurnId::now(),
String::from("NUDGE: check config"),
);
let entry = registry.lookup(&handle).unwrap();
entry.pending_injections.lock().unwrap().push(inj);
entry.injection_notify.notify_one();
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
let entry = registry.lookup(&handle).unwrap();
let status = entry.status.lock().unwrap().clone();
let pending_count = entry.pending_injections.lock().unwrap().len();
let msgs = entry.messages.lock().unwrap();
let has_nudge = msgs
.iter()
.any(|m| m.text_concat().contains("NUDGE: check config"));
assert!(
has_nudge,
"L1 nudge text should be in entry.messages, got: {:?}, pending: {}, status: {:?}",
msgs.iter().map(|m| m.text_concat()).collect::<Vec<_>>(),
pending_count,
status
);
let pending = entry.pending_injections.lock().unwrap();
assert!(
pending.is_empty(),
"pending_injections should be empty after drain"
);
}