1mod 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#[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#[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 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 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 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 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 pub fn has_hooks_for(&self, point: AgentHookPoint) -> bool {
168 self.hooks.iter().any(|h| h.point() == point)
169 }
170
171 pub fn len(&self) -> usize {
173 self.hooks.len()
174 }
175
176 pub fn is_empty(&self) -> bool {
178 self.hooks.is_empty()
179 }
180}
181
182pub(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
248pub(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
338pub(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 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 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 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)); 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 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}