Skip to main content

bamboo_engine/runtime/hooks/
mod.rs

1//! Hook runner — dispatches registered hooks at lifecycle points.
2
3mod shell_command;
4
5use std::sync::Arc;
6
7use bamboo_agent_core::{AgentError, AgentEvent, AgentHook, Message, Session};
8use bamboo_domain::{
9    AgentHookPoint, AgentRuntimeState, AgentStatusState, HookCheckpoint, HookPayload, HookResult,
10    SessionEndStatus, SuspensionState,
11};
12use chrono::Utc;
13use tokio::sync::mpsc;
14
15pub use shell_command::{
16    test_lifecycle_shell_command, ShellCommandHook, ShellHookEvent, ShellHookTestOutput,
17};
18
19/// Aggregate output from every hook registered at one seam.
20#[derive(Debug, Clone, PartialEq, Eq)]
21pub struct HookRunOutcome {
22    pub decision: HookResult,
23    pub injected_contexts: Vec<String>,
24}
25
26impl Default for HookRunOutcome {
27    fn default() -> Self {
28        Self {
29            decision: HookResult::Continue,
30            injected_contexts: Vec::new(),
31        }
32    }
33}
34
35/// Runs registered hooks at a given hook point.
36#[derive(Clone)]
37pub struct HookRunner {
38    hooks: Vec<Arc<dyn AgentHook>>,
39}
40
41impl HookRunner {
42    pub fn new() -> Self {
43        Self { hooks: Vec::new() }
44    }
45
46    /// Register a hook. Hooks are sorted by priority (lower runs first).
47    pub fn register(&mut self, hook: Arc<dyn AgentHook>) {
48        self.hooks.push(hook);
49        self.hooks.sort_by_key(|h| h.priority());
50    }
51
52    /// Clone this registry and append shell hooks from one frozen config
53    /// snapshot. The original registry remains reusable by future runs.
54    pub fn with_lifecycle_config(
55        &self,
56        config: &bamboo_config::LifecycleHooksConfig,
57        fallback_cwd: Option<std::path::PathBuf>,
58    ) -> Self {
59        let mut runner = self.clone();
60        shell_command::register_configured_shell_hooks(&mut runner, config, fallback_cwd);
61        runner
62    }
63
64    /// Run all hooks matching the given point.
65    ///
66    /// Records checkpoints in `runtime_state`. Returns the first
67    /// `Suspend` or `Abort` result, or the aggregate result otherwise.
68    pub async fn run_hooks(
69        &self,
70        point: AgentHookPoint,
71        payload: &HookPayload,
72        session: &Session,
73        runtime_state: &mut AgentRuntimeState,
74        event_tx: Option<&mpsc::Sender<AgentEvent>>,
75    ) -> HookRunOutcome {
76        self.run_hooks_with_control(point, payload, session, runtime_state, event_tx, true)
77            .await
78    }
79
80    /// Run every matching hook while recording checkpoints/events, but never
81    /// short-circuit on a control decision. Observer/advisory seams such as
82    /// `SessionEnd`, `PreCompact`, and `Notification` use this so a command's
83    /// control-shaped output cannot suppress later hooks or reverse an
84    /// operation that must proceed for correctness.
85    pub async fn run_observer_hooks(
86        &self,
87        point: AgentHookPoint,
88        payload: &HookPayload,
89        session: &Session,
90        runtime_state: &mut AgentRuntimeState,
91        event_tx: Option<&mpsc::Sender<AgentEvent>>,
92    ) -> HookRunOutcome {
93        self.run_hooks_with_control(point, payload, session, runtime_state, event_tx, false)
94            .await
95    }
96
97    async fn run_hooks_with_control(
98        &self,
99        point: AgentHookPoint,
100        payload: &HookPayload,
101        session: &Session,
102        runtime_state: &mut AgentRuntimeState,
103        event_tx: Option<&mpsc::Sender<AgentEvent>>,
104        honor_control_decisions: bool,
105    ) -> HookRunOutcome {
106        let mut outcome = HookRunOutcome::default();
107
108        for hook in &self.hooks {
109            if hook.point() != point || !hook.matches(payload) {
110                continue;
111            }
112
113            let start = std::time::Instant::now();
114            let result = hook.run(point, payload, session).await;
115            let elapsed = start.elapsed();
116
117            runtime_state.checkpoints.push(HookCheckpoint {
118                hook_point: format!("{:?}", point),
119                timestamp: Utc::now(),
120                result: format!("{:?}", result),
121                duration_ms: elapsed.as_millis() as u64,
122            });
123
124            if let Some(event_tx) = event_tx {
125                let _ = event_tx
126                    .send(AgentEvent::HookLifecycle {
127                        hook_name: hook.name().to_string(),
128                        point,
129                        phase: "completed".to_string(),
130                        duration_ms: elapsed.as_millis() as u64,
131                        decision: result.clone(),
132                    })
133                    .await;
134            }
135
136            let (result, mut contexts) = unwrap_context_result(result);
137            outcome.injected_contexts.append(&mut contexts);
138
139            match &result {
140                HookResult::Abort { .. }
141                | HookResult::Suspend { .. }
142                | HookResult::Deny { .. }
143                | HookResult::Ask => {
144                    if honor_control_decisions {
145                        outcome.decision = result;
146                        return outcome;
147                    }
148                }
149                HookResult::InjectContext { text } => {
150                    outcome.injected_contexts.push(text.clone());
151                }
152                HookResult::Mutated => {
153                    if matches!(outcome.decision, HookResult::Continue) {
154                        outcome.decision = HookResult::Mutated;
155                    }
156                }
157                HookResult::Allow => outcome.decision = HookResult::Allow,
158                HookResult::Continue => {}
159                HookResult::WithContext { .. } => unreachable!("context results are unwrapped"),
160            }
161        }
162
163        outcome
164    }
165
166    /// Check if any hooks are registered for the given point.
167    pub fn has_hooks_for(&self, point: AgentHookPoint) -> bool {
168        self.hooks.iter().any(|h| h.point() == point)
169    }
170
171    /// Number of registered hooks.
172    pub fn len(&self) -> usize {
173        self.hooks.len()
174    }
175
176    /// Whether any hooks are registered.
177    pub fn is_empty(&self) -> bool {
178        self.hooks.is_empty()
179    }
180}
181
182/// Fire cleanup/notification hooks after a terminal run. Decisions and context
183/// are intentionally ignored: `SessionEnd` observes a settled outcome and may
184/// not reverse it. Suspended runs are non-terminal and do not fire this event.
185pub(crate) async fn run_session_end_hooks(
186    runner: &HookRunner,
187    result: &Result<(), AgentError>,
188    session: &mut Session,
189    event_tx: &mpsc::Sender<AgentEvent>,
190) {
191    let suspended_non_terminal = result.is_ok()
192        && session
193            .metadata
194            .get("runtime.suspend_reason")
195            .is_some_and(|reason| !reason.trim().is_empty());
196    if suspended_non_terminal || !runner.has_hooks_for(AgentHookPoint::AfterSessionEnd) {
197        return;
198    }
199
200    let (status, completion_reason) = match result {
201        Ok(()) => (
202            SessionEndStatus::Completed,
203            session
204                .metadata
205                .get("runtime.completion_reason")
206                .cloned()
207                .or_else(|| Some("completed".to_string())),
208        ),
209        Err(error) if error.is_cancelled() => {
210            (SessionEndStatus::Cancelled, Some(error.to_string()))
211        }
212        Err(error) => (SessionEndStatus::Failed, Some(error.to_string())),
213    };
214    let mut runtime_state = session
215        .agent_runtime_state
216        .clone()
217        .unwrap_or_else(|| AgentRuntimeState::new(&session.id));
218    runner
219        .run_observer_hooks(
220            AgentHookPoint::AfterSessionEnd,
221            &HookPayload::SessionEnd {
222                status,
223                completion_reason,
224            },
225            session,
226            &mut runtime_state,
227            Some(event_tx),
228        )
229        .await;
230    session.agent_runtime_state = Some(runtime_state);
231}
232
233fn unwrap_context_result(mut result: HookResult) -> (HookResult, Vec<String>) {
234    let mut contexts = Vec::new();
235    while let HookResult::WithContext {
236        result: inner,
237        text,
238    } = result
239    {
240        if !text.trim().is_empty() {
241            contexts.push(text);
242        }
243        result = *inner;
244    }
245    (result, contexts)
246}
247
248/// Apply context injections and non-tool control decisions consistently across
249/// lifecycle seams.
250pub(crate) fn apply_hook_outcome(
251    point: AgentHookPoint,
252    outcome: HookRunOutcome,
253    session: &mut Session,
254    runtime_state: &mut AgentRuntimeState,
255) -> Result<(), AgentError> {
256    if matches!(point, AgentHookPoint::AfterSessionSetup) {
257        runtime_state.hook_contexts.extend(
258            outcome
259                .injected_contexts
260                .into_iter()
261                .filter(|text| !text.trim().is_empty()),
262        );
263    } else {
264        inject_contexts(session, point, outcome.injected_contexts);
265    }
266
267    match outcome.decision {
268        HookResult::Continue
269        | HookResult::Mutated
270        | HookResult::Allow
271        | HookResult::InjectContext { .. } => Ok(()),
272        HookResult::Suspend { reason } => {
273            let hook_point = format!("{point:?}");
274            runtime_state.status = AgentStatusState::Suspended;
275            runtime_state.suspension = Some(SuspensionState {
276                reason: reason.clone(),
277                suspended_at: Utc::now(),
278                resumable: true,
279                hook_point: Some(hook_point.clone()),
280            });
281            session.metadata.insert(
282                "runtime.suspend_reason".to_string(),
283                "hook_suspended".to_string(),
284            );
285            Err(AgentError::HookSuspended(format!("{hook_point}: {reason}")))
286        }
287        HookResult::Abort { reason } => Err(AgentError::Tool(format!(
288            "hook aborted at {point:?}: {reason}"
289        ))),
290        HookResult::Deny { reason } => Err(AgentError::Tool(format!(
291            "hook denied lifecycle seam {point:?}: {reason}"
292        ))),
293        HookResult::Ask => Err(AgentError::Tool(format!(
294            "hook requested parent approval at non-tool seam {point:?}"
295        ))),
296        HookResult::WithContext { result, text } => apply_hook_outcome(
297            point,
298            HookRunOutcome {
299                decision: *result,
300                injected_contexts: vec![text],
301            },
302            session,
303            runtime_state,
304        ),
305    }
306}
307
308pub(crate) fn inject_contexts(
309    session: &mut Session,
310    point: AgentHookPoint,
311    injected_contexts: Vec<String>,
312) {
313    for text in injected_contexts {
314        if text.trim().is_empty() {
315            continue;
316        }
317        let block =
318            format!("\n\n<agent_hook_context point=\"{point:?}\">\n{text}\n</agent_hook_context>");
319        if let Some(system_message) = session
320            .messages
321            .iter_mut()
322            .find(|message| matches!(message.role, bamboo_agent_core::Role::System))
323        {
324            system_message.content.push_str(&block);
325            system_message.never_compress = true;
326        } else {
327            let mut message = Message::system(block.trim().to_string());
328            message.never_compress = true;
329            message.metadata = Some(serde_json::json!({
330                "runtime_kind": "hook_context",
331                "hook_point": point,
332            }));
333            session.add_message(message);
334        }
335    }
336}
337
338/// Merge hook checkpoints produced through a session-local seam (notably
339/// compression) into the runner-owned state without losing checkpoints written
340/// directly by tool/round seams.
341pub(crate) fn merge_session_hook_checkpoints(
342    session: &Session,
343    runtime_state: &mut AgentRuntimeState,
344) {
345    let Some(session_state) = session.agent_runtime_state.as_ref() else {
346        return;
347    };
348    for checkpoint in &session_state.checkpoints {
349        if !runtime_state.checkpoints.contains(checkpoint) {
350            runtime_state.checkpoints.push(checkpoint.clone());
351        }
352    }
353    if matches!(session_state.status, AgentStatusState::Suspended) {
354        runtime_state.status = AgentStatusState::Suspended;
355        runtime_state.suspension = session_state.suspension.clone();
356    }
357}
358
359impl Default for HookRunner {
360    fn default() -> Self {
361        Self::new()
362    }
363}
364
365#[cfg(test)]
366mod tests {
367    use super::*;
368
369    /// A no-op hook that always returns Continue.
370    struct ContinueHook {
371        point: AgentHookPoint,
372        pri: u32,
373        name: String,
374    }
375
376    #[async_trait::async_trait]
377    impl AgentHook for ContinueHook {
378        fn point(&self) -> AgentHookPoint {
379            self.point
380        }
381
382        async fn run(
383            &self,
384            _point: AgentHookPoint,
385            _payload: &HookPayload,
386            _session: &Session,
387        ) -> HookResult {
388            HookResult::Continue
389        }
390
391        fn priority(&self) -> u32 {
392            self.pri
393        }
394
395        fn name(&self) -> &str {
396            &self.name
397        }
398    }
399
400    /// A hook that always returns Abort.
401    struct AbortHook;
402
403    #[async_trait::async_trait]
404    impl AgentHook for AbortHook {
405        fn point(&self) -> AgentHookPoint {
406            AgentHookPoint::BeforeLlmCall
407        }
408
409        async fn run(
410            &self,
411            _point: AgentHookPoint,
412            _payload: &HookPayload,
413            _session: &Session,
414        ) -> HookResult {
415            HookResult::Abort {
416                reason: "test abort".to_string(),
417            }
418        }
419
420        fn name(&self) -> &str {
421            "abort_hook"
422        }
423    }
424
425    fn test_session() -> Session {
426        Session::new("test", "test-model")
427    }
428
429    #[tokio::test]
430    async fn empty_runner_returns_continue() {
431        let runner = HookRunner::new();
432        let mut state = AgentRuntimeState::new("run-1");
433        let session = test_session();
434        let (tx, _rx) = mpsc::channel(4);
435
436        let result = runner
437            .run_hooks(
438                AgentHookPoint::BeforeRound,
439                &HookPayload::Round { round: 1 },
440                &session,
441                &mut state,
442                Some(&tx),
443            )
444            .await;
445
446        assert_eq!(result.decision, HookResult::Continue);
447        assert!(state.checkpoints.is_empty());
448    }
449
450    #[tokio::test]
451    async fn hooks_run_in_priority_order() {
452        let mut runner = HookRunner::new();
453        runner.register(Arc::new(ContinueHook {
454            point: AgentHookPoint::BeforeRound,
455            pri: 200,
456            name: "slow".to_string(),
457        }));
458        runner.register(Arc::new(ContinueHook {
459            point: AgentHookPoint::BeforeRound,
460            pri: 50,
461            name: "fast".to_string(),
462        }));
463
464        let mut state = AgentRuntimeState::new("run-2");
465        let session = test_session();
466        let (tx, mut rx) = mpsc::channel(4);
467
468        let result = runner
469            .run_hooks(
470                AgentHookPoint::BeforeRound,
471                &HookPayload::Round { round: 1 },
472                &session,
473                &mut state,
474                Some(&tx),
475            )
476            .await;
477
478        assert_eq!(result.decision, HookResult::Continue);
479        assert_eq!(state.checkpoints.len(), 2);
480        // Lower priority runs first
481        assert!(state.checkpoints[0].result.contains("Continue"));
482        assert!(matches!(
483            rx.recv().await,
484            Some(AgentEvent::HookLifecycle { hook_name, .. }) if hook_name == "fast"
485        ));
486    }
487
488    #[tokio::test]
489    async fn abort_short_circuits() {
490        let mut runner = HookRunner::new();
491        runner.register(Arc::new(AbortHook));
492
493        let mut state = AgentRuntimeState::new("run-3");
494        let session = test_session();
495        let (tx, _rx) = mpsc::channel(4);
496
497        let result = runner
498            .run_hooks(
499                AgentHookPoint::BeforeLlmCall,
500                &HookPayload::None,
501                &session,
502                &mut state,
503                Some(&tx),
504            )
505            .await;
506
507        assert!(matches!(result.decision, HookResult::Abort { .. }));
508        assert_eq!(state.checkpoints.len(), 1);
509    }
510
511    #[tokio::test]
512    async fn wrong_point_hooks_are_skipped() {
513        let mut runner = HookRunner::new();
514        runner.register(Arc::new(AbortHook)); // registered for BeforeLlmCall
515
516        let mut state = AgentRuntimeState::new("run-4");
517        let session = test_session();
518        let (tx, _rx) = mpsc::channel(4);
519
520        let result = runner
521            .run_hooks(
522                AgentHookPoint::AfterRound,
523                &HookPayload::Round { round: 1 },
524                &session,
525                &mut state,
526                Some(&tx),
527            )
528            .await;
529
530        assert_eq!(result.decision, HookResult::Continue);
531        assert!(state.checkpoints.is_empty());
532    }
533
534    struct RecordingSessionEndHook {
535        payloads: Arc<std::sync::Mutex<Vec<HookPayload>>>,
536    }
537
538    #[async_trait::async_trait]
539    impl AgentHook for RecordingSessionEndHook {
540        fn point(&self) -> AgentHookPoint {
541            AgentHookPoint::AfterSessionEnd
542        }
543
544        async fn run(
545            &self,
546            _point: AgentHookPoint,
547            payload: &HookPayload,
548            _session: &Session,
549        ) -> HookResult {
550            self.payloads.lock().unwrap().push(payload.clone());
551            // Decisions at SessionEnd are observability-only and must not
552            // change the already-settled terminal outcome.
553            HookResult::Deny {
554                reason: "ignored cleanup decision".to_string(),
555            }
556        }
557    }
558
559    #[tokio::test]
560    async fn session_end_fires_for_completed_failed_and_cancelled_and_ignores_decisions() {
561        for (result, expected_status) in [
562            (Ok(()), SessionEndStatus::Completed),
563            (
564                Err(AgentError::Tool("terminal failure".to_string())),
565                SessionEndStatus::Failed,
566            ),
567            (Err(AgentError::Cancelled), SessionEndStatus::Cancelled),
568        ] {
569            let payloads = Arc::new(std::sync::Mutex::new(Vec::new()));
570            let mut runner = HookRunner::new();
571            runner.register(Arc::new(RecordingSessionEndHook {
572                payloads: payloads.clone(),
573            }));
574            runner.register(Arc::new(RecordingSessionEndHook {
575                payloads: payloads.clone(),
576            }));
577            let mut session = test_session();
578            let (tx, _rx) = mpsc::channel(4);
579
580            run_session_end_hooks(&runner, &result, &mut session, &tx).await;
581
582            let recorded = payloads.lock().unwrap();
583            assert_eq!(
584                recorded.len(),
585                2,
586                "a denied observer must not suppress later cleanup hooks"
587            );
588            assert!(recorded.iter().all(|payload| matches!(
589                payload,
590                HookPayload::SessionEnd { status, .. } if *status == expected_status
591            )));
592            assert_eq!(
593                session
594                    .agent_runtime_state
595                    .as_ref()
596                    .map(|state| state.checkpoints.len()),
597                Some(2)
598            );
599        }
600    }
601
602    #[tokio::test]
603    async fn session_end_skips_suspended_non_terminal_runs() {
604        let payloads = Arc::new(std::sync::Mutex::new(Vec::new()));
605        let mut runner = HookRunner::new();
606        runner.register(Arc::new(RecordingSessionEndHook {
607            payloads: payloads.clone(),
608        }));
609        let mut session = test_session();
610        session.metadata.insert(
611            "runtime.suspend_reason".to_string(),
612            "waiting_for_children".to_string(),
613        );
614        let (tx, _rx) = mpsc::channel(4);
615
616        run_session_end_hooks(&runner, &Ok(()), &mut session, &tx).await;
617
618        assert!(payloads.lock().unwrap().is_empty());
619        assert!(session.agent_runtime_state.is_none());
620    }
621}