1use std::collections::HashMap;
9use std::sync::Arc;
10
11use tokio::task::JoinSet;
12
13use crate::error::Error;
14use crate::llm::LlmProvider;
15use crate::llm::types::TokenUsage;
16
17use super::{AgentOutput, AgentRunner};
18
19type VoteExtractor = Box<dyn Fn(&str) -> String + Send + Sync>;
21
22type TieBreaker = Box<dyn Fn(&[String]) -> String + Send + Sync>;
25
26#[derive(Debug)]
29pub struct VoteResult {
30 pub winner: String,
32 pub tally: HashMap<String, usize>,
34 pub output: AgentOutput,
36}
37
38pub struct VotingAgent<P: LlmProvider + 'static> {
40 voters: Vec<Arc<AgentRunner<P>>>,
41 vote_extractor: VoteExtractor,
42 tie_breaker: TieBreaker,
43}
44
45impl<P: LlmProvider + 'static> std::fmt::Debug for VotingAgent<P> {
46 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
47 f.debug_struct("VotingAgent")
48 .field("voter_count", &self.voters.len())
49 .finish()
50 }
51}
52
53pub struct VotingAgentBuilder<P: LlmProvider + 'static> {
55 voters: Vec<Arc<AgentRunner<P>>>,
56 vote_extractor: Option<VoteExtractor>,
57 tie_breaker: Option<TieBreaker>,
58}
59
60impl<P: LlmProvider + 'static> VotingAgent<P> {
61 pub fn builder() -> VotingAgentBuilder<P> {
63 VotingAgentBuilder {
64 voters: Vec::new(),
65 vote_extractor: None,
66 tie_breaker: None,
67 }
68 }
69
70 pub async fn execute(&self, task: &str) -> Result<VoteResult, Error> {
72 let mut set = JoinSet::new();
73
74 for (idx, voter) in self.voters.iter().enumerate() {
75 let voter = Arc::clone(voter);
76 let task = task.to_string();
77 set.spawn(async move {
78 let result = voter.execute(&task).await;
79 (idx, result)
80 });
81 }
82
83 let mut outputs: Vec<(usize, AgentOutput)> = Vec::with_capacity(self.voters.len());
86 let mut total_usage = TokenUsage::default();
87
88 while let Some(join_result) = set.join_next().await {
89 let (idx, agent_result) = join_result.map_err(|e| {
90 Error::Agent(format!("voting agent task panicked: {e}"))
92 .accumulate_usage(total_usage)
93 })?;
94 let output = agent_result.map_err(|e| e.accumulate_usage(total_usage))?;
95 total_usage += output.tokens_used;
96 outputs.push((idx, output));
97 }
98
99 outputs.sort_by_key(|(idx, _)| *idx);
101
102 let votes: Vec<String> = outputs
104 .iter()
105 .map(|(_, output)| (self.vote_extractor)(&output.result))
106 .collect();
107
108 let mut tally: HashMap<String, usize> = HashMap::new();
109 for vote in &votes {
110 *tally.entry(vote.clone()).or_insert(0) += 1;
111 }
112
113 let max_count = tally.values().copied().max().unwrap_or(0);
115
116 let mut top_votes: Vec<String> = tally
118 .iter()
119 .filter(|&(_, &count)| count == max_count)
120 .map(|(vote, _)| vote.clone())
121 .collect();
122 top_votes.sort();
123
124 let winner = if top_votes.len() == 1 {
125 top_votes.into_iter().next().expect("at least one vote")
126 } else {
127 (self.tie_breaker)(&top_votes)
128 };
129
130 let winner_idx = votes
132 .iter()
133 .position(|v| *v == winner)
134 .expect("winner must be among votes");
135
136 let (_, mut winning_output) = outputs.remove(winner_idx);
137
138 let mut total_tool_calls = 0usize;
141 let mut total_cost: Option<f64> = None;
142 for (_, output) in &outputs {
143 total_tool_calls += output.tool_calls_made;
144 if let Some(cost) = output.estimated_cost_usd {
145 *total_cost.get_or_insert(0.0) += cost;
146 }
147 }
148 total_tool_calls += winning_output.tool_calls_made;
149 if let Some(cost) = winning_output.estimated_cost_usd {
150 *total_cost.get_or_insert(0.0) += cost;
151 }
152
153 winning_output.tokens_used = total_usage;
154 winning_output.tool_calls_made = total_tool_calls;
155 winning_output.estimated_cost_usd = total_cost;
156
157 Ok(VoteResult {
158 winner,
159 tally,
160 output: winning_output,
161 })
162 }
163}
164
165impl<P: LlmProvider + 'static> VotingAgentBuilder<P> {
166 pub fn voter(mut self, agent: AgentRunner<P>) -> Self {
168 self.voters.push(Arc::new(agent));
169 self
170 }
171
172 pub fn voters(mut self, agents: Vec<AgentRunner<P>>) -> Self {
174 self.voters.extend(agents.into_iter().map(Arc::new));
175 self
176 }
177
178 pub fn vote_extractor(mut self, f: impl Fn(&str) -> String + Send + Sync + 'static) -> Self {
180 self.vote_extractor = Some(Box::new(f));
181 self
182 }
183
184 pub fn tie_breaker(mut self, f: impl Fn(&[String]) -> String + Send + Sync + 'static) -> Self {
186 self.tie_breaker = Some(Box::new(f));
187 self
188 }
189
190 pub fn build(self) -> Result<VotingAgent<P>, Error> {
192 if self.voters.len() < 2 {
193 return Err(Error::Config(
194 "VotingAgent requires at least 2 voters".into(),
195 ));
196 }
197 let vote_extractor = self
198 .vote_extractor
199 .ok_or_else(|| Error::Config("VotingAgent requires a vote_extractor".into()))?;
200 let tie_breaker = self.tie_breaker.unwrap_or_else(|| {
201 Box::new(|votes: &[String]| {
202 votes[0].clone()
204 })
205 });
206 Ok(VotingAgent {
207 voters: self.voters,
208 vote_extractor,
209 tie_breaker,
210 })
211 }
212}
213
214#[cfg(test)]
219mod tests {
220 use super::*;
221 use crate::agent::test_helpers::{MockProvider, make_agent};
222
223 fn yes_no_extractor(output: &str) -> String {
224 if output.contains("YES") {
225 "YES".to_string()
226 } else {
227 "NO".to_string()
228 }
229 }
230
231 #[test]
236 fn builder_rejects_fewer_than_two_voters() {
237 let provider = Arc::new(MockProvider::new(vec![MockProvider::text_response(
238 "YES", 10, 5,
239 )]));
240 let result = VotingAgent::builder()
241 .voter(make_agent(provider, "only-one"))
242 .vote_extractor(yes_no_extractor)
243 .build();
244 assert!(result.is_err());
245 assert!(result.unwrap_err().to_string().contains("at least 2"));
246 }
247
248 #[test]
249 fn builder_rejects_zero_voters() {
250 let result = VotingAgent::<MockProvider>::builder()
251 .vote_extractor(yes_no_extractor)
252 .build();
253 assert!(result.is_err());
254 assert!(result.unwrap_err().to_string().contains("at least 2"));
255 }
256
257 #[test]
258 fn builder_rejects_missing_vote_extractor() {
259 let p1 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
260 "YES", 10, 5,
261 )]));
262 let p2 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
263 "YES", 10, 5,
264 )]));
265 let result = VotingAgent::builder()
266 .voter(make_agent(p1, "a"))
267 .voter(make_agent(p2, "b"))
268 .build();
269 assert!(result.is_err());
270 assert!(result.unwrap_err().to_string().contains("vote_extractor"));
271 }
272
273 #[test]
274 fn builder_accepts_valid_config_without_tie_breaker() {
275 let p1 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
276 "YES", 10, 5,
277 )]));
278 let p2 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
279 "NO", 10, 5,
280 )]));
281 let result = VotingAgent::builder()
282 .voter(make_agent(p1, "a"))
283 .voter(make_agent(p2, "b"))
284 .vote_extractor(yes_no_extractor)
285 .build();
286 assert!(result.is_ok());
287 }
288
289 #[test]
290 fn builder_accepts_valid_config_with_tie_breaker() {
291 let p1 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
292 "YES", 10, 5,
293 )]));
294 let p2 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
295 "NO", 10, 5,
296 )]));
297 let result = VotingAgent::builder()
298 .voter(make_agent(p1, "a"))
299 .voter(make_agent(p2, "b"))
300 .vote_extractor(yes_no_extractor)
301 .tie_breaker(|votes| votes.last().unwrap().clone())
302 .build();
303 assert!(result.is_ok());
304 }
305
306 #[test]
311 fn builder_voters_bulk_method() {
312 let p1 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
313 "YES", 10, 5,
314 )]));
315 let p2 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
316 "NO", 10, 5,
317 )]));
318 let agents = vec![make_agent(p1, "a"), make_agent(p2, "b")];
319 let result = VotingAgent::builder()
320 .voters(agents)
321 .vote_extractor(yes_no_extractor)
322 .build();
323 assert!(result.is_ok());
324 }
325
326 #[tokio::test]
327 async fn unanimous_vote() {
328 let p1 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
329 "I vote YES",
330 100,
331 50,
332 )]));
333 let p2 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
334 "Definitely YES",
335 200,
336 80,
337 )]));
338 let p3 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
339 "YES please",
340 150,
341 60,
342 )]));
343
344 let voting = VotingAgent::builder()
345 .voter(make_agent(p1, "v1"))
346 .voter(make_agent(p2, "v2"))
347 .voter(make_agent(p3, "v3"))
348 .vote_extractor(yes_no_extractor)
349 .build()
350 .unwrap();
351
352 let result = voting.execute("should we?").await.unwrap();
353 assert_eq!(result.winner, "YES");
354 assert_eq!(result.tally["YES"], 3);
355 assert!(!result.tally.contains_key("NO"));
356 assert!(result.output.result.contains("YES"));
358 }
359
360 #[tokio::test]
361 async fn majority_vote_two_of_three() {
362 let p1 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
363 "I say YES",
364 100,
365 50,
366 )]));
367 let p2 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
368 "NO way", 200, 80,
369 )]));
370 let p3 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
371 "YES definitely",
372 150,
373 60,
374 )]));
375
376 let voting = VotingAgent::builder()
377 .voter(make_agent(p1, "v1"))
378 .voter(make_agent(p2, "v2"))
379 .voter(make_agent(p3, "v3"))
380 .vote_extractor(yes_no_extractor)
381 .build()
382 .unwrap();
383
384 let result = voting.execute("proceed?").await.unwrap();
385 assert_eq!(result.winner, "YES");
386 assert_eq!(result.tally["YES"], 2);
387 assert_eq!(result.tally["NO"], 1);
388 }
389
390 #[tokio::test]
391 async fn tie_broken_by_default_alphabetical() {
392 let p1 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
394 "NO thanks",
395 100,
396 50,
397 )]));
398 let p2 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
399 "YES sure", 200, 80,
400 )]));
401
402 let voting = VotingAgent::builder()
403 .voter(make_agent(p1, "v1"))
404 .voter(make_agent(p2, "v2"))
405 .vote_extractor(yes_no_extractor)
406 .build()
407 .unwrap();
408
409 let result = voting.execute("tie?").await.unwrap();
410 assert_eq!(result.winner, "NO");
412 assert_eq!(result.tally["YES"], 1);
413 assert_eq!(result.tally["NO"], 1);
414 }
415
416 #[tokio::test]
417 async fn tie_broken_by_custom_tie_breaker() {
418 let p1 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
419 "NO thanks",
420 100,
421 50,
422 )]));
423 let p2 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
424 "YES sure", 200, 80,
425 )]));
426
427 let voting = VotingAgent::builder()
428 .voter(make_agent(p1, "v1"))
429 .voter(make_agent(p2, "v2"))
430 .vote_extractor(yes_no_extractor)
431 .tie_breaker(|votes| votes.last().unwrap().clone()) .build()
433 .unwrap();
434
435 let result = voting.execute("tie?").await.unwrap();
436 assert_eq!(result.winner, "YES");
438 }
439
440 #[tokio::test]
441 async fn token_usage_accumulated_across_all_voters() {
442 let p1 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
443 "YES", 100, 50,
444 )]));
445 let p2 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
446 "YES", 200, 80,
447 )]));
448 let p3 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
449 "YES", 150, 60,
450 )]));
451
452 let voting = VotingAgent::builder()
453 .voter(make_agent(p1, "v1"))
454 .voter(make_agent(p2, "v2"))
455 .voter(make_agent(p3, "v3"))
456 .vote_extractor(yes_no_extractor)
457 .build()
458 .unwrap();
459
460 let result = voting.execute("go").await.unwrap();
461 assert_eq!(result.output.tokens_used.input_tokens, 450);
462 assert_eq!(result.output.tokens_used.output_tokens, 190);
463 }
464
465 #[tokio::test]
466 async fn error_carries_partial_usage() {
467 let p1 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
468 "YES", 100, 50,
469 )]));
470 let p2 = Arc::new(MockProvider::new(vec![]));
472
473 let voting = VotingAgent::builder()
474 .voter(make_agent(p1, "good"))
475 .voter(make_agent(p2, "bad"))
476 .vote_extractor(yes_no_extractor)
477 .build()
478 .unwrap();
479
480 let err = voting.execute("task").await.unwrap_err();
481 let partial = err.partial_usage();
482 assert!(
485 partial.input_tokens == 0 || partial.input_tokens >= 100,
486 "partial usage should be zero or include completed voter"
487 );
488 }
489
490 #[test]
491 fn debug_impl() {
492 let p1 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
493 "YES", 10, 5,
494 )]));
495 let p2 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
496 "NO", 10, 5,
497 )]));
498
499 let voting = VotingAgent::builder()
500 .voter(make_agent(p1, "a"))
501 .voter(make_agent(p2, "b"))
502 .vote_extractor(yes_no_extractor)
503 .build()
504 .unwrap();
505
506 let debug = format!("{voting:?}");
507 assert!(debug.contains("VotingAgent"));
508 assert!(debug.contains("voter_count"));
509 assert!(debug.contains("2"));
510 }
511
512 #[tokio::test]
513 async fn vote_result_contains_correct_tally() {
514 let p1 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
515 "YES agree",
516 10,
517 5,
518 )]));
519 let p2 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
520 "NO disagree",
521 10,
522 5,
523 )]));
524 let p3 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
525 "YES concur",
526 10,
527 5,
528 )]));
529 let p4 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
530 "NO object",
531 10,
532 5,
533 )]));
534 let p5 = Arc::new(MockProvider::new(vec![MockProvider::text_response(
535 "YES absolutely",
536 10,
537 5,
538 )]));
539
540 let voting = VotingAgent::builder()
541 .voter(make_agent(p1, "v1"))
542 .voter(make_agent(p2, "v2"))
543 .voter(make_agent(p3, "v3"))
544 .voter(make_agent(p4, "v4"))
545 .voter(make_agent(p5, "v5"))
546 .vote_extractor(yes_no_extractor)
547 .build()
548 .unwrap();
549
550 let result = voting.execute("vote").await.unwrap();
551 assert_eq!(result.winner, "YES");
552 assert_eq!(result.tally.len(), 2);
553 assert_eq!(result.tally["YES"], 3);
554 assert_eq!(result.tally["NO"], 2);
555 }
556}