Skip to main content

heartbit_core/agent/
voting.rs

1//! Majority-voting workflow agent.
2//!
3//! Runs N voter agents in parallel on the same task, extracts a vote from each
4//! output, and returns the full output of the first voter whose vote matches the
5//! majority. Ties are resolved by an optional `tie_breaker` (defaults to first
6//! alphabetically).
7
8use 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
19/// Extracts a vote string from an agent's output text.
20type VoteExtractor = Box<dyn Fn(&str) -> String + Send + Sync>;
21
22/// Resolves ties when multiple votes share the highest count.
23/// Receives the tied vote strings and must return one of them.
24type TieBreaker = Box<dyn Fn(&[String]) -> String + Send + Sync>;
25
26/// The result of a voting round, including the winning vote, the full tally,
27/// and the output from the first voter that cast the winning vote.
28#[derive(Debug)]
29pub struct VoteResult {
30    /// The vote string that won.
31    pub winner: String,
32    /// Vote string → number of voters that cast it.
33    pub tally: HashMap<String, usize>,
34    /// The full `AgentOutput` from the first voter whose vote matched the winner.
35    pub output: AgentOutput,
36}
37
38/// Orchestrates majority voting across N agents running in parallel.
39pub 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
53/// Builder for [`VotingAgent`].
54pub 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    /// Create a new [`VotingAgentBuilder`].
62    pub fn builder() -> VotingAgentBuilder<P> {
63        VotingAgentBuilder {
64            voters: Vec::new(),
65            vote_extractor: None,
66            tie_breaker: None,
67        }
68    }
69
70    /// Execute all voters in parallel, tally votes, and return the winning result.
71    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        // Collect results in completion order, tracking accumulated usage for
84        // partial-usage-on-error semantics.
85        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                // AP6: preserve completed siblings' usage on a panic.
91                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        // Sort by original index for deterministic vote ordering.
100        outputs.sort_by_key(|(idx, _)| *idx);
101
102        // Extract votes and build tally.
103        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        // Find the maximum vote count.
114        let max_count = tally.values().copied().max().unwrap_or(0);
115
116        // Collect all votes tied at the maximum.
117        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        // Find the first voter (by original index) whose vote matches the winner.
131        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        // Accumulate tool_calls and cost from all voters (usage already tracked
139        // in `total_usage` during JoinSet collection above).
140        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    /// Add a voter agent. Wraps it in `Arc` for concurrent sharing.
167    pub fn voter(mut self, agent: AgentRunner<P>) -> Self {
168        self.voters.push(Arc::new(agent));
169        self
170    }
171
172    /// Add multiple voter agents.
173    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    /// Set the vote extractor function.
179    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    /// Set an optional tie-breaker function.
185    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    /// Build the [`VotingAgent`]. Requires at least 2 voters and a vote extractor.
191    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                // Default: first alphabetically (votes are already sorted).
203                votes[0].clone()
204            })
205        });
206        Ok(VotingAgent {
207            voters: self.voters,
208            vote_extractor,
209            tie_breaker,
210        })
211    }
212}
213
214// ===========================================================================
215// Tests
216// ===========================================================================
217
218#[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    // -----------------------------------------------------------------------
232    // Builder validation tests
233    // -----------------------------------------------------------------------
234
235    #[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    // -----------------------------------------------------------------------
307    // Execution tests
308    // -----------------------------------------------------------------------
309
310    #[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        // Output should be from one of the YES voters
357        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        // 2 voters: one YES, one NO — tie. Default tie-breaker picks alphabetically first.
393        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        // "NO" < "YES" alphabetically
411        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()) // pick last alphabetically
432            .build()
433            .unwrap();
434
435        let result = voting.execute("tie?").await.unwrap();
436        // Custom tie-breaker picks last alphabetically: "YES"
437        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        // Second provider has no responses -> error
471        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        // JoinSet ordering is non-deterministic: the successful voter may
483        // or may not finish before the error is collected.
484        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}