Skip to main content

tokio_prompt_orchestrator/
model_selector.rs

1//! Intelligent model selection based on cost / quality / latency tradeoffs.
2//!
3//! ## Key types
4//!
5//! - [`ModelProfile`] — static description of a model's characteristics.
6//! - [`SelectionCriteria`] — caller-supplied hard constraints (max cost, min quality, …).
7//! - [`SelectionStrategy`] — ranking algorithm applied to the candidate set.
8//! - [`ModelSelector`] — registry + selection engine.
9//! - [`ModelUsageTracker`] — records actual runtime usage to compare with predictions.
10//!
11//! ## Example
12//!
13//! ```rust
14//! use tokio_prompt_orchestrator::model_selector::{
15//!     ModelProfile, ModelSelector, SelectionCriteria, SelectionStrategy,
16//! };
17//!
18//! let mut selector = ModelSelector::new(SelectionStrategy::CheapestFirst);
19//! selector.register(ModelProfile {
20//!     id: "fast-cheap".to_string(),
21//!     cost_per_1k_tokens: 0.001,
22//!     avg_latency_ms: 200.0,
23//!     quality_score: 0.75,
24//!     max_context_tokens: 8192,
25//!     supports_streaming: true,
26//!     supports_tools: false,
27//! });
28//! let criteria = SelectionCriteria::default();
29//! let model = selector.select(&criteria).unwrap();
30//! assert_eq!(model.id, "fast-cheap");
31//! ```
32
33use std::collections::HashMap;
34
35use dashmap::DashMap;
36
37// ── ModelProfile ──────────────────────────────────────────────────────────────
38
39/// Static description of a model's cost, performance, and capability profile.
40#[derive(Debug, Clone)]
41pub struct ModelProfile {
42    /// Unique model identifier (e.g. `"claude-3-haiku"`, `"gpt-4o"`).
43    pub id: String,
44    /// Cost in USD per 1 000 tokens (blended input + output, or input-only if
45    /// that is the pricing model; callers can always override via `estimate_cost`).
46    pub cost_per_1k_tokens: f64,
47    /// Average observed or documented end-to-end latency in milliseconds.
48    pub avg_latency_ms: f64,
49    /// Normalised quality score in `[0, 1]` (e.g. from benchmark results).
50    pub quality_score: f64,
51    /// Maximum context window size in tokens.
52    pub max_context_tokens: usize,
53    /// Whether the model supports streaming token output.
54    pub supports_streaming: bool,
55    /// Whether the model supports function / tool calling.
56    pub supports_tools: bool,
57}
58
59// ── SelectionCriteria ─────────────────────────────────────────────────────────
60
61/// Hard constraints that a model must satisfy to be eligible for selection.
62#[derive(Debug, Clone, Default)]
63pub struct SelectionCriteria {
64    /// Reject models whose `cost_per_1k_tokens` exceeds this value.
65    pub max_cost_per_1k: Option<f64>,
66    /// Reject models whose `avg_latency_ms` exceeds this value.
67    pub max_latency_ms: Option<u64>,
68    /// Reject models whose `quality_score` is below this value.
69    pub min_quality: Option<f64>,
70    /// Reject models whose `max_context_tokens` is below this value.
71    pub min_context_tokens: Option<usize>,
72    /// If `true`, only models with `supports_streaming = true` are eligible.
73    pub requires_streaming: bool,
74    /// If `true`, only models with `supports_tools = true` are eligible.
75    pub requires_tools: bool,
76}
77
78// ── SelectionStrategy ─────────────────────────────────────────────────────────
79
80/// Algorithm used to rank the eligible model candidates.
81#[derive(Debug, Clone)]
82pub enum SelectionStrategy {
83    /// Pick the model with the lowest `cost_per_1k_tokens`.
84    CheapestFirst,
85    /// Pick the model with the lowest `avg_latency_ms`.
86    FastestFirst,
87    /// Pick the model with the highest `quality_score`.
88    BestQuality,
89    /// Weighted linear combination of quality, cost, and latency.
90    ///
91    /// Score = `quality_weight * quality`
92    ///       − `cost_weight   * (cost    / max_cost_in_set)`
93    ///       − `latency_weight * (latency / max_latency_in_set)`
94    ///
95    /// Weights need not sum to 1.0; they are relative.
96    Balanced {
97        /// Weight applied to `quality_score`.
98        cost_weight: f64,
99        /// Weight applied to normalised latency (negatively).
100        latency_weight: f64,
101        /// Weight applied to normalised cost (negatively).
102        quality_weight: f64,
103    },
104}
105
106// ── ModelSelector ─────────────────────────────────────────────────────────────
107
108/// Registry and selection engine for [`ModelProfile`]s.
109pub struct ModelSelector {
110    profiles: HashMap<String, ModelProfile>,
111    strategy: SelectionStrategy,
112}
113
114impl ModelSelector {
115    /// Create a new selector with no registered profiles.
116    pub fn new(strategy: SelectionStrategy) -> Self {
117        Self {
118            profiles: HashMap::new(),
119            strategy,
120        }
121    }
122
123    /// Register (or replace) a model profile.
124    pub fn register(&mut self, profile: ModelProfile) {
125        self.profiles.insert(profile.id.clone(), profile);
126    }
127
128    /// Return the best model according to the active strategy that satisfies
129    /// all hard constraints in `criteria`, or `None` if no model qualifies.
130    pub fn select(&self, criteria: &SelectionCriteria) -> Option<&ModelProfile> {
131        self.rank_all(criteria).into_iter().next()
132    }
133
134    /// Return all eligible profiles ranked from best to worst by the active strategy.
135    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    /// Estimate the cost in USD for a single request on `model_id`.
206    ///
207    /// Uses the profile's blended `cost_per_1k_tokens` applied to the total
208    /// token count (`tokens_in + tokens_out`).  Returns `0.0` if the model ID
209    /// is not registered.
210    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    /// Return the cheapest registered model whose `quality_score >= min_quality`,
218    /// or `None` if no model meets the threshold.
219    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    // ── Private helpers ───────────────────────────────────────────────────────
231
232    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// ── ModelUsageTracker ─────────────────────────────────────────────────────────
264
265/// Accumulated runtime usage for a single model.
266#[derive(Debug, Default, Clone)]
267pub struct ModelUsage {
268    /// Total number of API calls recorded.
269    pub calls: u64,
270    /// Total tokens consumed across all calls.
271    pub total_tokens: u64,
272    /// Total cost in USD across all calls.
273    pub total_cost: f64,
274    /// Rolling average latency in milliseconds.
275    pub avg_latency_ms: f64,
276}
277
278/// Thread-safe tracker for actual model usage.
279///
280/// Useful for comparing predicted vs. actual costs and computing
281/// cost-efficiency ratios.
282pub struct ModelUsageTracker {
283    data: DashMap<String, ModelUsage>,
284    /// Reference back to profiles so we can compute efficiency ratios.
285    profiles: DashMap<String, ModelProfile>,
286}
287
288impl ModelUsageTracker {
289    /// Create an empty tracker.
290    pub fn new() -> Self {
291        Self {
292            data: DashMap::new(),
293            profiles: DashMap::new(),
294        }
295    }
296
297    /// Register a model profile so that efficiency ratios can be computed.
298    pub fn register_profile(&self, profile: ModelProfile) {
299        self.profiles.insert(profile.id.clone(), profile);
300    }
301
302    /// Record a completed API call for `model_id`.
303    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        // Update rolling average: new_avg = (prev_avg * prev_calls + new_val) / new_calls
310        entry.avg_latency_ms = (entry.avg_latency_ms * prev_calls as f64 + latency_ms as f64)
311            / entry.calls as f64;
312    }
313
314    /// Return the cost-efficiency ratio for `model_id`.
315    ///
316    /// Efficiency = `quality_score / actual_cost_per_1k_tokens`.
317    ///
318    /// Returns `None` if the model has no recorded usage or no registered profile.
319    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    /// Return the recorded usage for `model_id`, if any.
337    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// ── Tests ─────────────────────────────────────────────────────────────────────
349
350#[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        // "best" (0.020) should be excluded; top pick should be "fast" (highest quality in set)
416        assert_eq!(ranked[0].id, "fast");
417    }
418
419    #[test]
420    fn balanced_score_ranking() {
421        // With cost_weight and latency_weight very small and quality_weight large,
422        // "best" (highest quality) should dominate.  We use near-zero penalties so
423        // that cost/latency differences cannot overcome the quality advantage.
424        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        // Sanity: with cost dominant and quality ignored, "cheap" should win.
434        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        // "cheap" costs 0.001 per 1k tokens; 500 in + 500 out = 1000 tokens total
447        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        // "cheap" has quality 0.70 (below), "fast" has 0.80 (above) and is cheaper than "best"
456        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), // impossible
479            ..Default::default()
480        };
481        assert!(sel.select(&criteria).is_none());
482    }
483}