lean_ctx/core/context_kernel/
result_fusion.rs1use std::collections::HashSet;
4
5#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
7pub struct ChildResult {
8 pub agent_id: String,
9 pub receipt_ref: String,
10 pub evidence: Vec<String>,
11 pub confidence: f64,
12 pub quality_score: f64,
13 pub contradicts: Vec<String>,
14}
15
16#[derive(Debug, Clone, Copy, Default, serde::Serialize, serde::Deserialize)]
18pub enum FusionStrategy {
19 #[default]
21 BestConfidence,
22 MajorityVote,
24 WeightedMerge,
26}
27
28#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
30pub struct FusionReport {
31 pub winning_agent: String,
32 pub merged_evidence: Vec<String>,
33 pub conflicts: Vec<Conflict>,
34 pub total_confidence: f64,
35 pub attribution_per_child: Vec<ChildAttribution>,
36}
37
38#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
40pub struct Conflict {
41 pub agent_a: String,
42 pub agent_b: String,
43 pub reason: String,
44}
45
46#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
48pub struct ChildAttribution {
49 pub agent_id: String,
50 pub weight: f64,
51 pub evidence_contributed: usize,
52}
53
54pub fn fuse_results(results: &[ChildResult], strategy: FusionStrategy) -> FusionReport {
56 if results.is_empty() {
57 return FusionReport::default();
58 }
59
60 let conflicts = detect_contradictions(results);
61 let (winning_index, selected, weights, total_confidence) = match strategy {
62 FusionStrategy::BestConfidence => best_confidence(results),
63 FusionStrategy::MajorityVote => majority_vote(results),
64 FusionStrategy::WeightedMerge => weighted_merge(results),
65 };
66 let merged_evidence = merge_evidence(results, &selected);
67 let attribution_per_child = results
68 .iter()
69 .enumerate()
70 .map(|(index, result)| ChildAttribution {
71 agent_id: result.agent_id.clone(),
72 weight: weights[index],
73 evidence_contributed: if selected.contains(&index) {
74 result.evidence.iter().collect::<HashSet<_>>().len()
75 } else {
76 0
77 },
78 })
79 .collect();
80
81 FusionReport {
82 winning_agent: results[winning_index].agent_id.clone(),
83 merged_evidence,
84 conflicts,
85 total_confidence,
86 attribution_per_child,
87 }
88}
89
90pub fn detect_contradictions(results: &[ChildResult]) -> Vec<Conflict> {
92 let mut seen = HashSet::new();
93 let mut conflicts = Vec::new();
94 for result in results {
95 for contradicted in &result.contradicts {
96 if contradicted == &result.agent_id {
97 continue;
98 }
99 let pair = if result.agent_id <= *contradicted {
100 (result.agent_id.clone(), contradicted.clone())
101 } else {
102 (contradicted.clone(), result.agent_id.clone())
103 };
104 if seen.insert(pair) {
105 conflicts.push(Conflict {
106 agent_a: result.agent_id.clone(),
107 agent_b: contradicted.clone(),
108 reason: "explicit contradiction".to_owned(),
109 });
110 }
111 }
112 }
113 conflicts
114}
115
116fn best_confidence(results: &[ChildResult]) -> (usize, Vec<usize>, Vec<f64>, f64) {
117 let winner = highest_score_index(results, |result| result.confidence);
118 let mut weights = vec![0.0; results.len()];
119 weights[winner] = 1.0;
120 (winner, vec![winner], weights, results[winner].confidence)
121}
122
123fn majority_vote(results: &[ChildResult]) -> (usize, Vec<usize>, Vec<f64>, f64) {
124 let groups = overlap_groups(results);
125 let mut selected = groups[0].clone();
126 for group in groups.into_iter().skip(1) {
127 let group_confidence = confidence_sum(results, &group);
128 let selected_confidence = confidence_sum(results, &selected);
129 if group.len() > selected.len()
130 || (group.len() == selected.len() && group_confidence > selected_confidence)
131 {
132 selected = group;
133 }
134 }
135 let winner = selected.iter().copied().fold(selected[0], |best, index| {
136 if results[index].confidence > results[best].confidence {
137 index
138 } else {
139 best
140 }
141 });
142 let weight = 1.0 / selected.len() as f64;
143 let mut weights = vec![0.0; results.len()];
144 for index in &selected {
145 weights[*index] = weight;
146 }
147 let confidence = confidence_sum(results, &selected) / selected.len() as f64;
148 (winner, selected, weights, confidence)
149}
150
151fn weighted_merge(results: &[ChildResult]) -> (usize, Vec<usize>, Vec<f64>, f64) {
152 let raw: Vec<f64> = results
153 .iter()
154 .map(|result| result.quality_score * result.confidence)
155 .collect();
156 let total: f64 = raw.iter().sum();
157 let weights = if total > 0.0 {
158 raw.iter().map(|weight| weight / total).collect()
159 } else {
160 vec![1.0 / results.len() as f64; results.len()]
161 };
162 let winner = highest_score_index(results, |result| result.quality_score * result.confidence);
163 let confidence = results
164 .iter()
165 .zip(&weights)
166 .map(|(result, weight)| result.confidence * weight)
167 .sum();
168 (winner, (0..results.len()).collect(), weights, confidence)
169}
170
171fn highest_score_index(results: &[ChildResult], score: impl Fn(&ChildResult) -> f64) -> usize {
172 (1..results.len()).fold(0, |best, index| {
173 if score(&results[index]) > score(&results[best]) {
174 index
175 } else {
176 best
177 }
178 })
179}
180
181fn overlap_groups(results: &[ChildResult]) -> Vec<Vec<usize>> {
182 let evidence: Vec<HashSet<&str>> = results
183 .iter()
184 .map(|result| result.evidence.iter().map(String::as_str).collect())
185 .collect();
186 let mut visited = vec![false; results.len()];
187 let mut groups = Vec::new();
188 for start in 0..results.len() {
189 if visited[start] {
190 continue;
191 }
192 visited[start] = true;
193 let mut group = Vec::new();
194 let mut stack = vec![start];
195 while let Some(current) = stack.pop() {
196 group.push(current);
197 for candidate in 0..results.len() {
198 if !visited[candidate] && !evidence[current].is_disjoint(&evidence[candidate]) {
199 visited[candidate] = true;
200 stack.push(candidate);
201 }
202 }
203 }
204 group.sort_unstable();
205 groups.push(group);
206 }
207 groups
208}
209
210fn merge_evidence(results: &[ChildResult], selected: &[usize]) -> Vec<String> {
211 let mut seen = HashSet::new();
212 selected
213 .iter()
214 .flat_map(|index| &results[*index].evidence)
215 .filter(|evidence| seen.insert((*evidence).clone()))
216 .cloned()
217 .collect()
218}
219
220fn confidence_sum(results: &[ChildResult], indices: &[usize]) -> f64 {
221 indices.iter().map(|index| results[*index].confidence).sum()
222}
223
224#[cfg(test)]
225mod tests {
226 use super::{ChildResult, FusionStrategy, detect_contradictions, fuse_results};
227
228 fn result(id: &str, evidence: &[&str], confidence: f64, quality: f64) -> ChildResult {
229 ChildResult {
230 agent_id: id.to_owned(),
231 receipt_ref: format!("receipt-{id}"),
232 evidence: evidence.iter().map(|item| (*item).to_owned()).collect(),
233 confidence,
234 quality_score: quality,
235 contradicts: Vec::new(),
236 }
237 }
238
239 #[test]
240 fn best_confidence_picks_highest() {
241 let results = [
242 result("a", &["one"], 0.4, 1.0),
243 result("b", &["two", "two"], 0.9, 0.5),
244 result("c", &["three"], 0.7, 1.0),
245 ];
246 let report = fuse_results(&results, FusionStrategy::BestConfidence);
247 assert_eq!(report.winning_agent, "b");
248 }
249
250 #[test]
251 fn majority_vote_picks_majority() {
252 let results = [
253 result("a", &["shared", "a"], 0.5, 1.0),
254 result("b", &["shared", "b"], 0.8, 1.0),
255 result("c", &["other"], 0.9, 1.0),
256 ];
257 let report = fuse_results(&results, FusionStrategy::MajorityVote);
258 assert_eq!(report.winning_agent, "b");
259 }
260
261 #[test]
262 fn weighted_merge_attributes_proportionally() {
263 let results = [
264 result("a", &["one"], 0.5, 1.0),
265 result("b", &["two"], 1.0, 1.0),
266 ];
267 let report = fuse_results(&results, FusionStrategy::WeightedMerge);
268 let sum: f64 = report
269 .attribution_per_child
270 .iter()
271 .map(|item| item.weight)
272 .sum();
273 assert!((sum - 1.0).abs() < f64::EPSILON);
274 }
275
276 #[test]
277 fn contradictions_detected() {
278 let mut a = result("a", &["one"], 0.5, 1.0);
279 a.contradicts.push("b".to_owned());
280 let conflicts = detect_contradictions(&[a, result("b", &["two"], 0.5, 1.0)]);
281 assert_eq!(conflicts.len(), 1);
282 }
283
284 #[test]
285 fn single_result_returns_identity() {
286 let only = result("only", &["one"], 0.75, 0.8);
287 let report = fuse_results(std::slice::from_ref(&only), FusionStrategy::WeightedMerge);
288 assert_eq!(report.winning_agent, only.agent_id);
289 }
290
291 #[test]
292 fn empty_results_returns_empty_report() {
293 let report = fuse_results(&[], FusionStrategy::MajorityVote);
294 assert!(report.winning_agent.is_empty());
295 }
296}