1use crate::types::CompactStr;
7use crate::utils::current_timestamp;
8use hashbrown::HashMap;
9use serde::{Deserialize, Serialize};
10
11use crate::tools::result_metadata::ResultMetadata;
12
13#[derive(Debug, Clone, Serialize, Deserialize)]
15pub struct ToolEffectiveness {
16 pub tool_name: CompactStr,
17
18 pub success_rate: f32,
20
21 pub avg_result_quality: f32,
23
24 pub usage_count: usize,
26
27 pub success_count: usize,
29
30 pub last_used_timestamp: u64,
32
33 #[serde(default)]
35 pub failure_modes: Vec<ToolFailureMode>,
36
37 #[serde(default)]
39 pub avg_execution_time_ms: f32,
40}
41
42#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, Hash)]
43pub enum ToolFailureMode {
44 Timeout,
45 NoResults,
46 InvalidArgs,
47 ParseError,
48 PermissionDenied,
49 Unknown,
50}
51
52impl ToolEffectiveness {
53 pub fn new(tool_name: impl Into<CompactStr>) -> Self {
54 Self {
55 tool_name: tool_name.into(),
56 success_rate: 0.0,
57 avg_result_quality: 0.0,
58 usage_count: 0,
59 success_count: 0,
60 last_used_timestamp: 0,
61 failure_modes: vec![],
62 avg_execution_time_ms: 0.0,
63 }
64 }
65
66 pub fn record_success(&mut self, quality: f32, execution_time_ms: f32) {
68 self.usage_count += 1;
69 self.success_count += 1;
70 self.last_used_timestamp = current_timestamp();
71
72 self.avg_result_quality =
74 (self.avg_result_quality * (self.success_count - 1) as f32 + quality) / self.success_count as f32;
75
76 self.avg_execution_time_ms = (self.avg_execution_time_ms * (self.success_count - 1) as f32 + execution_time_ms)
78 / self.success_count as f32;
79
80 self.update_success_rate();
81 }
82
83 pub fn record_failure(&mut self, failure_mode: ToolFailureMode, execution_time_ms: f32) {
85 self.usage_count += 1;
86 self.last_used_timestamp = current_timestamp();
87
88 if !self.failure_modes.iter().any(|m| m == &failure_mode) {
90 self.failure_modes.push(failure_mode);
91 }
92
93 let success_count_f = (self.success_count + 1) as f32;
95 self.avg_execution_time_ms = (self.avg_execution_time_ms * (self.success_count as f32 / success_count_f))
96 + (execution_time_ms / success_count_f);
97
98 self.update_success_rate();
99 }
100
101 fn update_success_rate(&mut self) {
102 if self.usage_count > 0 {
103 self.success_rate = self.success_count as f32 / self.usage_count as f32;
104 }
105 }
106
107 pub fn effectiveness_score(&self) -> f32 {
109 if self.usage_count == 0 {
110 return 0.5; }
112
113 (self.success_rate * 0.6) + (self.avg_result_quality * 0.4)
115 }
116
117 pub fn is_reliable(&self) -> bool {
119 self.usage_count >= 3 && self.success_rate > 0.7
120 }
121
122 pub fn time_since_last_use_seconds(&self) -> u64 {
124 if self.last_used_timestamp == 0 {
125 u64::MAX
126 } else {
127 current_timestamp().saturating_sub(self.last_used_timestamp)
128 }
129 }
130}
131
132#[derive(Debug, Clone)]
134pub struct ToolSelectionContext {
135 pub task_description: String,
137
138 pub prior_tools_used: Vec<String>,
140
141 pub prior_result_qualities: Vec<f32>,
143
144 pub tool_effectiveness: HashMap<CompactStr, ToolEffectiveness>,
146}
147
148pub trait ToolSelector: Send + Sync {
150 fn select_tool(&self, context: &ToolSelectionContext, candidates: &[&str]) -> Option<String>;
151}
152
153pub struct AdaptiveToolSelector;
155
156impl ToolSelector for AdaptiveToolSelector {
157 fn select_tool(&self, context: &ToolSelectionContext, candidates: &[&str]) -> Option<String> {
158 if candidates.is_empty() {
159 return None;
160 }
161
162 let mut scored: Vec<(String, f32)> = candidates
164 .iter()
165 .map(|tool| {
166 let name = (*tool).to_owned();
167 let score = score_tool(&name, context);
168 (name, score)
169 })
170 .collect();
171
172 scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
174
175 scored.first().map(|(name, _)| name.clone())
176 }
177}
178
179fn score_tool(tool_name: &str, context: &ToolSelectionContext) -> f32 {
181 let mut score = 0.5; if let Some(eff) = context.tool_effectiveness.get(tool_name) {
185 score += eff.effectiveness_score() * 0.4;
187
188 let time_score = 1.0 - (eff.avg_execution_time_ms / 10000.0).min(1.0);
190 score += time_score * 0.1;
191
192 if !eff.failure_modes.is_empty() {
194 let failure_penalty = (eff.failure_modes.len() as f32) * 0.1;
195 score -= failure_penalty;
196 }
197 }
198
199 if context.prior_tools_used.iter().any(|s| s == tool_name) {
201 score -= 0.15; }
203
204 score.clamp(0.0, 1.0)
206}
207
208pub struct ToolEffectivenessTracker {
210 effectiveness: HashMap<CompactStr, ToolEffectiveness>,
211}
212
213impl ToolEffectivenessTracker {
214 pub fn new() -> Self {
215 Self { effectiveness: HashMap::new() }
216 }
217
218 fn get_or_create(&mut self, tool_name: &str) -> &mut ToolEffectiveness {
220 self.effectiveness
221 .entry(CompactStr::from(tool_name))
222 .or_insert_with(|| ToolEffectiveness::new(tool_name))
223 }
224
225 pub fn record_success(&mut self, tool_name: &str, metadata: &ResultMetadata, execution_time_ms: f32) {
227 let quality = metadata.quality_score();
228 self.get_or_create(tool_name).record_success(quality, execution_time_ms);
229 }
230
231 pub fn record_failure(&mut self, tool_name: &str, mode: ToolFailureMode, execution_time_ms: f32) {
233 self.get_or_create(tool_name).record_failure(mode, execution_time_ms);
234 }
235
236 pub fn snapshot(&self) -> HashMap<CompactStr, ToolEffectiveness> {
238 self.effectiveness.clone()
239 }
240
241 pub fn get(&self, tool_name: &str) -> Option<&ToolEffectiveness> {
243 self.effectiveness.get(tool_name)
244 }
245
246 pub fn sorted_by_effectiveness(&self) -> Vec<(CompactStr, f32)> {
248 let mut tools: Vec<_> = self
249 .effectiveness
250 .iter()
251 .map(|(name, eff)| (name.clone(), eff.effectiveness_score()))
252 .collect();
253
254 tools.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
255 tools
256 }
257}
258
259impl Default for ToolEffectivenessTracker {
260 fn default() -> Self {
261 Self::new()
262 }
263}
264
265#[cfg(test)]
266mod tests {
267 use super::*;
268
269 #[test]
270 fn test_tool_effectiveness_success() {
271 let mut eff = ToolEffectiveness::new("grep".to_owned());
272 eff.record_success(0.9, 100.0);
273 eff.record_success(0.85, 110.0);
274
275 assert_eq!(eff.success_count, 2);
276 assert_eq!(eff.usage_count, 2);
277 assert!((eff.success_rate - 1.0).abs() < f32::EPSILON);
278 assert!(eff.avg_result_quality > 0.8);
279 }
280
281 #[test]
282 fn test_tool_effectiveness_failure() {
283 let mut eff = ToolEffectiveness::new("find".to_owned());
284 eff.record_failure(ToolFailureMode::Timeout, 5000.0);
285
286 assert_eq!(eff.success_count, 0);
287 assert_eq!(eff.usage_count, 1);
288 assert!((eff.success_rate - 0.0).abs() < f32::EPSILON);
289 assert!(eff.failure_modes.contains(&ToolFailureMode::Timeout));
290 }
291
292 #[test]
293 fn test_adaptive_selector() {
294 let selector = AdaptiveToolSelector;
295 let mut effectiveness = HashMap::new();
296
297 let mut grep_eff = ToolEffectiveness::new("grep");
298 grep_eff.record_success(0.9, 100.0);
299 effectiveness.insert(CompactStr::from("grep"), grep_eff);
300
301 let mut find_eff = ToolEffectiveness::new("find");
302 find_eff.record_failure(ToolFailureMode::Timeout, 5000.0);
303 effectiveness.insert(CompactStr::from("find"), find_eff);
304
305 let context = ToolSelectionContext {
306 task_description: "find error patterns".to_owned(),
307 prior_tools_used: vec![],
308 prior_result_qualities: vec![],
309 tool_effectiveness: effectiveness,
310 };
311
312 let selected = selector.select_tool(&context, &["grep", "find"]);
313 assert_eq!(selected, Some("grep".to_owned()));
314 }
315
316 #[test]
317 fn test_tool_diversity_penalty() {
318 let selector = AdaptiveToolSelector;
319 let effectiveness = HashMap::new();
320
321 let context = ToolSelectionContext {
322 task_description: "find files".to_owned(),
323 prior_tools_used: vec!["grep".to_owned()],
324 prior_result_qualities: vec![],
325 tool_effectiveness: effectiveness,
326 };
327
328 let selected = selector.select_tool(&context, &["grep", "find"]);
329 assert_eq!(selected, Some("find".to_owned()));
331 }
332
333 #[test]
334 fn test_effectiveness_tracker() {
335 let mut tracker = ToolEffectivenessTracker::new();
336 let meta = ResultMetadata::success(0.8, 0.8);
337
338 tracker.record_success("grep", &meta, 100.0);
339 tracker.record_success("grep", &meta, 110.0);
340
341 let sorted = tracker.sorted_by_effectiveness();
342 assert_eq!(sorted[0].0, "grep");
343 assert!(sorted[0].1 > 0.7);
344 }
345}