use crate::{
agent::steering::AgentSteering,
cancellation::{AgentCancellation, AgentCancellationHandle},
output::ActivityId,
};
use std::{
collections::HashMap,
sync::{
Arc, Mutex,
atomic::{AtomicBool, Ordering},
},
};
const MAX_RETAINED_SUBAGENT_CONTROLS: usize = 64;
#[derive(Clone, Debug, Default)]
pub(crate) struct SubagentControls {
inner: Arc<Mutex<HashMap<ActivityId, SubagentControl>>>,
}
#[derive(Clone, Debug)]
pub(crate) struct SubagentControl {
pub(crate) steering: AgentSteering,
cancellation: AgentCancellationHandle,
cancellation_token: AgentCancellation,
active: Arc<AtomicBool>,
}
impl SubagentControl {
pub(crate) fn cancel(&self) {
self.cancellation.cancel();
self.steering.close();
}
pub(crate) fn is_active(&self) -> bool {
self.active.load(Ordering::SeqCst) && !self.cancellation_token.is_canceled()
}
}
impl SubagentControls {
pub(crate) fn has_pending_input(&self) -> bool {
self.inner
.lock()
.unwrap_or_else(|e| e.into_inner())
.values()
.any(|control| control.steering.try_pending_count() != Some(0))
}
pub(crate) fn get(&self, id: &ActivityId) -> Option<SubagentControl> {
self.inner
.lock()
.unwrap_or_else(|e| e.into_inner())
.get(id)
.cloned()
}
pub(super) fn register(
&self,
id: ActivityId,
steering: AgentSteering,
cancellation: AgentCancellationHandle,
cancellation_token: AgentCancellation,
) -> Option<SubagentControlRegistration> {
let mut entries = self.inner.lock().unwrap_or_else(|e| e.into_inner());
entries.retain(|_, control| {
control.active.load(Ordering::SeqCst) || control.steering.try_pending_count() != Some(0)
});
if entries.len() >= MAX_RETAINED_SUBAGENT_CONTROLS || entries.contains_key(&id) {
return None;
}
let control = SubagentControl {
steering,
cancellation,
cancellation_token,
active: Arc::new(AtomicBool::new(true)),
};
entries.insert(id.clone(), control.clone());
Some(SubagentControlRegistration {
registry: self.clone(),
id,
control,
})
}
}
pub(super) struct SubagentControlRegistration {
registry: SubagentControls,
id: ActivityId,
control: SubagentControl,
}
impl Drop for SubagentControlRegistration {
fn drop(&mut self) {
self.control.steering.close();
self.control.active.store(false, Ordering::SeqCst);
let has_pending_input = self.control.steering.pending_count() > 0;
let mut entries = self
.registry
.inner
.lock()
.unwrap_or_else(|e| e.into_inner());
if entries
.get(&self.id)
.is_some_and(|entry| Arc::ptr_eq(&entry.active, &self.control.active))
&& !has_pending_input
{
entries.remove(&self.id);
}
}
}