use std::sync::Arc;
use std::thread::ThreadId;
use super::{prepare, scope_placement, spawn_into, SubtaskPlacement};
use crate::call_budget::{charge_mcp_call, install_mcp_call_budget, mcp_calls_spent};
use crate::stdlib::pool::PoolRegistry;
struct Observation {
thread: ThreadId,
session: Option<String>,
event_log: Option<usize>,
}
const BRANCHES: usize = 1;
struct Measured {
parent_thread: ThreadId,
parent_session: Option<String>,
parent_log: usize,
observations: Vec<Observation>,
spent: Option<u64>,
}
#[test]
fn subtask_inherits_budget_session_and_event_log_across_threads() {
let runtime = tokio::runtime::Builder::new_multi_thread()
.worker_threads(4)
.enable_all()
.build()
.expect("multi-thread runtime");
crate::agent_sessions::reset_session_store();
crate::call_budget::reset_call_budget_state();
let measured = runtime.block_on(scope_placement(SubtaskPlacement::Worker, async {
let _budget = install_mcp_call_budget(BRANCHES as u64);
let _session = crate::agent_sessions::enter_current_session("parent-session");
let log = crate::event_log::install_memory_for_current_thread(64);
let parent_thread = std::thread::current().id();
let parent_session = crate::agent_sessions::current_session_id();
let parent_log = Arc::as_ptr(&log) as usize;
let registry = Arc::new(PoolRegistry::default());
let mut set: tokio::task::JoinSet<Observation> = tokio::task::JoinSet::new();
for _ in 0..BRANCHES {
let branch = prepare(Arc::clone(®istry), async {
charge_mcp_call().expect("branch charge is within the ceiling");
Observation {
thread: std::thread::current().id(),
session: crate::agent_sessions::current_session_id(),
event_log: crate::event_log::active_event_log()
.map(|log| Arc::as_ptr(&log) as usize),
}
});
spawn_into(&mut set, branch);
}
let mut observations = Vec::with_capacity(BRANCHES);
while let Some(joined) = set.join_next().await {
observations.push(joined.expect("branch joins"));
}
Measured {
parent_thread,
parent_session,
parent_log,
observations,
spent: mcp_calls_spent(),
}
}));
assert_eq!(measured.observations.len(), BRANCHES);
let migrated = measured
.observations
.iter()
.filter(|observed| observed.thread != measured.parent_thread)
.count();
assert!(
migrated > 0,
"no branch ran on a thread other than the parent's ({:?}); \
the inheritance assertions below would prove nothing",
measured.parent_thread
);
for observed in &measured.observations {
assert_eq!(
observed.session, measured.parent_session,
"a branch on {:?} lost the parent's agent session",
observed.thread
);
assert_eq!(
observed.event_log,
Some(measured.parent_log),
"a branch on {:?} wrote to a different event log than its parent",
observed.thread
);
}
assert_eq!(
measured.spent,
Some(BRANCHES as u64),
"the parent's call budget did not see its subtasks' charges"
);
}