diff --git a/src/bus/dispatcher.rs b/src/bus/dispatcher.rs index dbc6ebe..e2db5fc 100644 --- a/src/bus/dispatcher.rs +++ b/src/bus/dispatcher.rs @@ -7,6 +7,7 @@ use tokio::sync::mpsc; use crate::bus::{MessageBus, OutboundMessage}; use crate::channels::ChannelManager; use crate::channels::base::{Channel, ChannelError}; +use crate::task_supervisor::TaskSupervisor; const LANE_CAPACITY: usize = 64; const LANE_IDLE_TIMEOUT: Duration = Duration::from_secs(300); @@ -18,13 +19,19 @@ const SEND_TIMEOUT: Duration = Duration::from_secs(30); pub struct OutboundDispatcher { bus: Arc, channel_manager: ChannelManager, + task_supervisor: TaskSupervisor, } impl OutboundDispatcher { - pub fn new(bus: Arc, channel_manager: ChannelManager) -> Self { + pub fn new( + bus: Arc, + channel_manager: ChannelManager, + task_supervisor: TaskSupervisor, + ) -> Self { Self { bus, channel_manager, + task_supervisor, } } @@ -52,7 +59,7 @@ impl OutboundDispatcher { continue; }; let (new_sender, receiver) = mpsc::channel(LANE_CAPACITY); - Self::spawn_lane(channel, receiver, msg.channel.clone(), msg.chat_id.clone()); + self.spawn_lane(channel, receiver, msg.channel.clone(), msg.chat_id.clone()); lanes.insert(lane_key.clone(), new_sender.clone()); sender = Some(new_sender); } @@ -80,7 +87,7 @@ impl OutboundDispatcher { continue; }; let (new_sender, receiver) = mpsc::channel(LANE_CAPACITY); - Self::spawn_lane(channel, receiver, msg.channel.clone(), msg.chat_id.clone()); + self.spawn_lane(channel, receiver, msg.channel.clone(), msg.chat_id.clone()); if new_sender.try_send(msg).is_ok() { lanes.insert(lane_key, new_sender); } @@ -90,27 +97,31 @@ impl OutboundDispatcher { } fn spawn_lane( + &self, channel: Arc, mut receiver: mpsc::Receiver, channel_name: String, chat_id: String, ) { - tokio::spawn(async move { - loop { - let msg = match tokio::time::timeout(LANE_IDLE_TIMEOUT, receiver.recv()).await { - Ok(Some(msg)) => msg, - Ok(None) | Err(_) => break, - }; - if let Err(error) = Self::send_with_retry(&*channel, msg).await { - tracing::error!( - channel = %channel_name, - chat_id = %chat_id, - error = %error, - "Failed to send message after retries" - ); + self.task_supervisor.spawn( + format!("outbound-lane:{channel_name}:{chat_id}"), + async move { + loop { + let msg = match tokio::time::timeout(LANE_IDLE_TIMEOUT, receiver.recv()).await { + Ok(Some(msg)) => msg, + Ok(None) | Err(_) => break, + }; + if let Err(error) = Self::send_with_retry(&*channel, msg).await { + tracing::error!( + channel = %channel_name, + chat_id = %chat_id, + error = %error, + "Failed to send message after retries" + ); + } } - } - }); + }, + ); } async fn send_with_retry( @@ -206,7 +217,8 @@ mod tests { }); manager.register_channel("recording", channel.clone()).await; - let dispatcher = OutboundDispatcher::new(bus.clone(), manager); + let supervisor = TaskSupervisor::new(); + let dispatcher = OutboundDispatcher::new(bus.clone(), manager, supervisor.clone()); let task = tokio::spawn(async move { dispatcher.run().await }); bus.publish_outbound(outbound("slow", "slow-1")) .await @@ -233,5 +245,6 @@ mod tests { assert_eq!(sent[0], "fast-1"); assert_eq!(&sent[1..], &["slow-1", "slow-2"]); task.abort(); + supervisor.shutdown(Duration::from_secs(1)).await; } } diff --git a/src/channels/cli_chat.rs b/src/channels/cli_chat.rs index 6e88027..17a25e8 100644 --- a/src/channels/cli_chat.rs +++ b/src/channels/cli_chat.rs @@ -1,4 +1,5 @@ use async_trait::async_trait; +use std::collections::HashMap; use std::sync::Arc; use tokio::sync::{Mutex, mpsc}; @@ -18,13 +19,19 @@ pub(crate) struct Client { current_session_id: Mutex>, } +impl Client { + pub(crate) fn chat_id(&self) -> &str { + &self.chat_id + } +} + // ============================================================================ // CliChatChannel - Channel implementation for CLI chat // ============================================================================ pub struct CliChatChannel { bus: std::sync::Mutex>>, - clients: Mutex>>, + clients: Mutex>>, } impl Default for CliChatChannel { @@ -37,7 +44,7 @@ impl CliChatChannel { pub fn new() -> Self { Self { bus: std::sync::Mutex::new(None), - clients: Mutex::new(Vec::new()), + clients: Mutex::new(HashMap::new()), } } @@ -56,7 +63,10 @@ impl CliChatChannel { chat_id: chat_id.clone(), current_session_id: Mutex::new(None), }); - self.clients.lock().await.push(client.clone()); + self.clients + .lock() + .await + .insert(chat_id.clone(), client.clone()); // Create initial session via control message let session_id = match self.create_session_via_control(&chat_id, None).await { @@ -76,6 +86,10 @@ impl CliChatChannel { (session_id, client) } + pub(crate) async fn unregister_client(&self, chat_id: &str) { + self.clients.lock().await.remove(chat_id); + } + /// Handle an inbound message from a client pub(crate) async fn handle_inbound(&self, client: Arc, raw_msg: &str) { match parse_inbound(raw_msg) { @@ -574,29 +588,72 @@ impl Channel for CliChatChannel { async fn stop(&self) -> Result<(), ChannelError> { *self.bus.lock().unwrap() = None; + self.clients.lock().await.clear(); Ok(()) } async fn send(&self, msg: OutboundMessage) -> Result<(), ChannelError> { - let clients = self.clients.lock().await.clone(); - for client in clients { - if client.chat_id != msg.chat_id { - continue; + let client = self.clients.lock().await.get(&msg.chat_id).cloned(); + let Some(client) = client else { + tracing::debug!(chat_id = %msg.chat_id, "No active CLI client for outbound message"); + return Ok(()); + }; + let outbound = if msg.metadata.get("_type").map(|v| v.as_str()) == Some("notification") { + WsOutbound::SystemNotification { + content: msg.content, } - let outbound = if msg.metadata.get("_type").map(|v| v.as_str()) == Some("notification") + } else { + WsOutbound::AssistantResponse { + id: crate::util::short_id(), + content: msg.content, + role: "assistant".to_string(), + } + }; + if client.sender.send(outbound).await.is_err() { + let mut clients = self.clients.lock().await; + if clients + .get(&msg.chat_id) + .is_some_and(|registered| Arc::ptr_eq(registered, &client)) { - WsOutbound::SystemNotification { - content: msg.content.clone(), - } - } else { - WsOutbound::AssistantResponse { - id: crate::util::short_id(), - content: msg.content.clone(), - role: "assistant".to_string(), - } - }; - let _ = client.sender.send(outbound).await; + clients.remove(&msg.chat_id); + } } Ok(()) } } + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn failed_sender_is_pruned_from_client_registry() { + let channel = CliChatChannel::new(); + let (sender, receiver) = mpsc::channel(1); + drop(receiver); + let client = Arc::new(Client { + sender, + chat_id: "dead-client".to_string(), + current_session_id: Mutex::new(None), + }); + channel + .clients + .lock() + .await + .insert("dead-client".to_string(), client); + + channel + .send(OutboundMessage { + channel: "cli_chat".to_string(), + chat_id: "dead-client".to_string(), + content: "message".to_string(), + reply_to: None, + media: Vec::new(), + metadata: Default::default(), + }) + .await + .unwrap(); + + assert!(channel.clients.lock().await.is_empty()); + } +} diff --git a/src/config/mod.rs b/src/config/mod.rs index 2d613b7..f67284f 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -180,9 +180,12 @@ pub struct SchedulerConfig { /// Poll interval in seconds (how often to check for due jobs) #[serde(default = "default_poll_interval_secs")] pub poll_interval_secs: u64, - /// Maximum concurrent job executions (currently sequential, reserved for future) + /// Maximum concurrent job executions. #[serde(default = "default_max_concurrent")] pub max_concurrent: usize, + /// Hard timeout for one scheduled execution. + #[serde(default = "default_execution_timeout_secs")] + pub execution_timeout_secs: u64, } fn default_scheduler_enabled() -> bool { @@ -197,12 +200,17 @@ fn default_max_concurrent() -> usize { 1 } +fn default_execution_timeout_secs() -> u64 { + 900 +} + impl Default for SchedulerConfig { fn default() -> Self { Self { enabled: true, poll_interval_secs: 60, max_concurrent: 1, + execution_timeout_secs: 900, } } } diff --git a/src/gateway/mod.rs b/src/gateway/mod.rs index 2ecd6f2..0990e87 100644 --- a/src/gateway/mod.rs +++ b/src/gateway/mod.rs @@ -14,6 +14,7 @@ use crate::mcp; use crate::memory::MemoryManager; use crate::scheduler::Scheduler; use crate::session::SessionManager; +use crate::task_supervisor::TaskSupervisor; pub struct GatewayState { pub config: Config, @@ -21,11 +22,15 @@ pub struct GatewayState { pub session_manager: Arc, pub channel_manager: ChannelManager, pub storage: Arc, + pub task_supervisor: TaskSupervisor, + pub connection_shutdown: tokio_util::sync::CancellationToken, } impl GatewayState { pub async fn new() -> Result> { let config = Config::load_default()?; + let task_supervisor = TaskSupervisor::new(); + let connection_shutdown = tokio_util::sync::CancellationToken::new(); // Initialize workspace directory: expand path and ensure it exists let workspace_path = expand_path(&config.workspace_dir); @@ -98,6 +103,7 @@ impl GatewayState { memory_manager, browser_config, config.gateway.max_concurrent_background_tasks, + task_supervisor.clone(), )?; let session_manager = Arc::new(session_manager); @@ -171,6 +177,8 @@ impl GatewayState { session_manager: session_manager.clone(), channel_manager, storage, + task_supervisor, + connection_shutdown, }) } @@ -198,7 +206,7 @@ impl GatewayState { // Spawn unified message processor // This handles both inbound AI messages and control messages in one loop - tokio::spawn(async move { + self.task_supervisor.spawn("message-processor", async move { tracing::info!("Message processor started"); loop { @@ -267,12 +275,17 @@ impl GatewayState { }); // Spawn outbound dispatcher - let dispatcher = OutboundDispatcher::new(bus_for_outbound, self.channel_manager.clone()); + let dispatcher = OutboundDispatcher::new( + bus_for_outbound, + self.channel_manager.clone(), + self.task_supervisor.clone(), + ); - tokio::spawn(async move { - tracing::info!("Outbound dispatcher started"); - dispatcher.run().await; - }); + self.task_supervisor + .spawn("outbound-dispatcher", async move { + tracing::info!("Outbound dispatcher started"); + dispatcher.run().await; + }); // Spawn scheduler background task if enabled let scheduler_config = self.config.gateway.scheduler.clone().unwrap_or_default(); @@ -282,7 +295,7 @@ impl GatewayState { self.session_manager.clone(), scheduler_config, )); - tokio::spawn(async move { + self.task_supervisor.spawn("scheduler", async move { sched.run().await; }); tracing::info!("Scheduler background task spawned"); @@ -412,24 +425,27 @@ pub async fn run( let listener = TcpListener::bind(&addr).await?; tracing::info!(address = %addr, "Gateway listening"); - // Graceful shutdown using oneshot channel - let (shutdown_tx, shutdown_rx) = tokio::sync::oneshot::channel::<()>(); - let channel_manager = state.channel_manager.clone(); - - // Spawn ctrl_c handler - tokio::spawn(async move { - tokio::signal::ctrl_c().await.ok(); - tracing::info!("Shutdown signal received"); - let _ = channel_manager.stop_all().await; - let _ = shutdown_tx.send(()); - }); - - // Serve with graceful shutdown - axum::serve(listener, app) - .with_graceful_shutdown(async { - shutdown_rx.await.ok(); + let connection_shutdown = state.connection_shutdown.clone(); + let serve_result = axum::serve(listener, app) + .with_graceful_shutdown(async move { + if let Err(error) = tokio::signal::ctrl_c().await { + tracing::error!(error = %error, "Failed to listen for shutdown signal"); + } + tracing::info!("Shutdown signal received"); + connection_shutdown.cancel(); }) - .await?; + .await; + + // Stop external intake before waiting for internal work to finish. + if let Err(error) = state.channel_manager.stop_all().await { + tracing::error!(error = %error, "Failed to stop channels cleanly"); + } + state.task_supervisor.cancel(); + state + .task_supervisor + .shutdown(std::time::Duration::from_secs(10)) + .await; + serve_result?; Ok(()) } diff --git a/src/gateway/ws.rs b/src/gateway/ws.rs index eb85920..1a320c6 100644 --- a/src/gateway/ws.rs +++ b/src/gateway/ws.rs @@ -7,6 +7,7 @@ use axum::response::Response; use futures_util::{SinkExt, StreamExt}; use std::sync::Arc; use tokio::sync::mpsc; +use tokio::time::{Duration, timeout}; pub async fn ws_handler(ws: WebSocketUpgrade, State(state): State>) -> Response { ws.on_upgrade(|socket| async move { @@ -23,6 +24,7 @@ async fn handle_socket(ws: WebSocket, state: Arc) { // Register client with CliChatChannel and get initial session id let (session_id, client) = cli_chat_channel.register_client(sender.clone()).await; + let chat_id = client.chat_id().to_string(); // Send session established message let _ = sender @@ -36,7 +38,7 @@ async fn handle_socket(ws: WebSocket, state: Arc) { let (mut ws_sender, mut ws_receiver) = ws.split(); // Task: forward from receiver to WebSocket - tokio::spawn(async move { + let mut writer_task = tokio::spawn(async move { while let Some(msg) = receiver.recv().await { if let Ok(text) = serialize_outbound(&msg) && ws_sender.send(WsMessage::Text(text.into())).await.is_err() @@ -47,18 +49,43 @@ async fn handle_socket(ws: WebSocket, state: Arc) { }); // Main loop: receive WebSocket messages and forward to CliChatChannel - while let Some(msg) = ws_receiver.next().await { - match msg { - Ok(WsMessage::Text(text)) => { - cli_chat_channel.handle_inbound(client.clone(), &text).await; - } - Ok(WsMessage::Close(_)) | Err(_) => { - tracing::debug!(session_id = %session_id, "WebSocket closed"); + let cancellation = state.connection_shutdown.clone(); + let mut writer_finished = false; + loop { + tokio::select! { + _ = cancellation.cancelled() => break, + result = &mut writer_task => { + writer_finished = true; + if let Err(error) = result { + tracing::warn!(session_id = %session_id, error = %error, "WebSocket writer task failed"); + } break; } - _ => {} + msg = ws_receiver.next() => { + match msg { + Some(Ok(WsMessage::Text(text))) => { + cli_chat_channel.handle_inbound(client.clone(), &text).await; + } + Some(Ok(WsMessage::Close(_))) | Some(Err(_)) | None => { + tracing::debug!(session_id = %session_id, "WebSocket closed"); + break; + } + _ => {} + } + } } } + cli_chat_channel.unregister_client(&chat_id).await; + drop(client); + drop(sender); + if !writer_finished + && timeout(Duration::from_secs(2), &mut writer_task) + .await + .is_err() + { + writer_task.abort(); + let _ = writer_task.await; + } tracing::info!(session_id = %session_id, "CLI session ended"); } diff --git a/src/lib.rs b/src/lib.rs index 5b8edd4..d714629 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -14,5 +14,6 @@ pub mod scheduler; pub mod session; pub mod skills; pub mod storage; +pub mod task_supervisor; pub mod tools; pub mod util; diff --git a/src/scheduler/mod.rs b/src/scheduler/mod.rs index 502b1b1..14e768f 100644 --- a/src/scheduler/mod.rs +++ b/src/scheduler/mod.rs @@ -2,6 +2,8 @@ pub mod types; use std::sync::Arc; use std::time::Instant; + +use futures_util::stream::{self, StreamExt}; use tokio::time; use crate::config::SchedulerConfig; @@ -55,6 +57,7 @@ pub struct Scheduler { storage: Arc, session_manager: Arc, config: SchedulerConfig, + owner: String, } impl Scheduler { @@ -67,209 +70,150 @@ impl Scheduler { storage, session_manager, config, + owner: uuid::Uuid::new_v4().to_string(), } } - /// Run the scheduler loop. This is a long-running async function meant to be - /// spawned as a tokio background task. + /// Claim due jobs with a durable lease, then execute the claimed batch with + /// bounded concurrency. pub async fn run(self: Arc) { - let poll_duration = time::Duration::from_secs(self.config.poll_interval_secs); + let poll_duration = time::Duration::from_secs(self.config.poll_interval_secs.max(1)); let mut interval = time::interval(poll_duration); - - interval.tick().await; + interval.set_missed_tick_behavior(time::MissedTickBehavior::Skip); + // Keep accidental configuration values from claiming an unbounded + // batch and overwhelming the runtime or SQLite parameter conversion. + let max_concurrent = self.config.max_concurrent.clamp(1, 256); tracing::info!( - "Scheduler started (poll interval: {}s, max concurrent: {})", - self.config.poll_interval_secs, - self.config.max_concurrent, + poll_interval_secs = self.config.poll_interval_secs, + max_concurrent, + scheduler_owner = %self.owner, + "Scheduler started" ); loop { interval.tick().await; - let now = now_ms(); - - let due = match self + let lease_ms = self + .config + .execution_timeout_secs + .saturating_add(30) + .saturating_mul(1000) + .min(i64::MAX as u64) as i64; + let lease_until = now.saturating_add(lease_ms); + let jobs = match self .storage - .due_scheduled_jobs(now, self.config.max_concurrent) + .claim_due_scheduled_jobs(now, lease_until, &self.owner, max_concurrent) .await { Ok(jobs) => jobs, - Err(e) => { - tracing::error!("scheduler: failed to query due jobs: {}", e); + Err(error) => { + tracing::error!(error = %error, "scheduler: failed to claim due jobs"); continue; } }; - if due.is_empty() { + if jobs.is_empty() { continue; } + tracing::info!(count = jobs.len(), "scheduler: claimed due jobs"); - tracing::info!("scheduler: found {} due job(s)", due.len()); - - for job in &due { - let start = Instant::now(); - let started_at = now_ms(); - - if let Err(e) = self - .storage - .touch_scheduled_job_last_run(&job.id, started_at) - .await - { - tracing::error!(job_id = %job.id, "scheduler: failed to touch last_run_at: {}", e); - continue; - } - - tracing::info!( - job_id = %job.id, - job_name = %job.name, - "scheduler: executing cron job" - ); - - let result = self - .session_manager - .handle_cron_message( - &job.channel, - &job.chat_id, - &job.prompt, - &job.id, - &job.name, - ) - .await; - - let finished_at = now_ms(); - let duration_ms = start.elapsed().as_millis() as i64; - - match result { - Ok(HandleResult::AgentResponse(output)) => { - let output_truncated = if output.len() > 8000 { - format!( - "{}...[truncated]", - &output[..output.ceil_char_boundary(8000)] - ) - } else { - output.clone() - }; - - let run = JobRun { - id: 0, - job_id: job.id.clone(), - started_at, - finished_at, - status: "ok".to_string(), - output: Some(output_truncated), - error: None, - duration_ms, - }; - - if let Err(e) = self.storage.record_scheduled_job_run(&run).await { - tracing::error!(job_id = %job.id, "scheduler: failed to record run: {}", e); - } - - if let Err(e) = self - .storage - .set_scheduled_job_last_status(&job.id, "ok", None) - .await - { - tracing::error!(job_id = %job.id, "scheduler: failed to set last_status: {}", e); - } - - tracing::info!( - job_id = %job.id, - duration_ms = %duration_ms, - "scheduler: job completed successfully" - ); - } - Ok(HandleResult::CommandOutput(output)) => { - let run = JobRun { - id: 0, - job_id: job.id.clone(), - started_at, - finished_at, - status: "ok".to_string(), - output: Some(output), - error: None, - duration_ms, - }; - - let _ = self.storage.record_scheduled_job_run(&run).await; - } - Ok(HandleResult::AgentProcessing) => { - tracing::warn!(job_id = %job.id, "scheduler: unexpected AgentProcessing from cron — response sent via bus"); - } - Err(e) => { - let error_str = e.to_string(); - let run = JobRun { - id: 0, - job_id: job.id.clone(), - started_at, - finished_at, - status: "error".to_string(), - output: None, - error: Some(error_str.clone()), - duration_ms, - }; - - if let Err(e2) = self.storage.record_scheduled_job_run(&run).await { - tracing::error!(job_id = %job.id, "scheduler: failed to record error run: {}", e2); - } - - if let Err(e2) = self - .storage - .set_scheduled_job_last_status(&job.id, "error", Some(&error_str)) - .await - { - tracing::error!(job_id = %job.id, "scheduler: failed to set error status: {}", e2); - } - - tracing::error!( - job_id = %job.id, - duration_ms = %duration_ms, - error = %error_str, - "scheduler: job failed" - ); - } - } - - if let Err(e) = self.reschedule_after_run(job).await { - tracing::error!(job_id = %job.id, "scheduler: failed to reschedule: {}", e); - } - } + stream::iter(jobs) + .for_each_concurrent(max_concurrent, |job| { + let scheduler = self.clone(); + async move { scheduler.execute_claimed_job(job).await } + }) + .await; } } - /// After a job runs, compute its next execution time or disable/delete it. - async fn reschedule_after_run(&self, job: &ScheduledJob) -> anyhow::Result<()> { - let now = now_ms(); + async fn execute_claimed_job(self: Arc, job: ScheduledJob) { + let start = Instant::now(); + let started_at = now_ms(); + tracing::info!(job_id = %job.id, job_name = %job.name, "scheduler: executing claimed job"); - match &job.schedule { - Schedule::At { .. } => { - if job.delete_after_run { - self.storage.remove_scheduled_job(&job.id).await?; - tracing::info!(job_id = %job.id, "scheduler: one-shot job deleted after run"); + let execution = self.session_manager.handle_cron_message( + &job.channel, + &job.chat_id, + &job.prompt, + &job.id, + &job.name, + ); + let result = time::timeout( + time::Duration::from_secs(self.config.execution_timeout_secs.max(1)), + execution, + ) + .await; + let finished_at = now_ms(); + let duration_ms = start.elapsed().as_millis() as i64; + + let (status, output, error) = match result { + Ok(Ok(HandleResult::AgentResponse(output) | HandleResult::CommandOutput(output))) => { + let output = if output.len() > 8000 { + format!( + "{}...[truncated]", + &output[..output.ceil_char_boundary(8000)] + ) } else { - self.storage - .set_scheduled_job_enabled(&job.id, false) - .await?; - tracing::info!(job_id = %job.id, "scheduler: one-shot job disabled after run"); - } + output + }; + ("ok".to_string(), Some(output), None) } + Ok(Ok(HandleResult::AgentProcessing)) => ( + "error".to_string(), + None, + Some("cron execution returned asynchronous processing".to_string()), + ), + Ok(Err(error)) => ("error".to_string(), None, Some(error.to_string())), + Err(_) => ( + "timeout".to_string(), + None, + Some(format!( + "execution exceeded {} seconds", + self.config.execution_timeout_secs.max(1) + )), + ), + }; + + let (next_run_at, disable, delete) = match &job.schedule { + Schedule::At { .. } => (None, !job.delete_after_run, job.delete_after_run), Schedule::Every { .. } | Schedule::Cron { .. } => { - if let Some(next) = next_run_for_schedule(&job.schedule, now) { - self.storage - .set_scheduled_job_next_run(&job.id, next) - .await?; - tracing::info!(job_id = %job.id, next_run_at = %next, "scheduler: job rescheduled"); - } else { - tracing::error!(job_id = %job.id, "scheduler: could not compute next run -- disabling job"); - self.storage - .set_scheduled_job_enabled(&job.id, false) - .await?; + match next_run_for_schedule(&job.schedule, finished_at) { + Some(next) => (Some(next), false, false), + None => (None, true, false), } } + }; + let run = JobRun { + id: 0, + job_id: job.id.clone(), + started_at, + finished_at, + status, + output, + error, + duration_ms, + }; + + if let Err(error) = self + .storage + .complete_scheduled_job(&run, &self.owner, next_run_at, disable, delete) + .await + { + tracing::error!(job_id = %job.id, error = %error, "scheduler: failed to commit job completion"); + let _ = self + .storage + .release_scheduled_job_lease(&job.id, &self.owner) + .await; + return; } - Ok(()) + tracing::info!( + job_id = %job.id, + status = %run.status, + duration_ms, + "scheduler: job completed" + ); } } diff --git a/src/session/session.rs b/src/session/session.rs index 7e8b870..093abfd 100644 --- a/src/session/session.rs +++ b/src/session/session.rs @@ -109,6 +109,14 @@ struct AgentTask { media: Vec, } +#[derive(Clone)] +struct AgentWorkerDeps { + bus: Arc, + memory_manager: Arc, + skills_loader: Arc, + task_supervisor: crate::task_supervisor::TaskSupervisor, +} + impl Session { pub async fn new( id: UnifiedSessionId, @@ -927,6 +935,7 @@ pub struct SessionManager { bus: Arc, memory_manager: Arc, sub_agent_manager: Arc, + task_supervisor: crate::task_supervisor::TaskSupervisor, } struct SessionManagerInner { @@ -1017,6 +1026,15 @@ pub static SLASH_COMMANDS: &[SlashCommand] = &[ ]; impl SessionManager { + fn worker_deps(&self) -> AgentWorkerDeps { + AgentWorkerDeps { + bus: self.bus.clone(), + memory_manager: self.memory_manager.clone(), + skills_loader: self.skills_loader.clone(), + task_supervisor: self.task_supervisor.clone(), + } + } + pub fn new( provider_config: LLMProviderConfig, storage: Arc, @@ -1024,6 +1042,7 @@ impl SessionManager { memory_manager: Arc, browser_config: Option, max_concurrent_background_tasks: usize, + task_supervisor: crate::task_supervisor::TaskSupervisor, ) -> Result { let mut skills_loader = SkillsLoader::new(); skills_loader.load_skills(); @@ -1051,7 +1070,7 @@ impl SessionManager { // Start background task notification consumer let sm_bus = bus.clone(); - tokio::spawn(async move { + task_supervisor.spawn("background-task-notifications", async move { while let Some(notif) = notify_rx.recv().await { let content = format_task_notification(¬if.task_id, ¬if.status, ¬if.result_summary); @@ -1069,7 +1088,7 @@ impl SessionManager { // Start periodic background task cleanup (every hour, TTL 24h) let cleanup_storage = storage.clone(); - tokio::spawn(async move { + task_supervisor.spawn("background-task-cleanup", async move { let mut interval = tokio::time::interval(std::time::Duration::from_secs(3600)); interval.tick().await; // skip immediate first tick loop { @@ -1098,6 +1117,7 @@ impl SessionManager { bus, memory_manager, sub_agent_manager, + task_supervisor, }) } @@ -1918,9 +1938,7 @@ impl SessionManager { spawn_agent_worker( rx, session_clone.clone(), - self.bus.clone(), - self.memory_manager.clone(), - self.skills_loader.clone(), + self.worker_deps(), generation, unified_str.clone(), ); @@ -1948,9 +1966,7 @@ impl SessionManager { spawn_agent_worker( rx, session_clone.clone(), - self.bus.clone(), - self.memory_manager.clone(), - self.skills_loader.clone(), + self.worker_deps(), generation, unified_str.clone(), ); @@ -2070,13 +2086,18 @@ async fn persist_added_messages( fn spawn_agent_worker( mut task_rx: mpsc::Receiver, session: Arc>, - bus: Arc, - memory_manager: Arc, - skills_loader: Arc, + deps: AgentWorkerDeps, worker_gen: u64, unified_str: String, ) { - tokio::spawn(async move { + let AgentWorkerDeps { + bus, + memory_manager, + skills_loader, + task_supervisor, + } = deps; + let worker_supervisor = task_supervisor.clone(); + task_supervisor.spawn(format!("session-worker:{unified_str}"), async move { let unified_for_source = unified_str.clone(); let _scope = CURRENT_SOURCE_SESSION.scope(Some(unified_for_source), async { 'tasks: while let Some(task) = task_rx.recv().await { @@ -2090,21 +2111,24 @@ fn spawn_agent_worker( let bus = bus.clone(); let ch = task_chan.clone(); let cid = task_cid.clone(); - tokio::spawn(async move { - while let Some(notif) = notify_rx.recv().await { - let mut metadata = HashMap::new(); - metadata.insert("_type".to_string(), "notification".to_string()); - let outbound = OutboundMessage { - channel: ch.clone(), - chat_id: cid.clone(), - content: notif, - reply_to: None, - media: vec![], - metadata, - }; - let _ = bus.publish_outbound(outbound).await; - } - }); + worker_supervisor.spawn( + format!("session-notifications:{ch}:{cid}"), + async move { + while let Some(notif) = notify_rx.recv().await { + let mut metadata = HashMap::new(); + metadata.insert("_type".to_string(), "notification".to_string()); + let outbound = OutboundMessage { + channel: ch.clone(), + chat_id: cid.clone(), + content: notif, + reply_to: None, + media: vec![], + metadata, + }; + let _ = bus.publish_outbound(outbound).await; + } + }, + ); } // Phase 1: capture a stable session snapshot under lock. diff --git a/src/storage/error.rs b/src/storage/error.rs index 2809b52..b82e3a5 100644 --- a/src/storage/error.rs +++ b/src/storage/error.rs @@ -13,4 +13,26 @@ pub enum StorageError { #[error("serialization error: {0}")] Serialization(String), + + #[error("schema migration error: {0}")] + Migration(String), +} + +impl StorageError { + /// Only retry failures that can plausibly clear without changing the data. + pub fn is_transient(&self) -> bool { + let Self::Database(error) = self else { + return false; + }; + match error { + sqlx::Error::PoolTimedOut => true, + sqlx::Error::Database(database) => { + matches!(database.code().as_deref(), Some("5" | "6" | "261" | "262")) || { + let message = database.message().to_ascii_lowercase(); + message.contains("database is locked") || message.contains("database is busy") + } + } + _ => false, + } + } } diff --git a/src/storage/mod.rs b/src/storage/mod.rs index 78221a5..f42534a 100644 --- a/src/storage/mod.rs +++ b/src/storage/mod.rs @@ -9,10 +9,13 @@ pub use background_task::BackgroundTask; pub use error::StorageError; pub use scheduler::{JobRun, ScheduledJob}; -use sqlx::{Pool, Row, Sqlite, SqlitePool}; +use sqlx::sqlite::{SqliteConnectOptions, SqliteJournalMode, SqlitePoolOptions, SqliteSynchronous}; +use sqlx::{Pool, Row, Sqlite}; use std::path::Path; use tokio::time::{Duration, sleep}; +const SCHEMA_VERSION: i64 = 1; + pub struct Storage { pub(crate) pool: Pool, } @@ -20,8 +23,17 @@ pub struct Storage { impl Storage { /// 打开或创建数据库 pub async fn new(db_path: &Path) -> Result { - let database_url = format!("sqlite:{}?mode=rwc", db_path.display()); - let pool = SqlitePool::connect(&database_url).await?; + let options = SqliteConnectOptions::new() + .filename(db_path) + .create_if_missing(true) + .journal_mode(SqliteJournalMode::Wal) + .synchronous(SqliteSynchronous::Normal) + .busy_timeout(Duration::from_secs(5)) + .foreign_keys(true); + let pool = SqlitePoolOptions::new() + .max_connections(8) + .connect_with(options) + .await?; let storage = Self { pool }; storage.init_schema().await?; @@ -75,6 +87,7 @@ impl Storage { tool_name TEXT, tool_calls TEXT, source TEXT, + reasoning_content TEXT, created_at INTEGER NOT NULL, FOREIGN KEY (session_id) REFERENCES sessions(id) ON DELETE CASCADE ) @@ -92,18 +105,6 @@ impl Storage { .execute(&self.pool) .await?; - // Migration: add source column if upgrading from older schema - sqlx::query(r#"ALTER TABLE messages ADD COLUMN source TEXT"#) - .execute(&self.pool) - .await - .ok(); - - // Migration: add reasoning_content column if upgrading from older schema - sqlx::query(r#"ALTER TABLE messages ADD COLUMN reasoning_content TEXT"#) - .execute(&self.pool) - .await - .ok(); - // Background tasks table — for async sub-agent tasks. // Note: No FOREIGN KEY on session_id because sessions use soft delete (deleted_at IS NULL). // Session and task association is maintained at the application level. @@ -217,36 +218,6 @@ impl Storage { .execute(&self.pool) .await?; - // Migration: add last_consolidated_at column if not exists - sqlx::query( - r#" - ALTER TABLE sessions ADD COLUMN archived_at INTEGER - "#, - ) - .execute(&self.pool) - .await - .ok(); - - // Migration: add last_consolidated_at column if not exists - sqlx::query( - r#" - ALTER TABLE sessions ADD COLUMN last_consolidated_at INTEGER - "#, - ) - .execute(&self.pool) - .await - .ok(); - - // Migration: add last_compressed_message_at column if not exists - sqlx::query( - r#" - ALTER TABLE sessions ADD COLUMN last_compressed_message_at INTEGER - "#, - ) - .execute(&self.pool) - .await - .ok(); - sqlx::query( r#" CREATE TABLE IF NOT EXISTS llm_calls ( @@ -264,13 +235,88 @@ impl Storage { .execute(&self.pool) .await?; - if let Err(e) = Self::init_scheduler_schema(&self.pool).await { - tracing::warn!( - "Failed to init scheduler schema (tables may already exist): {}", - e - ); + Self::init_scheduler_schema(&self.pool).await?; + self.migrate_schema().await?; + + Ok(()) + } + + /// Apply ordered, atomic migrations. Existing installations predate + /// `user_version`, so each step also checks the actual table shape. + async fn migrate_schema(&self) -> Result<(), StorageError> { + let current: i64 = sqlx::query_scalar("PRAGMA user_version") + .fetch_one(&self.pool) + .await?; + if current > SCHEMA_VERSION { + return Err(StorageError::Migration(format!( + "database schema version {current} is newer than supported version {SCHEMA_VERSION}" + ))); + } + if current == SCHEMA_VERSION { + return Ok(()); } + let mut tx = self.pool.begin().await?; + for (table, column, definition) in [ + ("messages", "source", "source TEXT"), + ("messages", "reasoning_content", "reasoning_content TEXT"), + ("sessions", "archived_at", "archived_at INTEGER"), + ( + "sessions", + "last_consolidated_at", + "last_consolidated_at INTEGER", + ), + ( + "sessions", + "last_compressed_message_at", + "last_compressed_message_at INTEGER", + ), + ("scheduled_jobs", "locked_at", "locked_at INTEGER"), + ("scheduled_jobs", "lock_owner", "lock_owner TEXT"), + ("scheduled_jobs", "lease_until", "lease_until INTEGER"), + ] { + let pragma = format!("PRAGMA table_info({table})"); + let columns = sqlx::query(&pragma).fetch_all(&mut *tx).await?; + if !columns + .iter() + .any(|row| row.get::("name") == column) + { + let alter = format!("ALTER TABLE {table} ADD COLUMN {definition}"); + sqlx::query(&alter).execute(&mut *tx).await?; + } + } + + let duplicate: Option<(String, i64, i64)> = sqlx::query_as( + r#" + SELECT session_id, seq, COUNT(*) + FROM messages + GROUP BY session_id, seq + HAVING COUNT(*) > 1 + LIMIT 1 + "#, + ) + .fetch_optional(&mut *tx) + .await?; + if let Some((session_id, seq, count)) = duplicate { + return Err(StorageError::Migration(format!( + "cannot enforce unique message sequence: session {session_id} has {count} rows at seq {seq}" + ))); + } + + sqlx::query( + "CREATE UNIQUE INDEX IF NOT EXISTS idx_messages_session_seq_unique ON messages(session_id, seq)", + ) + .execute(&mut *tx) + .await?; + sqlx::query( + "CREATE INDEX IF NOT EXISTS idx_jobs_claimable ON scheduled_jobs(enabled, next_run_at, lease_until)", + ) + .execute(&mut *tx) + .await?; + sqlx::query(&format!("PRAGMA user_version = {SCHEMA_VERSION}")) + .execute(&mut *tx) + .await?; + tx.commit().await?; Ok(()) } @@ -292,6 +338,9 @@ impl Storage { last_run_at INTEGER, last_status TEXT, last_error TEXT, + locked_at INTEGER, + lock_owner TEXT, + lease_until INTEGER, created_at INTEGER NOT NULL, updated_at INTEGER NOT NULL ) @@ -612,11 +661,32 @@ impl Storage { session_id: &str, msgs: &[crate::storage::message::MessageMeta], ) -> Result, StorageError> { + let mut tx = self.pool.begin().await?; let mut seqs = Vec::with_capacity(msgs.len()); for msg in msgs { - let seq = self.append_message(session_id, msg).await?; - seqs.push(seq); + sqlx::query( + r#" + INSERT INTO messages (id, session_id, seq, role, content, reasoning_content, media_refs, tool_call_id, tool_name, tool_calls, source, created_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + "#, + ) + .bind(&msg.id) + .bind(session_id) + .bind(msg.seq) + .bind(&msg.role) + .bind(&msg.content) + .bind(&msg.reasoning_content) + .bind(&msg.media_refs) + .bind(&msg.tool_call_id) + .bind(&msg.tool_name) + .bind(&msg.tool_calls) + .bind(&msg.source) + .bind(msg.created_at) + .execute(&mut *tx) + .await?; + seqs.push(msg.seq); } + tx.commit().await?; Ok(seqs) } @@ -701,7 +771,7 @@ impl Storage { for (attempt, delay) in delays.iter().enumerate() { match self.persist_message_batch(session_id, msgs, meta).await { Ok(()) => return Ok(()), - Err(error) if attempt < delays.len() - 1 => { + Err(error) if attempt < delays.len() - 1 && error.is_transient() => { tracing::warn!(attempt = attempt + 1, error = %error, "Turn persistence failed; retrying"); sleep(Duration::from_millis(*delay)).await; } @@ -1145,6 +1215,134 @@ mod tests { (storage, dir) } + #[tokio::test] + async fn sqlite_runtime_guards_are_enabled() { + let (storage, _dir) = create_test_storage().await; + let journal_mode: String = sqlx::query_scalar("PRAGMA journal_mode") + .fetch_one(storage.pool()) + .await + .unwrap(); + let foreign_keys: i64 = sqlx::query_scalar("PRAGMA foreign_keys") + .fetch_one(storage.pool()) + .await + .unwrap(); + let busy_timeout: i64 = sqlx::query_scalar("PRAGMA busy_timeout") + .fetch_one(storage.pool()) + .await + .unwrap(); + let schema_version: i64 = sqlx::query_scalar("PRAGMA user_version") + .fetch_one(storage.pool()) + .await + .unwrap(); + + assert_eq!(journal_mode, "wal"); + assert_eq!(foreign_keys, 1); + assert_eq!(busy_timeout, 5000); + assert_eq!(schema_version, SCHEMA_VERSION); + + let orphan = sqlx::query( + r#" + INSERT INTO messages (id, session_id, seq, role, content, created_at) + VALUES ('orphan', 'missing', 1, 'user', 'no parent', 1) + "#, + ) + .execute(storage.pool()) + .await; + assert!(orphan.is_err()); + } + + #[tokio::test] + async fn legacy_schema_is_migrated_without_rebuild() { + let dir = tempfile::tempdir().unwrap(); + let db_path = dir.path().join("legacy.db"); + let pool = SqlitePoolOptions::new() + .connect_with( + SqliteConnectOptions::new() + .filename(&db_path) + .create_if_missing(true), + ) + .await + .unwrap(); + sqlx::query( + r#" + CREATE TABLE sessions ( + id TEXT PRIMARY KEY, channel TEXT NOT NULL, chat_id TEXT NOT NULL, + dialog_id TEXT NOT NULL, title TEXT NOT NULL DEFAULT 'new', + created_at INTEGER NOT NULL, last_active_at INTEGER NOT NULL, + message_count INTEGER DEFAULT 0, routing_info TEXT, deleted_at INTEGER, + UNIQUE(channel, chat_id, dialog_id) + ) + "#, + ) + .execute(&pool) + .await + .unwrap(); + sqlx::query( + r#" + CREATE TABLE messages ( + id TEXT PRIMARY KEY, session_id TEXT NOT NULL, seq INTEGER NOT NULL, + role TEXT NOT NULL, content TEXT NOT NULL, media_refs TEXT, + tool_call_id TEXT, tool_name TEXT, tool_calls TEXT, + created_at INTEGER NOT NULL + ) + "#, + ) + .execute(&pool) + .await + .unwrap(); + sqlx::query( + r#" + CREATE TABLE scheduled_jobs ( + id TEXT PRIMARY KEY, name TEXT NOT NULL, schedule TEXT NOT NULL, + prompt TEXT NOT NULL, channel TEXT NOT NULL, chat_id TEXT NOT NULL, + model TEXT, enabled INTEGER NOT NULL DEFAULT 1, + delete_after_run INTEGER NOT NULL DEFAULT 0, next_run_at INTEGER NOT NULL, + last_run_at INTEGER, last_status TEXT, last_error TEXT, + created_at INTEGER NOT NULL, updated_at INTEGER NOT NULL + ) + "#, + ) + .execute(&pool) + .await + .unwrap(); + drop(pool); + + let storage = Storage::new(&db_path).await.unwrap(); + for (table, expected) in [ + ("messages", vec!["source", "reasoning_content"]), + ( + "sessions", + vec![ + "archived_at", + "last_consolidated_at", + "last_compressed_message_at", + ], + ), + ( + "scheduled_jobs", + vec!["locked_at", "lock_owner", "lease_until"], + ), + ] { + let columns = sqlx::query(&format!("PRAGMA table_info({table})")) + .fetch_all(storage.pool()) + .await + .unwrap(); + for column in expected { + assert!( + columns + .iter() + .any(|row| row.get::("name") == column), + "missing migrated column {table}.{column}" + ); + } + } + let schema_version: i64 = sqlx::query_scalar("PRAGMA user_version") + .fetch_one(storage.pool()) + .await + .unwrap(); + assert_eq!(schema_version, SCHEMA_VERSION); + } + #[tokio::test] async fn test_upsert_and_get_session() { let (storage, _dir) = create_test_storage().await; diff --git a/src/storage/scheduler.rs b/src/storage/scheduler.rs index 3a2eff6..26805b1 100644 --- a/src/storage/scheduler.rs +++ b/src/storage/scheduler.rs @@ -224,12 +224,137 @@ impl crate::storage::Storage { "SELECT * FROM scheduled_jobs WHERE enabled = 1 AND next_run_at <= ? ORDER BY next_run_at ASC LIMIT ?", ) .bind(now) - .bind(limit as i64) + .bind(i64::try_from(limit).unwrap_or(i64::MAX)) .fetch_all(self.pool()) .await?; rows.iter().map(row_to_job).collect() } + /// Atomically claim due jobs for one scheduler instance. A crashed worker's + /// claims become eligible again after `lease_until`. + pub async fn claim_due_scheduled_jobs( + &self, + now: i64, + lease_until: i64, + owner: &str, + limit: usize, + ) -> anyhow::Result> { + if limit == 0 { + return Ok(Vec::new()); + } + let rows = sqlx::query( + r#" + UPDATE scheduled_jobs + SET locked_at = ?, lock_owner = ?, lease_until = ?, last_run_at = ?, updated_at = ? + WHERE id IN ( + SELECT id FROM scheduled_jobs + WHERE enabled = 1 + AND next_run_at <= ? + AND (lease_until IS NULL OR lease_until <= ?) + ORDER BY next_run_at ASC + LIMIT ? + ) + AND (lease_until IS NULL OR lease_until <= ?) + RETURNING * + "#, + ) + .bind(now) + .bind(owner) + .bind(lease_until) + .bind(now) + .bind(now) + .bind(now) + .bind(now) + .bind(limit as i64) + .bind(now) + .fetch_all(self.pool()) + .await?; + rows.iter().map(row_to_job).collect() + } + + /// Persist the run result, reschedule/disable the job, and release its + /// lease in one transaction. The owner check prevents a stale worker from + /// completing a claim that has already been recovered elsewhere. + pub async fn complete_scheduled_job( + &self, + run: &JobRun, + owner: &str, + next_run_at: Option, + disable: bool, + delete: bool, + ) -> anyhow::Result<()> { + let mut tx = self.pool().begin().await?; + + if delete { + let result = sqlx::query("DELETE FROM scheduled_jobs WHERE id = ? AND lock_owner = ?") + .bind(&run.job_id) + .bind(owner) + .execute(&mut *tx) + .await?; + if result.rows_affected() != 1 { + anyhow::bail!("scheduled job lease lost before delete: {}", run.job_id); + } + } else { + sqlx::query( + r#" + INSERT INTO job_runs (job_id, started_at, finished_at, status, output, error, duration_ms) + VALUES (?, ?, ?, ?, ?, ?, ?) + "#, + ) + .bind(&run.job_id) + .bind(run.started_at) + .bind(run.finished_at) + .bind(&run.status) + .bind(&run.output) + .bind(&run.error) + .bind(run.duration_ms) + .execute(&mut *tx) + .await?; + + let result = sqlx::query( + r#" + UPDATE scheduled_jobs + SET next_run_at = COALESCE(?, next_run_at), + enabled = CASE WHEN ? THEN 0 ELSE enabled END, + last_status = ?, last_error = ?, + locked_at = NULL, lock_owner = NULL, lease_until = NULL, + updated_at = ? + WHERE id = ? AND lock_owner = ? + "#, + ) + .bind(next_run_at) + .bind(disable) + .bind(&run.status) + .bind(&run.error) + .bind(run.finished_at) + .bind(&run.job_id) + .bind(owner) + .execute(&mut *tx) + .await?; + if result.rows_affected() != 1 { + anyhow::bail!("scheduled job lease lost before completion: {}", run.job_id); + } + } + + tx.commit().await?; + Ok(()) + } + + pub async fn release_scheduled_job_lease( + &self, + job_id: &str, + owner: &str, + ) -> anyhow::Result<()> { + sqlx::query( + "UPDATE scheduled_jobs SET locked_at = NULL, lock_owner = NULL, lease_until = NULL WHERE id = ? AND lock_owner = ?", + ) + .bind(job_id) + .bind(owner) + .execute(self.pool()) + .await?; + Ok(()) + } + /// Record a job execution run. pub async fn record_scheduled_job_run(&self, run: &JobRun) -> anyhow::Result<()> { sqlx::query( @@ -608,4 +733,100 @@ mod tests { let got = storage.get_scheduled_job("job-update").await.unwrap(); assert_eq!(got.prompt, "new prompt"); } + + #[tokio::test] + async fn claim_is_exclusive_until_lease_expires() { + let storage = setup_storage().await; + let t = now(); + let job = ScheduledJob { + id: "leased-job".into(), + name: "leased".into(), + schedule: Schedule::Every { every_ms: 1000 }, + prompt: "run".into(), + channel: "cli_chat".into(), + chat_id: "c".into(), + model: None, + enabled: true, + delete_after_run: false, + next_run_at: t, + last_run_at: None, + last_status: None, + last_error: None, + created_at: t, + updated_at: t, + }; + storage.add_scheduled_job(&job).await.unwrap(); + + let first = storage + .claim_due_scheduled_jobs(t, t + 100, "owner-1", 1) + .await + .unwrap(); + let duplicate = storage + .claim_due_scheduled_jobs(t, t + 100, "owner-2", 1) + .await + .unwrap(); + let recovered = storage + .claim_due_scheduled_jobs(t + 101, t + 201, "owner-2", 1) + .await + .unwrap(); + + assert_eq!(first.len(), 1); + assert!(duplicate.is_empty()); + assert_eq!(recovered.len(), 1); + } + + #[tokio::test] + async fn completion_is_atomic_and_releases_lease() { + let storage = setup_storage().await; + let t = now(); + let job = ScheduledJob { + id: "complete-job".into(), + name: "complete".into(), + schedule: Schedule::Every { every_ms: 1000 }, + prompt: "run".into(), + channel: "cli_chat".into(), + chat_id: "c".into(), + model: None, + enabled: true, + delete_after_run: false, + next_run_at: t, + last_run_at: None, + last_status: None, + last_error: None, + created_at: t, + updated_at: t, + }; + storage.add_scheduled_job(&job).await.unwrap(); + storage + .claim_due_scheduled_jobs(t, t + 1000, "owner", 1) + .await + .unwrap(); + let run = super::JobRun { + id: 0, + job_id: job.id.clone(), + started_at: t, + finished_at: t + 10, + status: "ok".into(), + output: Some("done".into()), + error: None, + duration_ms: 10, + }; + storage + .complete_scheduled_job(&run, "owner", Some(t + 2000), false, false) + .await + .unwrap(); + + let completed = storage.get_scheduled_job(&job.id).await.unwrap(); + let runs = storage.list_scheduled_job_runs(&job.id, 10).await.unwrap(); + let lease: (Option, Option) = + sqlx::query_as("SELECT lock_owner, lease_until FROM scheduled_jobs WHERE id = ?") + .bind(&job.id) + .fetch_one(storage.pool()) + .await + .unwrap(); + assert_eq!(completed.next_run_at, t + 2000); + assert_eq!(completed.last_status.as_deref(), Some("ok")); + assert_eq!(runs.len(), 1); + assert_eq!(lease, (None, None)); + } } diff --git a/src/task_supervisor.rs b/src/task_supervisor.rs new file mode 100644 index 0000000..d890a07 --- /dev/null +++ b/src/task_supervisor.rs @@ -0,0 +1,160 @@ +use std::future::Future; +use std::panic::AssertUnwindSafe; +use std::sync::{Arc, Mutex}; +use std::time::Duration; + +use futures_util::FutureExt; +use tokio::task::JoinHandle; +use tokio::time::{Instant, timeout_at}; +use tokio_util::sync::CancellationToken; + +#[derive(Clone)] +pub struct TaskSupervisor { + inner: Arc, +} + +struct Inner { + cancellation: CancellationToken, + state: Mutex, +} + +impl Drop for Inner { + fn drop(&mut self) { + self.cancellation.cancel(); + } +} + +#[derive(Default)] +struct State { + stopping: bool, + tasks: Vec, +} + +struct ManagedTask { + name: String, + handle: JoinHandle<()>, +} + +impl Default for TaskSupervisor { + fn default() -> Self { + Self::new() + } +} + +impl TaskSupervisor { + pub fn new() -> Self { + Self { + inner: Arc::new(Inner { + cancellation: CancellationToken::new(), + state: Mutex::new(State::default()), + }), + } + } + + pub fn cancellation_token(&self) -> CancellationToken { + self.inner.cancellation.clone() + } + + /// Register a task before shutdown begins. Cancellation drops the task + /// future, so task code should keep externally visible state transactional. + pub fn spawn(&self, name: impl Into, future: F) -> bool + where + F: Future + Send + 'static, + { + let name = name.into(); + let cancellation = self.inner.cancellation.clone(); + let mut state = self.inner.state.lock().unwrap_or_else(|e| e.into_inner()); + if state.stopping { + return false; + } + + // Completed handles no longer need to occupy the registry. Panics are + // observed and logged inside the wrapper below. + state.tasks.retain(|task| !task.handle.is_finished()); + let task_name = name.clone(); + let handle = tokio::spawn(async move { + tracing::debug!(task = %task_name, "Background task started"); + let outcome = tokio::select! { + _ = cancellation.cancelled() => None, + outcome = AssertUnwindSafe(future).catch_unwind() => Some(outcome), + }; + match outcome { + Some(Ok(())) => tracing::debug!(task = %task_name, "Background task finished"), + Some(Err(_)) => tracing::error!(task = %task_name, "Background task panicked"), + None => tracing::debug!(task = %task_name, "Background task cancelled"), + } + }); + state.tasks.push(ManagedTask { name, handle }); + true + } + + pub fn cancel(&self) { + let mut state = self.inner.state.lock().unwrap_or_else(|e| e.into_inner()); + state.stopping = true; + self.inner.cancellation.cancel(); + } + + /// Stop accepting tasks, broadcast cancellation, and wait up to `grace`. + /// Remaining tasks are aborted so shutdown has a deterministic upper bound. + pub async fn shutdown(&self, grace: Duration) { + let mut tasks = { + let mut state = self.inner.state.lock().unwrap_or_else(|e| e.into_inner()); + state.stopping = true; + self.inner.cancellation.cancel(); + std::mem::take(&mut state.tasks) + }; + let deadline = Instant::now() + grace; + + for index in 0..tasks.len() { + let result = timeout_at(deadline, &mut tasks[index].handle).await; + match result { + Ok(Ok(())) => {} + Ok(Err(error)) if error.is_cancelled() => {} + Ok(Err(error)) => { + tracing::error!(task = %tasks[index].name, error = %error, "Background task join failed"); + } + Err(_) => { + for task in &tasks[index..] { + if !task.handle.is_finished() { + tracing::warn!(task = %task.name, "Aborting background task after shutdown grace period"); + task.handle.abort(); + } + } + for task in &mut tasks[index..] { + let _ = (&mut task.handle).await; + } + break; + } + } + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::atomic::{AtomicBool, Ordering}; + + #[tokio::test] + async fn shutdown_cancels_registered_task() { + let supervisor = TaskSupervisor::new(); + let dropped = Arc::new(AtomicBool::new(false)); + let marker = dropped.clone(); + supervisor.spawn("pending", async move { + struct DropMarker(Arc); + impl Drop for DropMarker { + fn drop(&mut self) { + self.0.store(true, Ordering::SeqCst); + } + } + let _marker = DropMarker(marker); + std::future::pending::<()>().await; + }); + + tokio::task::yield_now().await; + supervisor.cancel(); + assert!(!supervisor.spawn("late", async {})); + supervisor.shutdown(Duration::from_secs(1)).await; + assert!(dropped.load(Ordering::SeqCst)); + } +}