1use crate::agent::{Agent, Conversation, RunContext, StopCause, Taint, ToolCallTrace};
8use crate::message::{Message, StopReason, Usage};
9use futures::stream::StreamExt;
10use serde::{Deserialize, Serialize};
11use std::sync::Arc;
12
13#[derive(Debug, Clone, Serialize, Deserialize)]
22#[serde(untagged)]
23pub enum Prompt {
24 One(String),
25 Many(Vec<String>),
26}
27
28impl Prompt {
29 pub fn turns(&self) -> &[String] {
30 match self {
31 Prompt::One(s) => std::slice::from_ref(s),
32 Prompt::Many(v) => v,
33 }
34 }
35
36 pub fn first(&self) -> &str {
38 self.turns().first().map(String::as_str).unwrap_or_default()
39 }
40
41 pub fn render(&self) -> String {
46 match self {
47 Prompt::One(s) => s.clone(),
48 Prompt::Many(v) => v
49 .iter()
50 .enumerate()
51 .map(|(i, t)| format!("[turn {}] {t}", i + 1))
52 .collect::<Vec<_>>()
53 .join("\n"),
54 }
55 }
56}
57
58impl From<String> for Prompt {
59 fn from(s: String) -> Self {
60 Prompt::One(s)
61 }
62}
63
64impl From<&str> for Prompt {
65 fn from(s: &str) -> Self {
66 Prompt::One(s.to_string())
67 }
68}
69
70#[derive(Debug, Clone, Serialize, Deserialize)]
71pub struct BatchItem {
72 pub id: String,
74 pub prompt: Prompt,
75 #[serde(default, skip_serializing_if = "Option::is_none")]
78 pub meta: Option<serde_json::Value>,
79}
80
81#[derive(Debug, Clone, Serialize, Deserialize)]
82pub struct BatchResult {
83 pub id: String,
84 pub ok: bool,
85 pub text: String,
86 #[serde(default, skip_serializing_if = "Option::is_none")]
87 pub error: Option<String>,
88 pub turns: u32,
89 pub usage: Usage,
90 pub stop_reason: Option<StopReason>,
91 #[serde(default, skip_serializing_if = "Option::is_none")]
92 pub meta: Option<serde_json::Value>,
93 pub elapsed_ms: u64,
94 #[serde(default)]
96 pub tool_calls: Vec<ToolCallTrace>,
97 #[serde(default)]
98 pub malformed_tool_args: u32,
99 #[serde(default, skip_serializing_if = "Option::is_none")]
101 pub stop_cause: Option<StopCause>,
102 #[serde(default)]
104 pub taint: Taint,
105 #[serde(default)]
107 pub blocked_sends: u32,
108 #[serde(default)]
110 pub compactions: u32,
111 #[serde(default)]
113 pub usage_complete: bool,
114}
115
116pub async fn run<F>(
121 agent: &Agent,
122 items: Vec<BatchItem>,
123 concurrency: usize,
124 on_result: F,
125) -> Vec<BatchResult>
126where
127 F: FnMut(&BatchResult),
128{
129 run_with(agent, items, concurrency, |_| None, on_result).await
130}
131
132pub async fn run_with<C, F>(
139 agent: &Agent,
140 items: Vec<BatchItem>,
141 concurrency: usize,
142 context_for: C,
143 mut on_result: F,
144) -> Vec<BatchResult>
145where
146 C: Fn(&BatchItem) -> Option<Arc<RunContext>> + Sync,
147 F: FnMut(&BatchResult),
148{
149 let concurrency = concurrency.max(1);
150 let context_for = &context_for;
151
152 let mut stream = futures::stream::iter(items.into_iter().map(|item| async move {
153 let started = std::time::Instant::now();
154 let mut convo = Conversation::new();
160 let cx = context_for(&item).unwrap_or_else(|| Arc::clone(agent.context()));
161
162 let mut totals = Totals::default();
166 let mut failure = None;
167 let mut last: Option<crate::agent::RunOutcome> = None;
168
169 for turn in item.prompt.turns() {
170 convo.push(Message::user(turn.clone()));
171 match agent.run_in(&cx, &mut convo, None).await {
172 Ok(outcome) => {
173 totals.absorb(&outcome);
174 last = Some(outcome);
175 }
176 Err(e) => {
177 failure = Some(format!("{e:#}"));
181 break;
182 }
183 }
184 }
185
186 let elapsed_ms = started.elapsed().as_millis() as u64;
187
188 match (last, failure) {
189 (Some(outcome), None) => BatchResult {
190 id: item.id,
191 ok: !outcome.exhausted
194 && outcome.stop_reason != StopReason::Refusal
195 && totals.malformed_tool_args == 0,
196 text: outcome.text,
197 error: outcome.refusal.map(|r| {
198 format!(
199 "refused ({}): {}",
200 r.category.unwrap_or_else(|| "unspecified".into()),
201 r.explanation.unwrap_or_default()
202 )
203 }),
204 turns: totals.turns,
205 usage: totals.usage,
206 stop_reason: Some(outcome.stop_reason),
207 meta: item.meta,
208 elapsed_ms,
209 tool_calls: totals.tool_calls,
210 malformed_tool_args: totals.malformed_tool_args,
211 stop_cause: Some(outcome.stop_cause),
212 taint: convo.taint,
214 blocked_sends: totals.blocked_sends,
215 compactions: totals.compactions,
216 usage_complete: totals.usage_complete,
217 },
218 (_, error) => BatchResult {
222 id: item.id,
223 ok: false,
224 text: String::new(),
225 error: error.or_else(|| Some("the item had no prompts".into())),
226 turns: totals.turns,
227 usage: totals.usage,
228 stop_reason: None,
229 meta: item.meta,
230 elapsed_ms,
231 tool_calls: totals.tool_calls,
232 malformed_tool_args: totals.malformed_tool_args,
233 stop_cause: None,
234 taint: convo.taint,
235 blocked_sends: totals.blocked_sends,
236 compactions: totals.compactions,
237 usage_complete: totals.usage_complete,
238 },
239 }
240 }))
241 .buffer_unordered(concurrency);
242
243 let mut results = Vec::new();
244 while let Some(result) = stream.next().await {
245 on_result(&result);
246 results.push(result);
247 }
248 results
249}
250
251struct Totals {
257 usage: Usage,
258 turns: u32,
259 tool_calls: Vec<ToolCallTrace>,
260 malformed_tool_args: u32,
261 blocked_sends: u32,
262 compactions: u32,
263 usage_complete: bool,
264}
265
266impl Default for Totals {
267 fn default() -> Self {
268 Totals {
269 usage: Usage::default(),
270 turns: 0,
271 tool_calls: Vec::new(),
272 malformed_tool_args: 0,
273 blocked_sends: 0,
274 compactions: 0,
275 usage_complete: true,
278 }
279 }
280}
281
282impl Totals {
283 fn absorb(&mut self, outcome: &crate::agent::RunOutcome) {
284 self.usage.add(&outcome.usage);
285 self.turns += outcome.turns;
286 self.tool_calls.extend(outcome.tool_calls.iter().cloned());
287 self.malformed_tool_args += outcome.malformed_tool_args;
288 self.blocked_sends += outcome.blocked_sends;
289 self.compactions += outcome.compactions;
290 self.usage_complete &= outcome.usage_complete;
291 }
292}
293
294#[derive(Debug, Clone, Serialize)]
296pub struct BatchSummary {
297 pub total: usize,
298 pub succeeded: usize,
299 pub failed: usize,
300 pub usage: Usage,
301 pub elapsed_ms: u64,
302}
303
304impl BatchSummary {
305 pub fn of(results: &[BatchResult], elapsed_ms: u64) -> Self {
306 let mut usage = Usage::default();
307 for r in results {
308 usage.add(&r.usage);
309 }
310 let succeeded = results.iter().filter(|r| r.ok).count();
311 BatchSummary {
312 total: results.len(),
313 succeeded,
314 failed: results.len() - succeeded,
315 usage,
316 elapsed_ms,
317 }
318 }
319}
320
321#[cfg(test)]
322mod tests {
323 use super::*;
324 use crate::config::{AgentConfig, PermissionMode};
325 use crate::message::{Block, CompletionRequest, CompletionResponse, Message};
326 use crate::provider::{Provider, StreamSink};
327 use crate::tool::{ModeApprover, Registry, ToolCtx};
328 use anyhow::Result;
329 use async_trait::async_trait;
330 use std::sync::atomic::{AtomicUsize, Ordering};
331 use std::sync::Mutex;
332
333 #[derive(Default)]
340 struct EchoProvider {
341 history_lengths: Mutex<Vec<usize>>,
344 in_flight: AtomicUsize,
345 max_in_flight: AtomicUsize,
346 delay: std::time::Duration,
347 }
348
349 #[async_trait]
350 impl Provider for EchoProvider {
351 fn id(&self) -> &str {
352 "echo"
353 }
354 fn default_model(&self) -> &str {
355 "echo-1"
356 }
357
358 async fn complete(
359 &self,
360 req: &CompletionRequest,
361 _sink: Option<&StreamSink>,
362 ) -> Result<CompletionResponse> {
363 let now = self.in_flight.fetch_add(1, Ordering::SeqCst) + 1;
364 self.max_in_flight.fetch_max(now, Ordering::SeqCst);
365 if !self.delay.is_zero() {
366 tokio::time::sleep(self.delay).await;
367 }
368 self.in_flight.fetch_sub(1, Ordering::SeqCst);
369
370 self.history_lengths
371 .lock()
372 .unwrap()
373 .push(req.messages.len());
374 let prompt = req.messages.last().map(|m| m.text()).unwrap_or_default();
375
376 anyhow::ensure!(!prompt.contains("boom"), "the provider exploded");
379
380 Ok(CompletionResponse {
381 message: Message::assistant(vec![Block::text(format!("answered: {prompt}"))]),
382 stop_reason: StopReason::EndTurn,
383 usage: Usage {
384 input_tokens: 10,
385 output_tokens: 5,
386 ..Usage::default()
387 },
388 refusal: None,
389 model: "echo-1".into(),
390 malformed_tool_args: 0,
391 })
392 }
393 }
394
395 fn agent_with(provider: Arc<EchoProvider>) -> Agent {
396 struct Shared(Arc<EchoProvider>);
397 #[async_trait]
398 impl Provider for Shared {
399 fn id(&self) -> &str {
400 self.0.id()
401 }
402 fn default_model(&self) -> &str {
403 self.0.default_model()
404 }
405 async fn complete(
406 &self,
407 req: &CompletionRequest,
408 sink: Option<&StreamSink>,
409 ) -> Result<CompletionResponse> {
410 self.0.complete(req, sink).await
411 }
412 }
413
414 Agent::new(
415 Box::new(Shared(provider)),
416 Registry::new(),
417 Arc::new(ModeApprover {
418 mode: PermissionMode::Allow,
419 }),
420 ToolCtx {
421 workspace: std::env::temp_dir(),
422 ..Default::default()
423 },
424 AgentConfig::default(),
425 None,
426 )
427 .unwrap()
428 }
429
430 fn items(prompts: &[&str]) -> Vec<BatchItem> {
431 prompts
432 .iter()
433 .enumerate()
434 .map(|(i, p)| BatchItem {
435 id: format!("item-{i}"),
436 prompt: (*p).to_string().into(),
437 meta: Some(serde_json::json!({"index": i})),
438 })
439 .collect()
440 }
441
442 #[tokio::test]
443 async fn results_are_matched_by_id_not_by_position() {
444 let provider = Arc::new(EchoProvider {
445 delay: std::time::Duration::from_millis(20),
446 ..Default::default()
447 });
448 let agent = agent_with(Arc::clone(&provider));
449
450 let results = run(
451 &agent,
452 items(&["alpha", "beta", "gamma", "delta"]),
453 4,
454 |_| {},
455 )
456 .await;
457
458 assert_eq!(results.len(), 4);
461 for r in &results {
462 let index = r.meta.as_ref().unwrap()["index"].as_u64().unwrap();
463 assert_eq!(r.id, format!("item-{index}"));
464 let expected = ["alpha", "beta", "gamma", "delta"][index as usize];
465 assert_eq!(
466 r.text,
467 format!("answered: {expected}"),
468 "{} got another item's answer",
469 r.id
470 );
471 }
472 }
473
474 #[tokio::test]
475 async fn every_item_gets_its_own_conversation() {
476 let provider = Arc::new(EchoProvider::default());
477 let agent = agent_with(Arc::clone(&provider));
478
479 run(&agent, items(&["one", "two", "three"]), 1, |_| {}).await;
480
481 let lengths = provider.history_lengths.lock().unwrap();
486 assert_eq!(*lengths, vec![1, 1, 1], "history leaked between items");
487 }
488
489 #[tokio::test]
490 async fn a_failing_item_is_recorded_rather_than_sinking_the_batch() {
491 let provider = Arc::new(EchoProvider::default());
492 let agent = agent_with(Arc::clone(&provider));
493
494 let results = run(&agent, items(&["fine", "boom", "also fine"]), 1, |_| {}).await;
495
496 assert_eq!(results.len(), 3, "a failure took other items down with it");
497 let failed: Vec<_> = results.iter().filter(|r| !r.ok).collect();
498 assert_eq!(failed.len(), 1);
499 assert_eq!(failed[0].id, "item-1");
500 assert!(failed[0].error.as_ref().unwrap().contains("exploded"));
501 assert!(failed[0].meta.is_some());
504
505 assert!(results.iter().filter(|r| r.ok).count() == 2);
506 }
507
508 #[tokio::test]
509 async fn concurrency_is_bounded_by_what_was_asked_for() {
510 let provider = Arc::new(EchoProvider {
511 delay: std::time::Duration::from_millis(50),
512 ..Default::default()
513 });
514 let agent = agent_with(Arc::clone(&provider));
515
516 run(&agent, items(&["a", "b", "c", "d", "e", "f"]), 2, |_| {}).await;
517
518 let peak = provider.max_in_flight.load(Ordering::SeqCst);
519 assert!(peak <= 2, "{peak} items ran at once against a limit of 2");
520 assert_eq!(
521 peak, 2,
522 "the limit was never actually reached; the test proves nothing"
523 );
524 }
525
526 #[tokio::test]
527 async fn a_concurrency_of_zero_still_makes_progress() {
528 let provider = Arc::new(EchoProvider::default());
529 let agent = agent_with(Arc::clone(&provider));
530
531 let results = run(&agent, items(&["only"]), 0, |_| {}).await;
534 assert_eq!(results.len(), 1);
535 assert!(results[0].ok);
536 }
537
538 #[tokio::test]
539 async fn each_result_is_announced_as_it_lands() {
540 let provider = Arc::new(EchoProvider::default());
541 let agent = agent_with(Arc::clone(&provider));
542
543 let mut announced = Vec::new();
546 let results = run(&agent, items(&["a", "b", "c"]), 1, |r| {
547 announced.push(r.id.clone())
548 })
549 .await;
550
551 assert_eq!(announced.len(), 3);
552 assert_eq!(
553 announced,
554 results.iter().map(|r| r.id.clone()).collect::<Vec<_>>()
555 );
556 }
557
558 #[test]
559 fn a_summary_totals_usage_and_counts_both_outcomes() {
560 let result = |id: &str, ok: bool| BatchResult {
561 id: id.into(),
562 ok,
563 text: String::new(),
564 error: None,
565 turns: 1,
566 usage: Usage {
567 input_tokens: 10,
568 output_tokens: 5,
569 ..Usage::default()
570 },
571 stop_reason: Some(StopReason::EndTurn),
572 meta: None,
573 elapsed_ms: 1,
574 tool_calls: Vec::new(),
575 malformed_tool_args: 0,
576 stop_cause: None,
577 taint: Taint::default(),
578 blocked_sends: 0,
579 compactions: 0,
580 usage_complete: true,
581 };
582
583 let summary = BatchSummary::of(
584 &[result("a", true), result("b", false), result("c", true)],
585 99,
586 );
587
588 assert_eq!(summary.total, 3);
589 assert_eq!(summary.succeeded, 2);
590 assert_eq!(summary.failed, 1);
591 assert_eq!(summary.usage.input_tokens, 30);
594 assert_eq!(summary.usage.output_tokens, 15);
595 assert_eq!(summary.elapsed_ms, 99);
596 }
597}