Skip to main content

aether_core/testing/
utils.rs

1use std::collections::BTreeMap;
2use std::sync::{Arc, Mutex};
3use std::time::Duration;
4use tokio::sync::{Notify, mpsc};
5
6use crate::context::CompactionConfig;
7use crate::core::{AgentError, Prompt, RetryConfig, agent};
8use crate::events::{
9    AgentCommand, AgentEvent, AgentObserver, Command, ContextEvent, ToolEvent, TurnEvent, UserCommand,
10};
11use crate::mcp::mcp;
12use crate::testing::{AgentTrace, FakeAgentObserver, FakeMcpServer, McpBuilderTestExt};
13use llm::{ChatMessage, Context, LlmError, LlmModel, LlmResponse, ModelSettings, StreamingModelProvider};
14
15use llm::testing::FakeLlmProvider;
16
17pub async fn drain_until(
18    receiver: &mut mpsc::Receiver<AgentEvent>,
19    predicate: impl Fn(&AgentEvent) -> bool,
20) -> Vec<AgentEvent> {
21    let mut events = Vec::new();
22    while let Some(event) = receiver.recv().await {
23        let matched = predicate(&event);
24        events.push(event);
25        if matched {
26            return events;
27        }
28    }
29    panic!("agent event channel closed before predicate matched");
30}
31
32pub fn content_events(events: Vec<AgentEvent>) -> Vec<AgentEvent> {
33    events
34        .into_iter()
35        .filter(|event| {
36            !matches!(
37                event,
38                AgentEvent::Turn(
39                    TurnEvent::Started { .. }
40                        | TurnEvent::UserMessageInserted { .. }
41                        | TurnEvent::UserMessageDiscarded { .. }
42                        | TurnEvent::LlmCallStarted { .. }
43                        | TurnEvent::LlmCallEnded { .. }
44                ) | AgentEvent::Tool(ToolEvent::DefinitionsUpdated { .. })
45            )
46        })
47        .collect()
48}
49
50pub fn mcp_instructions(entries: &[(&str, &str)]) -> BTreeMap<String, String> {
51    entries.iter().map(|(k, v)| ((*k).to_string(), (*v).to_string())).collect()
52}
53
54/// Millisecond-scale retry delays so retry tests run fast under virtual time.
55pub fn fast_retry(max_attempts: u32) -> RetryConfig {
56    RetryConfig { max_attempts, base_delay: Duration::from_millis(1), max_delay: Duration::from_millis(5) }
57}
58
59pub fn test_agent() -> TestAgentBuilder {
60    TestAgentBuilder::new()
61}
62
63/// An ordered interaction with a test agent.
64pub enum TestAgentStep {
65    Send(Command),
66    WaitFor(Box<dyn Fn(&AgentEvent) -> bool + Send>),
67    /// Run an arbitrary side effect (e.g. releasing a paused LLM stream) between steps.
68    Perform(Box<dyn FnOnce() + Send>),
69}
70
71impl TestAgentStep {
72    pub fn send(command: Command) -> Self {
73        Self::Send(command)
74    }
75
76    pub fn user_text(text: impl Into<String>) -> Self {
77        Self::send(Command::text(text))
78    }
79
80    pub fn cancel() -> Self {
81        Self::send(Command::UserCommand(UserCommand::Cancel))
82    }
83
84    pub fn switch_model(provider: impl StreamingModelProvider + 'static) -> Self {
85        Self::send(Command::AgentCommand(AgentCommand::SwitchModel(Box::new(provider))))
86    }
87
88    pub fn replace_conversation(messages: Vec<ChatMessage>) -> Self {
89        Self::send(Command::AgentCommand(AgentCommand::ReplaceConversation(messages)))
90    }
91
92    pub fn perform(action: impl FnOnce() + Send + 'static) -> Self {
93        Self::Perform(Box::new(action))
94    }
95
96    pub fn wait_for(predicate: impl Fn(&AgentEvent) -> bool + Send + 'static) -> Self {
97        Self::WaitFor(Box::new(predicate))
98    }
99
100    pub fn wait_for_turn_end() -> Self {
101        Self::wait_for(|event| matches!(event, AgentEvent::Turn(TurnEvent::Ended { .. })))
102    }
103
104    pub fn wait_for_compaction_start() -> Self {
105        Self::wait_for(|event| matches!(event, AgentEvent::Context(ContextEvent::CompactionStarted { .. })))
106    }
107
108    pub fn wait_for_retry(attempt: u32) -> Self {
109        Self::wait_for(
110            move |event| matches!(event, AgentEvent::Turn(TurnEvent::RetryScheduled { attempt: actual, .. }) if *actual == attempt),
111        )
112    }
113}
114
115/// A fluent sequence of commands and synchronization points for a test agent.
116#[derive(Default)]
117pub struct TestScenario {
118    steps: Vec<TestAgentStep>,
119}
120
121impl TestScenario {
122    pub fn new() -> Self {
123        Self::default()
124    }
125
126    pub fn send(mut self, command: Command) -> Self {
127        self.steps.push(TestAgentStep::send(command));
128        self
129    }
130
131    pub fn user_text(mut self, text: impl Into<String>) -> Self {
132        self.steps.push(TestAgentStep::user_text(text));
133        self
134    }
135
136    pub fn cancel(mut self) -> Self {
137        self.steps.push(TestAgentStep::cancel());
138        self
139    }
140
141    pub fn switch_model(mut self, provider: impl StreamingModelProvider + 'static) -> Self {
142        self.steps.push(TestAgentStep::switch_model(provider));
143        self
144    }
145
146    pub fn replace_conversation(mut self, messages: Vec<ChatMessage>) -> Self {
147        self.steps.push(TestAgentStep::replace_conversation(messages));
148        self
149    }
150
151    pub fn wait_for(mut self, predicate: impl Fn(&AgentEvent) -> bool + Send + 'static) -> Self {
152        self.steps.push(TestAgentStep::wait_for(predicate));
153        self
154    }
155
156    pub fn wait_for_turn_end(mut self) -> Self {
157        self.steps.push(TestAgentStep::wait_for_turn_end());
158        self
159    }
160
161    pub fn wait_for_compaction_start(mut self) -> Self {
162        self.steps.push(TestAgentStep::wait_for_compaction_start());
163        self
164    }
165
166    pub fn wait_for_retry(mut self, attempt: u32) -> Self {
167        self.steps.push(TestAgentStep::wait_for_retry(attempt));
168        self
169    }
170
171    /// Run an arbitrary side effect between scenario steps, e.g. releasing a
172    /// paused LLM stream so a queued message can be injected mid-turn.
173    pub fn perform(mut self, action: impl FnOnce() + Send + 'static) -> Self {
174        self.steps.push(TestAgentStep::perform(action));
175        self
176    }
177}
178
179impl From<Vec<TestAgentStep>> for TestScenario {
180    fn from(steps: Vec<TestAgentStep>) -> Self {
181        Self { steps }
182    }
183}
184
185/// Result of running a test agent, including messages and captured contexts.
186pub struct TestAgentResult {
187    pub messages: Vec<AgentEvent>,
188    pub captured_contexts: Arc<Mutex<Vec<Context>>>,
189}
190
191/// Error returned by [`TestAgentBuilder::run`], [`TestAgentBuilder::run_trace`], and
192/// [`TestAgentBuilder::run_with_context`].
193#[derive(Debug, thiserror::Error)]
194pub enum TestAgentError {
195    #[error(transparent)]
196    Agent(#[from] AgentError),
197    #[error("failed to send command to the test agent: {0}")]
198    SendCommand(#[from] mpsc::error::SendError<Command>),
199}
200
201pub type TestResult<T> = std::result::Result<T, TestAgentError>;
202
203struct ProviderTestConfig {
204    responses: Vec<Vec<Result<LlmResponse, LlmError>>>,
205    model: Option<LlmModel>,
206    context_window: Option<u32>,
207    pause: Option<(usize, usize, Arc<Notify>)>,
208}
209
210struct AgentTestConfig {
211    context_window_override: Option<u32>,
212    timeout: Option<Duration>,
213    max_auto_continues: Option<u32>,
214    retry_config: Option<RetryConfig>,
215    observers: Vec<Box<dyn AgentObserver>>,
216    mcp_server: Option<(String, FakeMcpServer)>,
217    initial_messages: Vec<ChatMessage>,
218    system_prompts: Vec<Prompt>,
219    session_affinity_key: Option<String>,
220    compaction: Option<CompactionConfig>,
221    model_settings: Option<ModelSettings>,
222}
223
224enum TestExecution {
225    CommandsUntilTurnEnd(Vec<Command>),
226    Scenario(TestScenario),
227}
228
229pub struct TestAgentBuilder {
230    provider: ProviderTestConfig,
231    agent: AgentTestConfig,
232    execution: Option<TestExecution>,
233}
234
235impl Default for TestAgentBuilder {
236    fn default() -> Self {
237        Self::new()
238    }
239}
240
241impl TestAgentBuilder {
242    pub fn new() -> Self {
243        Self {
244            provider: ProviderTestConfig { responses: Vec::new(), model: None, context_window: None, pause: None },
245            agent: AgentTestConfig {
246                context_window_override: None,
247                timeout: None,
248                max_auto_continues: None,
249                retry_config: None,
250                observers: Vec::new(),
251                mcp_server: Some(("test".to_string(), FakeMcpServer::new())),
252                initial_messages: Vec::new(),
253                system_prompts: Vec::new(),
254                session_affinity_key: None,
255                compaction: None,
256                model_settings: None,
257            },
258            execution: None,
259        }
260    }
261
262    pub fn commands(self, commands: Vec<Command>) -> Self {
263        self.with_execution(TestExecution::CommandsUntilTurnEnd(commands))
264    }
265
266    pub fn scenario(self, scenario: impl Into<TestScenario>) -> Self {
267        self.with_execution(TestExecution::Scenario(scenario.into()))
268    }
269
270    pub fn user_text(self, text: &str) -> Self {
271        self.commands(vec![Command::text(text)])
272    }
273
274    pub fn llm_responses(mut self, llm_responses: &[Vec<LlmResponse>]) -> Self {
275        self.provider.responses = llm_responses.iter().map(|turn| turn.iter().cloned().map(Ok).collect()).collect();
276        self
277    }
278
279    pub fn llm_result_responses(mut self, llm_responses: &[Vec<Result<LlmResponse, LlmError>>]) -> Self {
280        self.provider.responses = Vec::from(llm_responses);
281        self
282    }
283
284    pub fn model(mut self, model: LlmModel) -> Self {
285        self.provider.model = Some(model);
286        self
287    }
288
289    pub fn provider_context_window(mut self, window: Option<u32>) -> Self {
290        self.provider.context_window = window;
291        self
292    }
293
294    pub fn context_window_override(mut self, window: u32) -> Self {
295        self.agent.context_window_override = Some(window);
296        self
297    }
298
299    pub fn tool_timeout(mut self, timeout: Duration) -> Self {
300        self.agent.timeout = Some(timeout);
301        self
302    }
303
304    pub fn max_auto_continues(mut self, max: u32) -> Self {
305        self.agent.max_auto_continues = Some(max);
306        self
307    }
308
309    pub fn retry_config(mut self, config: RetryConfig) -> Self {
310        self.agent.retry_config = Some(config);
311        self
312    }
313
314    /// Run without the default fake MCP server when the scenario does not exercise tools.
315    pub fn without_mcp(mut self) -> Self {
316        self.agent.mcp_server = None;
317        self
318    }
319
320    /// Replace the default fake MCP server with a scripted server.
321    pub fn fake_mcp_server(mut self, name: &str, server: FakeMcpServer) -> Self {
322        self.agent.mcp_server = Some((name.to_string(), server));
323        self
324    }
325
326    /// Pre-populate the context with conversation history.
327    pub fn messages(mut self, messages: Vec<ChatMessage>) -> Self {
328        self.agent.initial_messages = messages;
329        self
330    }
331
332    /// Set the system prompt. Multiple prompts are concatenated with double
333    /// newlines, mirroring [`crate::core::AgentBuilder::system_prompt`].
334    pub fn system_prompt(mut self, prompt: Prompt) -> Self {
335        self.agent.system_prompts.push(prompt);
336        self
337    }
338
339    /// Set the session affinity key attached to every LLM call; defaults to a
340    /// fresh random key per agent.
341    pub fn session_affinity_key(mut self, key: impl Into<String>) -> Self {
342        self.agent.session_affinity_key = Some(key.into());
343        self
344    }
345
346    /// Configure context compaction settings.
347    pub fn compaction_config(mut self, config: CompactionConfig) -> Self {
348        self.agent.compaction = Some(config);
349        self
350    }
351
352    /// Set the model settings applied to every LLM call.
353    pub fn model_settings(mut self, settings: ModelSettings) -> Self {
354        self.agent.model_settings = Some(settings);
355        self
356    }
357
358    /// Pause the fake LLM stream at `turn_index` / `chunk_index` until
359    /// `release.notify_one()` is called. Used for deterministic timing tests.
360    pub fn pause_turn_after(mut self, turn_index: usize, chunk_index: usize, release: Arc<Notify>) -> Self {
361        self.provider.pause = Some((turn_index, chunk_index, release));
362        self
363    }
364
365    /// Attach an observer of the test agent's event stream.
366    pub fn observer(mut self, observer: Box<dyn AgentObserver>) -> Self {
367        self.agent.observers.push(observer);
368        self
369    }
370
371    pub async fn run(self) -> TestResult<Vec<AgentEvent>> {
372        let result = self.run_with_context().await?;
373        Ok(result.messages)
374    }
375
376    /// Runs the test agent with a recording observer attached and returns the
377    /// full event trace, including internal events.
378    pub async fn run_trace(self) -> TestResult<AgentTrace> {
379        let observer = FakeAgentObserver::new();
380        let events = observer.events();
381        self.observer(Box::new(observer)).run().await?;
382        Ok(AgentTrace::from_observer_events(&events))
383    }
384
385    /// Runs the test agent and returns both messages and captured contexts.
386    ///
387    /// Use this when you need to verify what context was passed to the LLM,
388    /// for example when testing that file attachments are properly formatted.
389    pub async fn run_with_context(self) -> TestResult<TestAgentResult> {
390        let Self { provider, agent: config, execution } = self;
391        let mut llm = FakeLlmProvider::from_results(provider.responses).with_context_window(provider.context_window);
392        if let Some(model) = provider.model {
393            llm = llm.with_model(model);
394        }
395        if let Some((turn_index, chunk_index, release)) = provider.pause {
396            llm = llm.pause_turn_after(turn_index, chunk_index, release);
397        }
398        let captured_contexts = llm.captured_contexts();
399
400        let mut mcp_spawn = match config.mcp_server {
401            Some((name, server)) => {
402                Some(mcp("/workspace").with_fake_mcp(name, server).spawn().await.map_err(AgentError::from)?)
403            }
404            None => None,
405        };
406
407        let mut builder = agent(llm);
408        if let Some(spawn) = &mut mcp_spawn {
409            let snapshot = spawn.block_until_ready().await.expect("bootstrap completes");
410            builder = builder.tools(spawn.handle().clone(), snapshot.tool_definitions());
411        }
412        if let Some(timeout) = config.timeout {
413            builder = builder.tool_timeout(timeout);
414        }
415        if let Some(max) = config.max_auto_continues {
416            builder = builder.max_auto_continues(max);
417        }
418        if let Some(retry) = config.retry_config {
419            builder = builder.retry(retry);
420        } else {
421            builder = builder.retry(RetryConfig::disabled());
422        }
423        for prompt in config.system_prompts {
424            builder = builder.system_prompt(prompt);
425        }
426        if let Some(key) = config.session_affinity_key {
427            builder = builder.session_affinity_key(key);
428        }
429        if let Some(compaction) = config.compaction {
430            builder = builder.compaction(compaction);
431        }
432        if let Some(settings) = config.model_settings {
433            builder = builder.model_settings(settings);
434        }
435        builder = builder.context_window(config.context_window_override);
436        if !config.initial_messages.is_empty() {
437            builder = builder.messages(config.initial_messages);
438        }
439        for observer in config.observers {
440            builder = builder.observer(observer);
441        }
442
443        let steps = match execution.expect("test agent requires commands(), user_text(), or scenario()") {
444            TestExecution::CommandsUntilTurnEnd(commands) => {
445                assert!(!commands.is_empty(), "commands() requires at least one command");
446                let mut steps = commands.into_iter().map(TestAgentStep::send).collect::<Vec<_>>();
447                steps.push(TestAgentStep::wait_for_turn_end());
448                steps
449            }
450            TestExecution::Scenario(scenario) => {
451                assert!(!scenario.steps.is_empty(), "scenario() requires at least one step");
452                scenario.steps
453            }
454        };
455        let (tx, mut rx, handle) = builder.spawn().await?;
456        let mut messages = Vec::new();
457
458        for step in steps {
459            match step {
460                TestAgentStep::Send(command) => tx.send(command).await?,
461                TestAgentStep::WaitFor(predicate) => loop {
462                    let message = rx.recv().await.expect("agent event channel closed before scenario step matched");
463                    let matched = predicate(&message);
464                    messages.push(message);
465                    if matched {
466                        break;
467                    }
468                },
469                TestAgentStep::Perform(action) => action(),
470            }
471        }
472        drop(tx);
473
474        handle.await_completion().await;
475
476        Ok(TestAgentResult { messages, captured_contexts })
477    }
478
479    fn with_execution(mut self, execution: TestExecution) -> Self {
480        assert!(self.execution.is_none(), "commands(), user_text(), and scenario() are mutually exclusive");
481        self.execution = Some(execution);
482        self
483    }
484}