use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use parking_lot::RwLock;
use tokio::sync::Notify;
use tokio_util::sync::CancellationToken;
use crate::lifecycle::AgentSupervisor;
use oxi_agent::AgentConfig;
pub const DEFAULT_MAX_SUBAGENT_DEPTH: u32 = 2;
#[derive(Debug, Clone)]
pub enum SubagentState {
Pending {
registered_at_ms: u64,
},
Active {
started_at_ms: u64,
},
Completed {
finished_at_ms: u64,
response: String,
},
Failed {
finished_at_ms: u64,
error: String,
},
Cancelled {
finished_at_ms: u64,
},
}
impl SubagentState {
pub fn is_terminal(&self) -> bool {
matches!(
self,
SubagentState::Completed { .. }
| SubagentState::Failed { .. }
| SubagentState::Cancelled { .. }
)
}
pub fn is_active(&self) -> bool {
matches!(self, SubagentState::Active { .. })
}
}
#[derive(Debug, Clone)]
pub struct SubagentSpawnRequest {
pub agent_id: String,
pub config: AgentConfig,
pub task: String,
pub run_in_background: bool,
pub resume_from: Option<String>,
pub depth: u32,
}
#[derive(Debug)]
pub struct SubagentTracker {
cancel_token: CancellationToken,
completion: Arc<Notify>,
state: Arc<RwLock<SubagentState>>,
run_in_background: bool,
resume_from: Option<String>,
spawned_at_ms: u64,
}
impl SubagentTracker {
pub fn state(&self) -> SubagentState {
self.state.read().clone()
}
pub fn run_in_background(&self) -> bool {
self.run_in_background
}
pub fn resume_from(&self) -> Option<&str> {
self.resume_from.as_deref()
}
pub fn spawned_at_ms(&self) -> u64 {
self.spawned_at_ms
}
pub fn cancel(&self) {
self.cancel_token.cancel();
}
pub async fn wait_for_completion(&self, timeout: Duration) -> Option<SubagentState> {
let current = self.state();
if current.is_terminal() {
return Some(current);
}
match tokio::time::timeout(timeout, self.completion.notified()).await {
Ok(()) => Some(self.state()),
Err(_) => None,
}
}
}
#[derive(Debug, thiserror::Error)]
pub enum SubagentCoordinatorError {
#[error("subagent depth {depth} exceeds maximum {max}")]
MaxDepthExceeded {
depth: u32,
max: u32,
},
#[error("subagent agent_id '{0}' already in use")]
DuplicateId(String),
#[error("supervisor spawn failed: {0}")]
SpawnFailed(String),
#[error("resume_from agent '{0}' not found")]
ResumeFromNotFound(String),
}
pub type Result<T, E = SubagentCoordinatorError> = std::result::Result<T, E>;
#[derive(Clone)]
pub struct SubagentCoordinator {
supervisor: AgentSupervisor,
trackers: Arc<RwLock<HashMap<String, Arc<SubagentTracker>>>>,
last_responses: Arc<RwLock<HashMap<String, String>>>,
max_depth: u32,
}
impl SubagentCoordinator {
pub fn new(supervisor: AgentSupervisor) -> Self {
Self::with_max_depth(supervisor, DEFAULT_MAX_SUBAGENT_DEPTH)
}
pub fn with_max_depth(supervisor: AgentSupervisor, max_depth: u32) -> Self {
Self {
supervisor,
trackers: Arc::new(RwLock::new(HashMap::new())),
last_responses: Arc::new(RwLock::new(HashMap::new())),
max_depth,
}
}
pub fn max_depth(&self) -> u32 {
self.max_depth
}
pub fn supervisor(&self) -> &AgentSupervisor {
&self.supervisor
}
pub fn tracked_count(&self) -> usize {
self.trackers.read().len()
}
pub fn tracker(&self, agent_id: &str) -> Option<Arc<SubagentTracker>> {
self.trackers.read().get(agent_id).cloned()
}
pub fn state(&self, agent_id: &str) -> Option<SubagentState> {
self.tracker(agent_id).map(|t| t.state())
}
pub fn snapshot(&self) -> HashMap<String, SubagentState> {
self.trackers
.read()
.iter()
.map(|(id, t)| (id.clone(), t.state()))
.collect()
}
pub fn spawn(&self, req: SubagentSpawnRequest) -> Result<String> {
if req.depth > self.max_depth {
return Err(SubagentCoordinatorError::MaxDepthExceeded {
depth: req.depth,
max: self.max_depth,
});
}
if self.trackers.read().contains_key(&req.agent_id) {
return Err(SubagentCoordinatorError::DuplicateId(req.agent_id));
}
let task = if let Some(parent_id) = &req.resume_from {
let parent_response = self
.last_responses
.read()
.get(parent_id)
.cloned()
.ok_or_else(|| SubagentCoordinatorError::ResumeFromNotFound(parent_id.clone()))?;
format!(
"Previous context from agent '{parent_id}':\n---\n{parent_response}\n---\n\n{task}",
task = req.task
)
} else {
req.task.clone()
};
let now = now_ms();
let state = Arc::new(RwLock::new(SubagentState::Pending {
registered_at_ms: now,
}));
let completion = Arc::new(Notify::new());
let cancel_token = CancellationToken::new();
let tracker = Arc::new(SubagentTracker {
cancel_token: cancel_token.clone(),
completion: completion.clone(),
state: state.clone(),
run_in_background: req.run_in_background,
resume_from: req.resume_from.clone(),
spawned_at_ms: now,
});
self.trackers
.write()
.insert(req.agent_id.clone(), tracker.clone());
let handle = self.supervisor.spawn(req.config).map_err(|e| {
self.trackers.write().remove(&req.agent_id);
SubagentCoordinatorError::SpawnFailed(e.to_string())
})?;
let agent_id = req.agent_id.clone();
let last_responses = self.last_responses.clone();
let state_for_task = state.clone();
let completion_for_task = completion.clone();
tokio::spawn(async move {
{
let mut s = state_for_task.write();
*s = SubagentState::Active {
started_at_ms: now_ms(),
};
}
let outcome = tokio::select! {
_ = cancel_token.cancelled() => {
let mut s = state_for_task.write();
*s = SubagentState::Cancelled { finished_at_ms: now_ms() };
None
}
r = handle.run(task) => Some(r),
};
if let Some(res) = outcome {
let mut s = state_for_task.write();
match res {
Ok((response, _)) => {
last_responses
.write()
.insert(agent_id.clone(), response.content.clone());
*s = SubagentState::Completed {
finished_at_ms: now_ms(),
response: response.content,
};
}
Err(e) => {
*s = SubagentState::Failed {
finished_at_ms: now_ms(),
error: e.to_string(),
};
}
}
}
completion_for_task.notify_waiters();
});
Ok(req.agent_id)
}
pub fn cancel(&self, agent_id: &str) -> bool {
if let Some(t) = self.tracker(agent_id) {
t.cancel();
true
} else {
false
}
}
pub async fn block_wait_slot(
&self,
agent_id: &str,
timeout: Duration,
) -> Option<SubagentState> {
let tracker = self.tracker(agent_id)?;
tracker.wait_for_completion(timeout).await
}
}
fn now_ms() -> u64 {
use std::time::{SystemTime, UNIX_EPOCH};
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_millis() as u64)
.unwrap_or(0)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::error::SdkError;
use crate::lifecycle::SnapshotStore;
use std::future::Future;
use std::pin::Pin;
struct NoopSnapshotStore;
impl SnapshotStore for NoopSnapshotStore {
fn save<'a>(
&'a self,
_snapshot: &'a crate::lifecycle::AgentSnapshot,
) -> Pin<Box<dyn Future<Output = anyhow::Result<()>> + Send + 'a>> {
Box::pin(async { Ok(()) })
}
fn load<'a>(
&'a self,
_agent_id: &'a str,
) -> Pin<
Box<
dyn Future<Output = anyhow::Result<Option<crate::lifecycle::AgentSnapshot>>>
+ Send
+ 'a,
>,
> {
Box::pin(async { Ok(None) })
}
fn list(&self) -> Pin<Box<dyn Future<Output = anyhow::Result<Vec<String>>> + Send + '_>> {
Box::pin(async { Ok(vec![]) })
}
fn delete<'a>(
&'a self,
_agent_id: &'a str,
) -> Pin<Box<dyn Future<Output = anyhow::Result<()>> + Send + 'a>> {
Box::pin(async { Ok(()) })
}
}
struct FailingResolver;
impl oxi_agent::ProviderResolver for FailingResolver {
fn resolve_model(&self, _id: &str) -> Option<oxi_ai::Model> {
None
}
fn resolve_provider(&self, _provider: &str) -> Option<Arc<dyn oxi_ai::Provider>> {
None
}
}
fn make_coordinator(max_depth: u32) -> SubagentCoordinator {
let resolver: Arc<dyn oxi_agent::ProviderResolver> = Arc::new(FailingResolver);
let store: Arc<dyn SnapshotStore> = Arc::new(NoopSnapshotStore);
let supervisor = AgentSupervisor::new(resolver, store);
SubagentCoordinator::with_max_depth(supervisor, max_depth)
}
fn basic_request(id: &str, depth: u32) -> SubagentSpawnRequest {
SubagentSpawnRequest {
agent_id: id.to_string(),
config: AgentConfig {
model_id: "anthropic/claude-3-5-sonnet".into(),
..Default::default()
},
task: "do nothing".into(),
run_in_background: true,
resume_from: None,
depth,
}
}
#[test]
fn rejects_depth_above_max() {
let coord = make_coordinator(2);
let req = basic_request("a", 3);
let err = coord.spawn(req).unwrap_err();
assert!(
matches!(
err,
SubagentCoordinatorError::MaxDepthExceeded { depth: 3, max: 2 }
),
"got: {err:?}"
);
}
#[test]
fn spawn_fails_when_resolver_fails() {
let coord = make_coordinator(2);
let err = coord.spawn(basic_request("a", 0)).unwrap_err();
assert!(
matches!(err, SubagentCoordinatorError::SpawnFailed(_)),
"got: {err:?}"
);
assert_eq!(
coord.tracked_count(),
0,
"tracker must roll back on spawn failure"
);
}
#[test]
fn rejects_unknown_resume_from() {
let coord = make_coordinator(2);
let mut req = basic_request("a", 0);
req.resume_from = Some("nonexistent".into());
let err = coord.spawn(req).unwrap_err();
assert!(
matches!(err, SubagentCoordinatorError::ResumeFromNotFound(_)),
"got: {err:?}"
);
}
#[test]
fn tracked_count_starts_zero() {
let coord = make_coordinator(2);
assert_eq!(coord.tracked_count(), 0);
assert_eq!(coord.max_depth(), 2);
}
#[test]
fn cancel_for_unknown_returns_false() {
let coord = make_coordinator(2);
assert!(!coord.cancel("ghost"));
}
#[test]
fn block_wait_slot_unknown_returns_none() {
let coord = make_coordinator(2);
let rt = tokio::runtime::Builder::new_current_thread()
.enable_time()
.build()
.unwrap();
let r = rt.block_on(coord.block_wait_slot("ghost", Duration::from_millis(10)));
assert!(r.is_none());
}
#[test]
fn snapshot_of_empty_coordinator() {
let coord = make_coordinator(2);
assert!(coord.snapshot().is_empty());
}
#[test]
fn default_max_depth_is_two() {
let resolver: Arc<dyn oxi_agent::ProviderResolver> = Arc::new(FailingResolver);
let store: Arc<dyn SnapshotStore> = Arc::new(NoopSnapshotStore);
let supervisor = AgentSupervisor::new(resolver, store);
let coord = SubagentCoordinator::new(supervisor);
assert_eq!(coord.max_depth(), DEFAULT_MAX_SUBAGENT_DEPTH);
assert_eq!(coord.max_depth(), 2);
}
#[test]
fn error_type_is_displayable() {
let e1 = SubagentCoordinatorError::MaxDepthExceeded { depth: 3, max: 2 };
let e2 = SubagentCoordinatorError::DuplicateId("x".into());
let e3 = SubagentCoordinatorError::SpawnFailed("nope".into());
let e4 = SubagentCoordinatorError::ResumeFromNotFound("p".into());
assert!(!e1.to_string().is_empty());
assert!(!e2.to_string().is_empty());
assert!(!e3.to_string().is_empty());
assert!(!e4.to_string().is_empty());
}
#[test]
fn sdkerror_unused() {
let _ = std::marker::PhantomData::<SdkError>;
}
}