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 { .. } | TurnEvent::LlmCallStarted { .. } | TurnEvent::LlmCallEnded { .. }
40 ) | AgentEvent::Tool(ToolEvent::ExecutionStarted { .. } | ToolEvent::DefinitionsUpdated { .. })
41 )
42 })
43 .collect()
44}
45
46pub fn mcp_instructions(entries: &[(&str, &str)]) -> BTreeMap<String, String> {
47 entries.iter().map(|(k, v)| ((*k).to_string(), (*v).to_string())).collect()
48}
49
50pub fn fast_retry(max_attempts: u32) -> RetryConfig {
52 RetryConfig { max_attempts, base_delay: Duration::from_millis(1), max_delay: Duration::from_millis(5) }
53}
54
55pub fn test_agent() -> TestAgentBuilder {
56 TestAgentBuilder::new()
57}
58
59pub enum TestAgentStep {
61 Send(Command),
62 WaitFor(Box<dyn Fn(&AgentEvent) -> bool + Send>),
63 Perform(Box<dyn FnOnce() + Send>),
65}
66
67impl TestAgentStep {
68 pub fn send(command: Command) -> Self {
69 Self::Send(command)
70 }
71
72 pub fn user_text(text: impl Into<String>) -> Self {
73 Self::send(Command::text(text))
74 }
75
76 pub fn cancel() -> Self {
77 Self::send(Command::UserCommand(UserCommand::Cancel))
78 }
79
80 pub fn switch_model(provider: impl StreamingModelProvider + 'static) -> Self {
81 Self::send(Command::AgentCommand(AgentCommand::SwitchModel(Box::new(provider))))
82 }
83
84 pub fn replace_conversation(messages: Vec<ChatMessage>) -> Self {
85 Self::send(Command::AgentCommand(AgentCommand::ReplaceConversation(messages)))
86 }
87
88 pub fn perform(action: impl FnOnce() + Send + 'static) -> Self {
89 Self::Perform(Box::new(action))
90 }
91
92 pub fn wait_for(predicate: impl Fn(&AgentEvent) -> bool + Send + 'static) -> Self {
93 Self::WaitFor(Box::new(predicate))
94 }
95
96 pub fn wait_for_turn_end() -> Self {
97 Self::wait_for(|event| matches!(event, AgentEvent::Turn(TurnEvent::Ended { .. })))
98 }
99
100 pub fn wait_for_compaction_start() -> Self {
101 Self::wait_for(|event| matches!(event, AgentEvent::Context(ContextEvent::CompactionStarted { .. })))
102 }
103
104 pub fn wait_for_retry(attempt: u32) -> Self {
105 Self::wait_for(
106 move |event| matches!(event, AgentEvent::Turn(TurnEvent::RetryScheduled { attempt: actual, .. }) if *actual == attempt),
107 )
108 }
109}
110
111#[derive(Default)]
113pub struct TestScenario {
114 steps: Vec<TestAgentStep>,
115}
116
117impl TestScenario {
118 pub fn new() -> Self {
119 Self::default()
120 }
121
122 pub fn send(mut self, command: Command) -> Self {
123 self.steps.push(TestAgentStep::send(command));
124 self
125 }
126
127 pub fn user_text(mut self, text: impl Into<String>) -> Self {
128 self.steps.push(TestAgentStep::user_text(text));
129 self
130 }
131
132 pub fn cancel(mut self) -> Self {
133 self.steps.push(TestAgentStep::cancel());
134 self
135 }
136
137 pub fn switch_model(mut self, provider: impl StreamingModelProvider + 'static) -> Self {
138 self.steps.push(TestAgentStep::switch_model(provider));
139 self
140 }
141
142 pub fn replace_conversation(mut self, messages: Vec<ChatMessage>) -> Self {
143 self.steps.push(TestAgentStep::replace_conversation(messages));
144 self
145 }
146
147 pub fn wait_for(mut self, predicate: impl Fn(&AgentEvent) -> bool + Send + 'static) -> Self {
148 self.steps.push(TestAgentStep::wait_for(predicate));
149 self
150 }
151
152 pub fn wait_for_turn_end(mut self) -> Self {
153 self.steps.push(TestAgentStep::wait_for_turn_end());
154 self
155 }
156
157 pub fn wait_for_compaction_start(mut self) -> Self {
158 self.steps.push(TestAgentStep::wait_for_compaction_start());
159 self
160 }
161
162 pub fn wait_for_retry(mut self, attempt: u32) -> Self {
163 self.steps.push(TestAgentStep::wait_for_retry(attempt));
164 self
165 }
166
167 pub fn perform(mut self, action: impl FnOnce() + Send + 'static) -> Self {
170 self.steps.push(TestAgentStep::perform(action));
171 self
172 }
173}
174
175impl From<Vec<TestAgentStep>> for TestScenario {
176 fn from(steps: Vec<TestAgentStep>) -> Self {
177 Self { steps }
178 }
179}
180
181pub struct TestAgentResult {
183 pub messages: Vec<AgentEvent>,
184 pub captured_contexts: Arc<Mutex<Vec<Context>>>,
185}
186
187#[derive(Debug, thiserror::Error)]
190pub enum TestAgentError {
191 #[error(transparent)]
192 Agent(#[from] AgentError),
193 #[error("failed to send command to the test agent: {0}")]
194 SendCommand(#[from] mpsc::error::SendError<Command>),
195}
196
197pub type TestResult<T> = std::result::Result<T, TestAgentError>;
198
199struct ProviderTestConfig {
200 responses: Vec<Vec<Result<LlmResponse, LlmError>>>,
201 model: Option<LlmModel>,
202 context_window: Option<u32>,
203 pause: Option<(usize, usize, Arc<Notify>)>,
204}
205
206struct AgentTestConfig {
207 context_window_override: Option<u32>,
208 timeout: Option<Duration>,
209 max_auto_continues: Option<u32>,
210 retry_config: Option<RetryConfig>,
211 observers: Vec<Box<dyn AgentObserver>>,
212 mcp_server: Option<(String, FakeMcpServer)>,
213 initial_messages: Vec<ChatMessage>,
214 system_prompts: Vec<Prompt>,
215 session_affinity_key: Option<String>,
216 compaction: Option<CompactionConfig>,
217 model_settings: Option<ModelSettings>,
218}
219
220enum TestExecution {
221 CommandsUntilTurnEnd(Vec<Command>),
222 Scenario(TestScenario),
223}
224
225pub struct TestAgentBuilder {
226 provider: ProviderTestConfig,
227 agent: AgentTestConfig,
228 execution: Option<TestExecution>,
229}
230
231impl Default for TestAgentBuilder {
232 fn default() -> Self {
233 Self::new()
234 }
235}
236
237impl TestAgentBuilder {
238 pub fn new() -> Self {
239 Self {
240 provider: ProviderTestConfig { responses: Vec::new(), model: None, context_window: None, pause: None },
241 agent: AgentTestConfig {
242 context_window_override: None,
243 timeout: None,
244 max_auto_continues: None,
245 retry_config: None,
246 observers: Vec::new(),
247 mcp_server: Some(("test".to_string(), FakeMcpServer::new())),
248 initial_messages: Vec::new(),
249 system_prompts: Vec::new(),
250 session_affinity_key: None,
251 compaction: None,
252 model_settings: None,
253 },
254 execution: None,
255 }
256 }
257
258 pub fn commands(self, commands: Vec<Command>) -> Self {
259 self.with_execution(TestExecution::CommandsUntilTurnEnd(commands))
260 }
261
262 pub fn scenario(self, scenario: impl Into<TestScenario>) -> Self {
263 self.with_execution(TestExecution::Scenario(scenario.into()))
264 }
265
266 pub fn user_text(self, text: &str) -> Self {
267 self.commands(vec![Command::text(text)])
268 }
269
270 pub fn llm_responses(mut self, llm_responses: &[Vec<LlmResponse>]) -> Self {
271 self.provider.responses = llm_responses.iter().map(|turn| turn.iter().cloned().map(Ok).collect()).collect();
272 self
273 }
274
275 pub fn llm_result_responses(mut self, llm_responses: &[Vec<Result<LlmResponse, LlmError>>]) -> Self {
276 self.provider.responses = Vec::from(llm_responses);
277 self
278 }
279
280 pub fn model(mut self, model: LlmModel) -> Self {
281 self.provider.model = Some(model);
282 self
283 }
284
285 pub fn provider_context_window(mut self, window: Option<u32>) -> Self {
286 self.provider.context_window = window;
287 self
288 }
289
290 pub fn context_window_override(mut self, window: u32) -> Self {
291 self.agent.context_window_override = Some(window);
292 self
293 }
294
295 pub fn tool_timeout(mut self, timeout: Duration) -> Self {
296 self.agent.timeout = Some(timeout);
297 self
298 }
299
300 pub fn max_auto_continues(mut self, max: u32) -> Self {
301 self.agent.max_auto_continues = Some(max);
302 self
303 }
304
305 pub fn retry_config(mut self, config: RetryConfig) -> Self {
306 self.agent.retry_config = Some(config);
307 self
308 }
309
310 pub fn without_mcp(mut self) -> Self {
312 self.agent.mcp_server = None;
313 self
314 }
315
316 pub fn fake_mcp_server(mut self, name: &str, server: FakeMcpServer) -> Self {
318 self.agent.mcp_server = Some((name.to_string(), server));
319 self
320 }
321
322 pub fn messages(mut self, messages: Vec<ChatMessage>) -> Self {
324 self.agent.initial_messages = messages;
325 self
326 }
327
328 pub fn system_prompt(mut self, prompt: Prompt) -> Self {
331 self.agent.system_prompts.push(prompt);
332 self
333 }
334
335 pub fn session_affinity_key(mut self, key: impl Into<String>) -> Self {
338 self.agent.session_affinity_key = Some(key.into());
339 self
340 }
341
342 pub fn compaction_config(mut self, config: CompactionConfig) -> Self {
344 self.agent.compaction = Some(config);
345 self
346 }
347
348 pub fn model_settings(mut self, settings: ModelSettings) -> Self {
350 self.agent.model_settings = Some(settings);
351 self
352 }
353
354 pub fn pause_turn_after(mut self, turn_index: usize, chunk_index: usize, release: Arc<Notify>) -> Self {
357 self.provider.pause = Some((turn_index, chunk_index, release));
358 self
359 }
360
361 pub fn observer(mut self, observer: Box<dyn AgentObserver>) -> Self {
363 self.agent.observers.push(observer);
364 self
365 }
366
367 pub async fn run(self) -> TestResult<Vec<AgentEvent>> {
368 let result = self.run_with_context().await?;
369 Ok(result.messages)
370 }
371
372 pub async fn run_trace(self) -> TestResult<AgentTrace> {
375 let observer = FakeAgentObserver::new();
376 let events = observer.events();
377 self.observer(Box::new(observer)).run().await?;
378 Ok(AgentTrace::from_observer_events(&events))
379 }
380
381 pub async fn run_with_context(self) -> TestResult<TestAgentResult> {
386 let Self { provider, agent: config, execution } = self;
387 let mut llm = FakeLlmProvider::from_results(provider.responses).with_context_window(provider.context_window);
388 if let Some(model) = provider.model {
389 llm = llm.with_model(model);
390 }
391 if let Some((turn_index, chunk_index, release)) = provider.pause {
392 llm = llm.pause_turn_after(turn_index, chunk_index, release);
393 }
394 let captured_contexts = llm.captured_contexts();
395
396 let mut mcp_spawn = match config.mcp_server {
397 Some((name, server)) => {
398 Some(mcp("/workspace").with_fake_mcp(name, server).spawn().await.map_err(AgentError::from)?)
399 }
400 None => None,
401 };
402
403 let mut builder = agent(llm);
404 if let Some(spawn) = &mut mcp_spawn {
405 let snapshot = spawn.block_until_ready().await.expect("bootstrap completes");
406 builder = builder.tools(spawn.handle().clone(), snapshot.tool_definitions());
407 }
408 if let Some(timeout) = config.timeout {
409 builder = builder.tool_timeout(timeout);
410 }
411 if let Some(max) = config.max_auto_continues {
412 builder = builder.max_auto_continues(max);
413 }
414 if let Some(retry) = config.retry_config {
415 builder = builder.retry(retry);
416 } else {
417 builder = builder.retry(RetryConfig::disabled());
418 }
419 for prompt in config.system_prompts {
420 builder = builder.system_prompt(prompt);
421 }
422 if let Some(key) = config.session_affinity_key {
423 builder = builder.session_affinity_key(key);
424 }
425 if let Some(compaction) = config.compaction {
426 builder = builder.compaction(compaction);
427 }
428 if let Some(settings) = config.model_settings {
429 builder = builder.model_settings(settings);
430 }
431 builder = builder.context_window(config.context_window_override);
432 if !config.initial_messages.is_empty() {
433 builder = builder.messages(config.initial_messages);
434 }
435 for observer in config.observers {
436 builder = builder.observer(observer);
437 }
438
439 let steps = match execution.expect("test agent requires commands(), user_text(), or scenario()") {
440 TestExecution::CommandsUntilTurnEnd(commands) => {
441 assert!(!commands.is_empty(), "commands() requires at least one command");
442 let mut steps = commands.into_iter().map(TestAgentStep::send).collect::<Vec<_>>();
443 steps.push(TestAgentStep::wait_for_turn_end());
444 steps
445 }
446 TestExecution::Scenario(scenario) => {
447 assert!(!scenario.steps.is_empty(), "scenario() requires at least one step");
448 scenario.steps
449 }
450 };
451 let (tx, mut rx, handle) = builder.spawn().await?;
452 let mut messages = Vec::new();
453
454 for step in steps {
455 match step {
456 TestAgentStep::Send(command) => tx.send(command).await?,
457 TestAgentStep::WaitFor(predicate) => loop {
458 let message = rx.recv().await.expect("agent event channel closed before scenario step matched");
459 let matched = predicate(&message);
460 messages.push(message);
461 if matched {
462 break;
463 }
464 },
465 TestAgentStep::Perform(action) => action(),
466 }
467 }
468 drop(tx);
469
470 handle.await_completion().await;
471
472 Ok(TestAgentResult { messages, captured_contexts })
473 }
474
475 fn with_execution(mut self, execution: TestExecution) -> Self {
476 assert!(self.execution.is_none(), "commands(), user_text(), and scenario() are mutually exclusive");
477 self.execution = Some(execution);
478 self
479 }
480}