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
54pub 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
63pub enum TestAgentStep {
65 Send(Command),
66 WaitFor(Box<dyn Fn(&AgentEvent) -> bool + Send>),
67 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#[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 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
185pub struct TestAgentResult {
187 pub messages: Vec<AgentEvent>,
188 pub captured_contexts: Arc<Mutex<Vec<Context>>>,
189}
190
191#[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 pub fn without_mcp(mut self) -> Self {
316 self.agent.mcp_server = None;
317 self
318 }
319
320 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 pub fn messages(mut self, messages: Vec<ChatMessage>) -> Self {
328 self.agent.initial_messages = messages;
329 self
330 }
331
332 pub fn system_prompt(mut self, prompt: Prompt) -> Self {
335 self.agent.system_prompts.push(prompt);
336 self
337 }
338
339 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 pub fn compaction_config(mut self, config: CompactionConfig) -> Self {
348 self.agent.compaction = Some(config);
349 self
350 }
351
352 pub fn model_settings(mut self, settings: ModelSettings) -> Self {
354 self.agent.model_settings = Some(settings);
355 self
356 }
357
358 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 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 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 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}