tokio_prompt_orchestrator/
model_selector.rs1use std::collections::HashMap;
34
35use dashmap::DashMap;
36
37#[derive(Debug, Clone)]
41pub struct ModelProfile {
42 pub id: String,
44 pub cost_per_1k_tokens: f64,
47 pub avg_latency_ms: f64,
49 pub quality_score: f64,
51 pub max_context_tokens: usize,
53 pub supports_streaming: bool,
55 pub supports_tools: bool,
57}
58
59#[derive(Debug, Clone, Default)]
63pub struct SelectionCriteria {
64 pub max_cost_per_1k: Option<f64>,
66 pub max_latency_ms: Option<u64>,
68 pub min_quality: Option<f64>,
70 pub min_context_tokens: Option<usize>,
72 pub requires_streaming: bool,
74 pub requires_tools: bool,
76}
77
78#[derive(Debug, Clone)]
82pub enum SelectionStrategy {
83 CheapestFirst,
85 FastestFirst,
87 BestQuality,
89 Balanced {
97 cost_weight: f64,
99 latency_weight: f64,
101 quality_weight: f64,
103 },
104}
105
106pub struct ModelSelector {
110 profiles: HashMap<String, ModelProfile>,
111 strategy: SelectionStrategy,
112}
113
114impl ModelSelector {
115 pub fn new(strategy: SelectionStrategy) -> Self {
117 Self {
118 profiles: HashMap::new(),
119 strategy,
120 }
121 }
122
123 pub fn register(&mut self, profile: ModelProfile) {
125 self.profiles.insert(profile.id.clone(), profile);
126 }
127
128 pub fn select(&self, criteria: &SelectionCriteria) -> Option<&ModelProfile> {
131 self.rank_all(criteria).into_iter().next()
132 }
133
134 pub fn rank_all(&self, criteria: &SelectionCriteria) -> Vec<&ModelProfile> {
136 let mut candidates: Vec<&ModelProfile> = self
137 .profiles
138 .values()
139 .filter(|p| self.satisfies(p, criteria))
140 .collect();
141
142 match &self.strategy {
143 SelectionStrategy::CheapestFirst => {
144 candidates.sort_by(|a, b| {
145 a.cost_per_1k_tokens
146 .partial_cmp(&b.cost_per_1k_tokens)
147 .unwrap_or(std::cmp::Ordering::Equal)
148 });
149 }
150 SelectionStrategy::FastestFirst => {
151 candidates.sort_by(|a, b| {
152 a.avg_latency_ms
153 .partial_cmp(&b.avg_latency_ms)
154 .unwrap_or(std::cmp::Ordering::Equal)
155 });
156 }
157 SelectionStrategy::BestQuality => {
158 candidates.sort_by(|a, b| {
159 b.quality_score
160 .partial_cmp(&a.quality_score)
161 .unwrap_or(std::cmp::Ordering::Equal)
162 });
163 }
164 SelectionStrategy::Balanced {
165 cost_weight,
166 latency_weight,
167 quality_weight,
168 } => {
169 let max_cost = candidates
170 .iter()
171 .map(|p| p.cost_per_1k_tokens)
172 .fold(f64::NEG_INFINITY, f64::max);
173 let max_latency = candidates
174 .iter()
175 .map(|p| p.avg_latency_ms)
176 .fold(f64::NEG_INFINITY, f64::max);
177
178 let score = |p: &ModelProfile| -> f64 {
179 let norm_cost = if max_cost > 0.0 {
180 p.cost_per_1k_tokens / max_cost
181 } else {
182 0.0
183 };
184 let norm_latency = if max_latency > 0.0 {
185 p.avg_latency_ms / max_latency
186 } else {
187 0.0
188 };
189 quality_weight * p.quality_score
190 - cost_weight * norm_cost
191 - latency_weight * norm_latency
192 };
193
194 candidates.sort_by(|a, b| {
195 score(b)
196 .partial_cmp(&score(a))
197 .unwrap_or(std::cmp::Ordering::Equal)
198 });
199 }
200 }
201
202 candidates
203 }
204
205 pub fn estimate_cost(&self, model_id: &str, tokens_in: usize, tokens_out: usize) -> f64 {
211 self.profiles
212 .get(model_id)
213 .map(|p| p.cost_per_1k_tokens * (tokens_in + tokens_out) as f64 / 1000.0)
214 .unwrap_or(0.0)
215 }
216
217 pub fn cheapest_for_quality(&self, min_quality: f64) -> Option<&ModelProfile> {
220 self.profiles
221 .values()
222 .filter(|p| p.quality_score >= min_quality)
223 .min_by(|a, b| {
224 a.cost_per_1k_tokens
225 .partial_cmp(&b.cost_per_1k_tokens)
226 .unwrap_or(std::cmp::Ordering::Equal)
227 })
228 }
229
230 fn satisfies(&self, profile: &ModelProfile, criteria: &SelectionCriteria) -> bool {
233 if let Some(max_cost) = criteria.max_cost_per_1k {
234 if profile.cost_per_1k_tokens > max_cost {
235 return false;
236 }
237 }
238 if let Some(max_latency) = criteria.max_latency_ms {
239 if profile.avg_latency_ms > max_latency as f64 {
240 return false;
241 }
242 }
243 if let Some(min_quality) = criteria.min_quality {
244 if profile.quality_score < min_quality {
245 return false;
246 }
247 }
248 if let Some(min_ctx) = criteria.min_context_tokens {
249 if profile.max_context_tokens < min_ctx {
250 return false;
251 }
252 }
253 if criteria.requires_streaming && !profile.supports_streaming {
254 return false;
255 }
256 if criteria.requires_tools && !profile.supports_tools {
257 return false;
258 }
259 true
260 }
261}
262
263#[derive(Debug, Default, Clone)]
267pub struct ModelUsage {
268 pub calls: u64,
270 pub total_tokens: u64,
272 pub total_cost: f64,
274 pub avg_latency_ms: f64,
276}
277
278pub struct ModelUsageTracker {
283 data: DashMap<String, ModelUsage>,
284 profiles: DashMap<String, ModelProfile>,
286}
287
288impl ModelUsageTracker {
289 pub fn new() -> Self {
291 Self {
292 data: DashMap::new(),
293 profiles: DashMap::new(),
294 }
295 }
296
297 pub fn register_profile(&self, profile: ModelProfile) {
299 self.profiles.insert(profile.id.clone(), profile);
300 }
301
302 pub fn record(&self, model_id: &str, tokens: u64, cost: f64, latency_ms: u64) {
304 let mut entry = self.data.entry(model_id.to_string()).or_default();
305 let prev_calls = entry.calls;
306 entry.calls += 1;
307 entry.total_tokens += tokens;
308 entry.total_cost += cost;
309 entry.avg_latency_ms = (entry.avg_latency_ms * prev_calls as f64 + latency_ms as f64)
311 / entry.calls as f64;
312 }
313
314 pub fn cost_efficiency(&self, model_id: &str) -> Option<f64> {
320 let usage = self.data.get(model_id)?;
321 let profile = self.profiles.get(model_id)?;
322
323 let actual_cost_per_1k = if usage.total_tokens > 0 {
324 usage.total_cost / (usage.total_tokens as f64 / 1000.0)
325 } else {
326 return None;
327 };
328
329 if actual_cost_per_1k == 0.0 {
330 return None;
331 }
332
333 Some(profile.quality_score / actual_cost_per_1k)
334 }
335
336 pub fn usage(&self, model_id: &str) -> Option<ModelUsage> {
338 self.data.get(model_id).map(|u| u.clone())
339 }
340}
341
342impl Default for ModelUsageTracker {
343 fn default() -> Self {
344 Self::new()
345 }
346}
347
348#[cfg(test)]
351mod tests {
352 use super::*;
353
354 fn profile(id: &str, cost: f64, latency: f64, quality: f64) -> ModelProfile {
355 ModelProfile {
356 id: id.to_string(),
357 cost_per_1k_tokens: cost,
358 avg_latency_ms: latency,
359 quality_score: quality,
360 max_context_tokens: 8192,
361 supports_streaming: true,
362 supports_tools: true,
363 }
364 }
365
366 fn make_selector(strategy: SelectionStrategy) -> ModelSelector {
367 let mut sel = ModelSelector::new(strategy);
368 sel.register(profile("cheap", 0.001, 500.0, 0.70));
369 sel.register(profile("fast", 0.005, 100.0, 0.80));
370 sel.register(profile("best", 0.020, 800.0, 0.95));
371 sel
372 }
373
374 #[test]
375 fn cheapest_model_selected() {
376 let sel = make_selector(SelectionStrategy::CheapestFirst);
377 let m = sel.select(&SelectionCriteria::default()).unwrap();
378 assert_eq!(m.id, "cheap");
379 }
380
381 #[test]
382 fn fastest_model_selected() {
383 let sel = make_selector(SelectionStrategy::FastestFirst);
384 let m = sel.select(&SelectionCriteria::default()).unwrap();
385 assert_eq!(m.id, "fast");
386 }
387
388 #[test]
389 fn best_quality_selected() {
390 let sel = make_selector(SelectionStrategy::BestQuality);
391 let m = sel.select(&SelectionCriteria::default()).unwrap();
392 assert_eq!(m.id, "best");
393 }
394
395 #[test]
396 fn quality_threshold_filtering() {
397 let sel = make_selector(SelectionStrategy::CheapestFirst);
398 let criteria = SelectionCriteria {
399 min_quality: Some(0.85),
400 ..Default::default()
401 };
402 let m = sel.select(&criteria).unwrap();
403 assert_eq!(m.id, "best");
404 }
405
406 #[test]
407 fn cost_filtering_excludes_expensive() {
408 let sel = make_selector(SelectionStrategy::BestQuality);
409 let criteria = SelectionCriteria {
410 max_cost_per_1k: Some(0.006),
411 ..Default::default()
412 };
413 let ranked = sel.rank_all(&criteria);
414 assert!(ranked.iter().all(|p| p.cost_per_1k_tokens <= 0.006));
415 assert_eq!(ranked[0].id, "fast");
417 }
418
419 #[test]
420 fn balanced_score_ranking() {
421 let sel = make_selector(SelectionStrategy::Balanced {
425 cost_weight: 0.01,
426 latency_weight: 0.01,
427 quality_weight: 10.0,
428 });
429 let ranked = sel.rank_all(&SelectionCriteria::default());
430 assert!(!ranked.is_empty());
431 assert_eq!(ranked[0].id, "best");
432
433 let sel2 = make_selector(SelectionStrategy::Balanced {
435 cost_weight: 10.0,
436 latency_weight: 0.01,
437 quality_weight: 0.01,
438 });
439 let ranked2 = sel2.rank_all(&SelectionCriteria::default());
440 assert_eq!(ranked2[0].id, "cheap");
441 }
442
443 #[test]
444 fn cost_estimation() {
445 let sel = make_selector(SelectionStrategy::CheapestFirst);
446 let cost = sel.estimate_cost("cheap", 500, 500);
448 assert!((cost - 0.001).abs() < 1e-9);
449 }
450
451 #[test]
452 fn cheapest_for_quality_threshold() {
453 let sel = make_selector(SelectionStrategy::CheapestFirst);
454 let m = sel.cheapest_for_quality(0.79).unwrap();
455 assert_eq!(m.id, "fast");
457 }
458
459 #[test]
460 fn usage_tracker_records_and_efficiency() {
461 let tracker = ModelUsageTracker::new();
462 tracker.register_profile(profile("fast", 0.005, 100.0, 0.80));
463 tracker.record("fast", 1000, 0.005, 100);
464 tracker.record("fast", 1000, 0.005, 200);
465 let usage = tracker.usage("fast").unwrap();
466 assert_eq!(usage.calls, 2);
467 assert_eq!(usage.total_tokens, 2000);
468 assert!((usage.avg_latency_ms - 150.0).abs() < 1e-6);
469 let eff = tracker.cost_efficiency("fast");
470 assert!(eff.is_some());
471 assert!(eff.unwrap() > 0.0);
472 }
473
474 #[test]
475 fn no_match_returns_none() {
476 let sel = make_selector(SelectionStrategy::CheapestFirst);
477 let criteria = SelectionCriteria {
478 min_quality: Some(1.1), ..Default::default()
480 };
481 assert!(sel.select(&criteria).is_none());
482 }
483}