Skip to main content

mentra/
background.rs

1mod hook;
2mod observer;
3mod store;
4
5use std::{
6    collections::HashMap,
7    path::PathBuf,
8    sync::{
9        Arc, Mutex,
10        atomic::{AtomicU64, Ordering},
11    },
12};
13
14use serde::{Deserialize, Serialize};
15use strum::Display;
16
17use crate::agent::AgentEvent;
18use crate::runtime::{
19    RuntimeStore,
20    control::{CommandOutput, CommandRequest, RuntimeExecutor},
21};
22
23pub(crate) use hook::BackgroundHookSink;
24pub(crate) use observer::{BackgroundObserverSink, BackgroundRegistration};
25pub use store::BackgroundStore;
26
27const OUTPUT_PREVIEW_MAX_CHARS: usize = 500;
28const NOTIFICATION_PENDING: i64 = 0;
29const NOTIFICATION_ACKED: i64 = 2;
30
31#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Display)]
32#[strum(serialize_all = "snake_case")]
33#[serde(rename_all = "snake_case")]
34pub enum BackgroundTaskStatus {
35    Running,
36    Finished,
37    Failed,
38    Interrupted,
39}
40
41#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
42pub struct BackgroundTaskSummary {
43    pub id: String,
44    pub command: String,
45    pub cwd: PathBuf,
46    pub status: BackgroundTaskStatus,
47    pub output_preview: Option<String>,
48}
49
50#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
51pub struct BackgroundNotification {
52    pub task_id: String,
53    pub command: String,
54    pub cwd: PathBuf,
55    pub status: BackgroundTaskStatus,
56    pub output_preview: String,
57}
58
59#[derive(Clone)]
60pub(crate) struct BackgroundTaskManager {
61    inner: Arc<BackgroundTaskManagerInner>,
62}
63
64struct BackgroundTaskManagerInner {
65    store: Arc<dyn RuntimeStore>,
66    executor: Arc<dyn RuntimeExecutor>,
67    hooks: Arc<dyn BackgroundHookSink>,
68    next_task_id: AtomicU64,
69    state: Mutex<BackgroundTaskManagerState>,
70}
71
72#[derive(Default)]
73struct BackgroundTaskManagerState {
74    agents: HashMap<String, AgentBackgroundState>,
75}
76
77#[derive(Default)]
78struct AgentBackgroundState {
79    tasks: Vec<BackgroundTaskSummary>,
80    observer: Option<BackgroundObserver>,
81}
82
83#[derive(Clone)]
84struct BackgroundObserver {
85    sink: Arc<dyn BackgroundObserverSink>,
86}
87
88impl BackgroundTaskManager {
89    pub(crate) fn new(
90        store: Arc<dyn RuntimeStore>,
91        executor: Arc<dyn RuntimeExecutor>,
92        hooks: Arc<dyn BackgroundHookSink>,
93    ) -> Self {
94        Self {
95            inner: Arc::new(BackgroundTaskManagerInner {
96                store,
97                executor,
98                hooks,
99                next_task_id: AtomicU64::default(),
100                state: Mutex::new(BackgroundTaskManagerState::default()),
101            }),
102        }
103    }
104
105    pub(crate) fn register_agent(&self, registration: BackgroundRegistration) {
106        let BackgroundRegistration { agent_id, observer } = registration;
107        let tasks = {
108            let mut state = self
109                .inner
110                .state
111                .lock()
112                .expect("background manager poisoned");
113            let agent = state.agents.entry(agent_id.clone()).or_default();
114            agent.tasks = self
115                .inner
116                .store
117                .load_background_tasks(&agent_id)
118                .unwrap_or_default();
119            agent.observer = Some(BackgroundObserver {
120                sink: observer.clone(),
121            });
122            agent.tasks.clone()
123        };
124
125        observer.publish_snapshot(&tasks);
126    }
127
128    pub(crate) fn start_task(
129        &self,
130        agent_id: &str,
131        request: CommandRequest,
132    ) -> Result<BackgroundTaskSummary, String> {
133        let task_id = format!(
134            "bg-{}",
135            self.inner.next_task_id.fetch_add(1, Ordering::Relaxed) + 1
136        );
137        let summary = BackgroundTaskSummary {
138            id: task_id.clone(),
139            command: request.spec.display().to_string(),
140            cwd: request.cwd.clone(),
141            status: BackgroundTaskStatus::Running,
142            output_preview: None,
143        };
144        let _ = self
145            .inner
146            .store
147            .upsert_background_task(agent_id, &summary, NOTIFICATION_ACKED);
148
149        let (observer, tasks) = {
150            let mut state = self
151                .inner
152                .state
153                .lock()
154                .expect("background manager poisoned");
155            let agent = state.agents.entry(agent_id.to_string()).or_default();
156            agent.tasks.push(summary.clone());
157            (agent.observer.clone(), agent.tasks.clone())
158        };
159        self.publish_observer(
160            observer,
161            tasks,
162            AgentEvent::BackgroundTaskStarted {
163                task: summary.clone(),
164            },
165        );
166        let _ =
167            self.inner
168                .hooks
169                .task_started(agent_id, &summary.id, &summary.command, &summary.cwd);
170
171        let manager = self.clone();
172        let agent_id = agent_id.to_string();
173        let executor = self.inner.executor.clone();
174        tokio::spawn(async move {
175            let completed = execute_task(task_id, request, executor).await;
176            manager.finish_task(&agent_id, completed);
177        });
178
179        Ok(summary)
180    }
181
182    pub(crate) fn running_task_count(&self, agent_id: &str) -> usize {
183        let state = self
184            .inner
185            .state
186            .lock()
187            .expect("background manager poisoned");
188        state
189            .agents
190            .get(agent_id)
191            .map(|agent| {
192                agent
193                    .tasks
194                    .iter()
195                    .filter(|task| task.status == BackgroundTaskStatus::Running)
196                    .count()
197            })
198            .unwrap_or(0)
199    }
200
201    pub(crate) fn drain_notifications(&self, agent_id: &str) -> Vec<BackgroundNotification> {
202        self.inner
203            .store
204            .drain_background_notifications(agent_id)
205            .unwrap_or_default()
206    }
207
208    pub(crate) fn has_deliverable_notifications(&self, agent_id: &str) -> bool {
209        self.inner
210            .store
211            .has_deliverable_background_notifications(agent_id)
212            .unwrap_or(false)
213    }
214
215    pub(crate) fn requeue_notifications(
216        &self,
217        agent_id: &str,
218        notifications: Vec<BackgroundNotification>,
219    ) {
220        if notifications.is_empty() {
221            return;
222        }
223        let _ = self.inner.store.requeue_background_notifications(agent_id);
224    }
225
226    pub(crate) fn acknowledge_notifications(&self, agent_id: &str) {
227        let _ = self.inner.store.ack_background_notifications(agent_id);
228    }
229
230    pub(crate) fn check_task(
231        &self,
232        agent_id: &str,
233        task_id: Option<&str>,
234    ) -> Result<String, String> {
235        let state = self
236            .inner
237            .state
238            .lock()
239            .expect("background manager poisoned");
240        let Some(agent) = state.agents.get(agent_id) else {
241            return Ok("No background tasks.".to_string());
242        };
243
244        if let Some(task_id) = task_id {
245            let task = agent
246                .tasks
247                .iter()
248                .find(|task| task.id == task_id)
249                .ok_or_else(|| format!("Unknown background task {task_id}"))?;
250            return Ok(render_task_detail(task));
251        }
252
253        if agent.tasks.is_empty() {
254            return Ok("No background tasks.".to_string());
255        }
256
257        Ok(agent
258            .tasks
259            .iter()
260            .map(render_task_summary)
261            .collect::<Vec<_>>()
262            .join("\n"))
263    }
264
265    fn finish_task(&self, agent_id: &str, completed: CompletedBackgroundTask) {
266        let summary = BackgroundTaskSummary {
267            id: completed.id.clone(),
268            command: completed.command.clone(),
269            cwd: completed.cwd.clone(),
270            status: completed.status.clone(),
271            output_preview: Some(completed.output_preview.clone()),
272        };
273        let (observer, tasks) = {
274            let mut state = self
275                .inner
276                .state
277                .lock()
278                .expect("background manager poisoned");
279            let agent = state.agents.entry(agent_id.to_string()).or_default();
280            if let Some(existing) = agent.tasks.iter_mut().find(|task| task.id == summary.id) {
281                *existing = summary.clone();
282            } else {
283                agent.tasks.push(summary.clone());
284            }
285            (agent.observer.clone(), agent.tasks.clone())
286        };
287        let _ = self
288            .inner
289            .store
290            .upsert_background_task(agent_id, &summary, NOTIFICATION_PENDING);
291        let status = summary.status.to_string();
292        let _ = self
293            .inner
294            .hooks
295            .task_finished(agent_id, &summary.id, &status);
296
297        self.publish_observer(
298            observer,
299            tasks,
300            AgentEvent::BackgroundTaskFinished { task: summary },
301        );
302    }
303    fn publish_observer(
304        &self,
305        observer: Option<BackgroundObserver>,
306        tasks: Vec<BackgroundTaskSummary>,
307        event: AgentEvent,
308    ) {
309        let Some(observer) = observer else {
310            return;
311        };
312
313        observer.sink.publish_snapshot(&tasks);
314        observer.sink.publish_event(event);
315    }
316}
317
318struct CompletedBackgroundTask {
319    id: String,
320    command: String,
321    cwd: PathBuf,
322    status: BackgroundTaskStatus,
323    output_preview: String,
324}
325
326async fn execute_task(
327    id: String,
328    request: CommandRequest,
329    executor: Arc<dyn RuntimeExecutor>,
330) -> CompletedBackgroundTask {
331    let command = request.spec.display().to_string();
332    let cwd = request.cwd.clone();
333    match executor.run(request).await {
334        Ok(output) => completed_task_from_output(id, command, cwd, output),
335        Err(error) => CompletedBackgroundTask {
336            id,
337            command,
338            cwd,
339            status: BackgroundTaskStatus::Failed,
340            output_preview: truncate_preview(&error),
341        },
342    }
343}
344
345fn completed_task_from_output(
346    id: String,
347    command: String,
348    cwd: PathBuf,
349    output: CommandOutput,
350) -> CompletedBackgroundTask {
351    let combined = format!("{} {}", output.stdout, output.stderr);
352    let preview = if combined.trim().is_empty() {
353        "(no output)".to_string()
354    } else {
355        truncate_preview(&combined)
356    };
357    let status = if output.success() {
358        BackgroundTaskStatus::Finished
359    } else {
360        BackgroundTaskStatus::Failed
361    };
362
363    CompletedBackgroundTask {
364        id,
365        command,
366        cwd,
367        status,
368        output_preview: preview,
369    }
370}
371
372fn truncate_preview(text: &str) -> String {
373    let mut compact = String::new();
374    for (index, chunk) in text.split_whitespace().enumerate() {
375        if index > 0 {
376            compact.push(' ');
377        }
378        compact.push_str(chunk);
379    }
380
381    let mut truncated = compact
382        .chars()
383        .take(OUTPUT_PREVIEW_MAX_CHARS)
384        .collect::<String>();
385    if compact.chars().count() > OUTPUT_PREVIEW_MAX_CHARS {
386        truncated.push_str("...");
387    }
388    truncated
389}
390
391fn render_task_summary(task: &BackgroundTaskSummary) -> String {
392    format!(
393        "{}: [{}] cwd={} {}",
394        task.id,
395        task.status,
396        task.cwd.display(),
397        task.command
398    )
399}
400
401fn render_task_detail(task: &BackgroundTaskSummary) -> String {
402    let output = task.output_preview.as_deref().unwrap_or("(running)");
403    format!(
404        "[{}] cwd={}\n{}\n{}",
405        task.status,
406        task.cwd.display(),
407        task.command,
408        output
409    )
410}