use super::{Task, TaskResult, TaskState};
use crate::config::RuntimeConfig;
use crate::error::RuntimeError;
use crate::state::StateTracker;
use crate::utils::topological_sort;
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::oneshot;
use tokio::task::JoinSet;
use tracing::{debug, error, info, warn};
pub struct TaskManager {
tasks: Vec<Box<dyn Task>>,
state_tracker: Arc<StateTracker>,
config: RuntimeConfig,
}
impl TaskManager {
pub fn new() -> Self {
Self {
tasks: Vec::new(),
state_tracker: Arc::new(StateTracker::new()),
config: RuntimeConfig::default(),
}
}
pub fn with_config(config: RuntimeConfig) -> Self {
Self {
tasks: Vec::new(),
state_tracker: Arc::new(StateTracker::new()),
config,
}
}
pub(crate) fn set_config(&mut self, config: RuntimeConfig) {
self.config = config;
}
pub(crate) fn shutdown_timeout(&self) -> Duration {
self.config.shutdown_timeout
}
pub fn add_task(&mut self, task: Box<dyn Task>) {
debug!(task_name = %task.name(), "Adding task to manager");
self.tasks.push(task);
}
pub fn task_count(&self) -> usize {
self.tasks.len()
}
pub fn state_tracker(&self) -> Arc<StateTracker> {
Arc::clone(&self.state_tracker)
}
pub async fn start_all(
&mut self,
) -> Result<(JoinSet<TaskResult>, Vec<oneshot::Sender<()>>), RuntimeError> {
info!(task_count = self.tasks.len(), "Starting all tasks");
let sorted_tasks = self.sort_tasks()?;
for task in &sorted_tasks {
self.state_tracker
.register_task(task.name(), TaskState::Pending)
.await;
}
let mut join_set = JoinSet::new();
let mut shutdown_txs = Vec::new();
for task in sorted_tasks {
let task_name = task.name().to_string();
let (shutdown_tx, shutdown_rx) = oneshot::channel();
shutdown_txs.push(shutdown_tx);
self.state_tracker
.update_state(&task_name, TaskState::Starting)
.await;
let state_tracker = Arc::clone(&self.state_tracker);
let task_future = task.run(shutdown_rx);
join_set.spawn(async move {
state_tracker
.update_state(&task_name, TaskState::Running)
.await;
let result = task_future.await;
match &result {
Ok(_) => {
debug!(task_name = %task_name, "Task completed");
state_tracker
.update_state(&task_name, TaskState::Stopped)
.await;
}
Err(e) => {
error!(task_name = %task_name, error = %e, "❌ Task failed");
state_tracker
.update_state_with_error(&task_name, TaskState::Failed, e.to_string())
.await;
}
}
result
});
}
info!("All tasks started");
Ok((join_set, shutdown_txs))
}
pub async fn stop_all(
&self,
mut join_set: JoinSet<TaskResult>,
shutdown_txs: Vec<oneshot::Sender<()>>,
) {
info!("Stopping all tasks");
for tx in shutdown_txs {
let _ = tx.send(());
}
match tokio::time::timeout(self.config.shutdown_timeout, async {
while let Some(result) = join_set.join_next().await {
match result {
Ok(Ok(_)) => {
debug!("Task completed gracefully");
}
Ok(Err(e)) => {
warn!("Task completed with error: {}", e);
}
Err(e) => {
warn!("Task join error: {}", e);
}
}
}
})
.await
{
Ok(_) => {
info!("All tasks completed");
}
Err(_) => {
warn!("Tasks shutdown timeout, forcing exit");
join_set.abort_all();
}
}
}
pub async fn wait_for_ready(&self) -> Result<(), RuntimeError> {
if !self.config.task_startup.enable_ready_check {
debug!("Task ready check is disabled, skipping");
return Ok(());
}
info!("Waiting for all tasks to be ready");
let timeout = self.config.task_startup.ready_check_timeout;
let start = std::time::Instant::now();
loop {
if self.state_tracker.all_ready().await {
info!("✅ All tasks are ready");
return Ok(());
}
if start.elapsed() > timeout {
return Err(RuntimeError::StartupTimeout {
name: "all tasks".to_string(),
timeout,
});
}
tokio::time::sleep(Duration::from_millis(100)).await;
}
}
fn sort_tasks(&mut self) -> Result<Vec<Box<dyn Task>>, RuntimeError> {
let items: Vec<(String, Vec<String>)> = self
.tasks
.iter()
.map(|task| (task.name().to_string(), task.dependencies()))
.collect();
let sorted_names = topological_sort(items)
.map_err(|cycle| RuntimeError::CircularDependency { tasks: cycle })?;
let mut by_name: HashMap<String, Box<dyn Task>> = HashMap::new();
for task in self.tasks.drain(..) {
by_name.insert(task.name().to_string(), task);
}
let mut sorted_tasks = Vec::with_capacity(sorted_names.len());
for name in sorted_names {
if let Some(task) = by_name.remove(&name) {
sorted_tasks.push(task);
}
}
for (_, task) in by_name {
warn!(
task_name = %task.name(),
"Task missing from topological order output, appending at end"
);
sorted_tasks.push(task);
}
Ok(sorted_tasks)
}
}
impl Default for TaskManager {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::task::SpawnTask;
#[tokio::test]
async fn test_task_manager_new() {
let manager = TaskManager::new();
assert_eq!(manager.task_count(), 0);
}
#[tokio::test]
async fn test_task_manager_add_task() {
let mut manager = TaskManager::new();
manager.add_task(Box::new(SpawnTask::new("task-1", async { Ok(()) })));
assert_eq!(manager.task_count(), 1);
}
#[tokio::test]
async fn test_task_manager_state_tracker() {
let manager = TaskManager::new();
let _tracker = manager.state_tracker();
}
#[tokio::test]
async fn test_task_manager_start_all_three_independent_no_panic() {
let mut manager = TaskManager::new();
manager.add_task(Box::new(SpawnTask::new("conversation-grpc", async {
tokio::time::sleep(std::time::Duration::from_millis(5)).await;
Ok(())
})));
manager.add_task(Box::new(SpawnTask::new("read-receipt-consumer", async {
tokio::time::sleep(std::time::Duration::from_millis(5)).await;
Ok(())
})));
manager.add_task(Box::new(SpawnTask::new(
"conversation-ensure-consumer",
async {
tokio::time::sleep(std::time::Duration::from_millis(5)).await;
Ok(())
},
)));
let (mut join_set, shutdown_txs) = manager.start_all().await.expect("start_all");
for tx in shutdown_txs {
let _ = tx.send(());
}
while join_set.join_next().await.is_some() {}
}
}