Skip to main content

aether_core/testing/
fake_agent_observer.rs

1use crate::events::{AgentEvent, AgentObserver, TraceContext};
2use std::sync::{Arc, Mutex};
3
4/// In-memory [`AgentObserver`] that records every event it receives, for
5/// asserting on the stream an agent emits.
6#[derive(Default)]
7pub struct FakeAgentObserver {
8    events: Arc<Mutex<Vec<AgentEvent>>>,
9    system_prompts: Arc<Mutex<Vec<String>>>,
10    panic_on_event: Option<fn(&AgentEvent) -> bool>,
11    panic_on_system_prompt: bool,
12    panic_on_tool_trace_context: bool,
13}
14
15impl FakeAgentObserver {
16    pub fn new() -> Self {
17        Self::default()
18    }
19
20    pub fn with_event_panic(mut self, predicate: fn(&AgentEvent) -> bool) -> Self {
21        self.panic_on_event = Some(predicate);
22        self
23    }
24
25    pub fn with_system_prompt_panic(mut self) -> Self {
26        self.panic_on_system_prompt = true;
27        self
28    }
29
30    pub fn with_tool_trace_context_panic(mut self) -> Self {
31        self.panic_on_tool_trace_context = true;
32        self
33    }
34
35    /// Shared handle to the recorded events; clones observe future events too.
36    pub fn events(&self) -> Arc<Mutex<Vec<AgentEvent>>> {
37        Arc::clone(&self.events)
38    }
39
40    /// Shared handle to the system prompts reported for each LLM request.
41    pub fn system_prompts(&self) -> Arc<Mutex<Vec<String>>> {
42        Arc::clone(&self.system_prompts)
43    }
44}
45
46impl AgentObserver for FakeAgentObserver {
47    fn on_event(&mut self, message: &AgentEvent) {
48        self.events.lock().unwrap().push(message.clone());
49        assert!(!self.panic_on_event.is_some_and(|predicate| predicate(message)), "observer exploded");
50    }
51
52    fn on_system_prompt(&mut self, prompt: &str) {
53        self.system_prompts.lock().unwrap().push(prompt.to_string());
54        assert!(!self.panic_on_system_prompt, "system prompt observer exploded");
55    }
56
57    fn tool_trace_context(&self, _tool_id: &str) -> Option<TraceContext> {
58        assert!(!self.panic_on_tool_trace_context, "tool trace observer exploded");
59        None
60    }
61}