Skip to main content

tokio_prompt_orchestrator/
smart_router.rs

1//! # Cost-Aware Smart Router
2//!
3//! Routes prompts to the cheapest provider meeting the request's SLA:
4//!   - Latency budget (max acceptable latency in ms)
5//!   - Quality floor (minimum acceptable quality score)
6//!   - Cost ceiling (maximum USD per request)
7//!
8//! ## Algorithm
9//!
10//! 1. Filter providers: remove unhealthy and those exceeding the latency budget.
11//! 2. Filter providers: remove those below the quality floor.
12//! 3. Filter providers: remove those exceeding the cost ceiling.
13//! 4. If `prefer_quality` is set, select the highest-quality remaining provider.
14//! 5. Otherwise select the cheapest provider that meets all constraints.
15//! 6. The remaining providers (sorted cheapest-first) become the fallback list.
16//!
17//! Pricing is configurable and can be hot-reloaded from a TOML string at
18//! runtime via [`SmartRouter::load_pricing_toml`].
19
20// ---------------------------------------------------------------------------
21// ModelPricing
22// ---------------------------------------------------------------------------
23
24/// Per-token pricing and capability metadata for one provider/model combination.
25#[derive(Debug, Clone)]
26pub struct ModelPricing {
27    /// Stable provider identifier (e.g. `"anthropic"`, `"openai"`).
28    pub provider_id: String,
29    /// Model name (e.g. `"claude-3-5-sonnet-20241022"`).
30    pub model: String,
31    /// USD cost per **1 million** input tokens.
32    pub input_per_million: f64,
33    /// USD cost per **1 million** output tokens.
34    pub output_per_million: f64,
35    /// Estimated quality score in `[0.0, 1.0]` derived from public benchmarks.
36    pub quality_score: f32,
37    /// Typical p95 latency for this model in milliseconds.
38    pub typical_p95_ms: u64,
39}
40
41impl ModelPricing {
42    /// Estimate the total cost of a single request in USD.
43    ///
44    /// # Panics
45    ///
46    /// Never panics.
47    pub fn estimate_cost(&self, input_tokens: usize, estimated_output_tokens: usize) -> f64 {
48        let input_cost = (input_tokens as f64 / 1_000_000.0) * self.input_per_million;
49        let output_cost =
50            (estimated_output_tokens as f64 / 1_000_000.0) * self.output_per_million;
51        input_cost + output_cost
52    }
53}
54
55// ---------------------------------------------------------------------------
56// RoutingRequirements
57// ---------------------------------------------------------------------------
58
59/// SLA constraints for a single routing request.
60#[derive(Debug, Clone)]
61pub struct RoutingRequirements {
62    /// Maximum acceptable p95 latency in milliseconds. `None` = no constraint.
63    pub max_latency_ms: Option<u64>,
64    /// Minimum acceptable quality score. `None` = no constraint.
65    pub min_quality: Option<f32>,
66    /// Maximum cost per request in USD. `None` = no constraint.
67    pub max_cost_usd: Option<f64>,
68    /// Estimated number of input tokens (used for cost estimation).
69    pub estimated_input_tokens: usize,
70    /// Estimated output tokens (rough guess for cost estimation).
71    pub estimated_output_tokens: usize,
72    /// When `true`, ignore cost constraints and always pick the highest-quality
73    /// provider that satisfies latency requirements.
74    pub prefer_quality: bool,
75}
76
77impl Default for RoutingRequirements {
78    /// Construct permissive requirements: no latency/quality/cost constraints,
79    /// 512 estimated input tokens, 256 estimated output tokens.
80    ///
81    /// # Panics
82    ///
83    /// Never panics.
84    fn default() -> Self {
85        Self {
86            max_latency_ms: None,
87            min_quality: None,
88            max_cost_usd: None,
89            estimated_input_tokens: 512,
90            estimated_output_tokens: 256,
91            prefer_quality: false,
92        }
93    }
94}
95
96// ---------------------------------------------------------------------------
97// RoutingDecision
98// ---------------------------------------------------------------------------
99
100/// The result of a routing call: primary provider plus ordered fallbacks.
101#[derive(Debug, Clone)]
102pub struct RoutingDecision {
103    /// Chosen provider identifier.
104    pub provider_id: String,
105    /// Model name for the chosen provider.
106    pub model: String,
107    /// Estimated cost for this request in USD.
108    pub estimated_cost_usd: f64,
109    /// Human-readable explanation of why this provider was selected.
110    pub reason: String,
111    /// Ordered list of fallback provider IDs to try if the primary fails.
112    /// Sorted cheapest-first (or quality-first when `prefer_quality` is set).
113    pub fallbacks: Vec<String>,
114}
115
116// ---------------------------------------------------------------------------
117// SmartRouter
118// ---------------------------------------------------------------------------
119
120/// Routes incoming requests to the cheapest (or highest-quality) provider
121/// that satisfies all SLA constraints.
122pub struct SmartRouter {
123    pricing: Vec<ModelPricing>,
124    health: crate::provider_health::ProviderHealthMonitor,
125}
126
127impl SmartRouter {
128    /// Create a new router with the given pricing table and health monitor.
129    ///
130    /// # Panics
131    ///
132    /// Never panics.
133    pub fn new(
134        pricing: Vec<ModelPricing>,
135        health: crate::provider_health::ProviderHealthMonitor,
136    ) -> Self {
137        Self { pricing, health }
138    }
139
140    /// Route a request, returning the best [`RoutingDecision`] or `None` if no
141    /// suitable provider exists after applying all filters.
142    ///
143    /// # Panics
144    ///
145    /// Never panics.
146    pub async fn route(&self, req: &RoutingRequirements) -> Option<RoutingDecision> {
147        // Step 1 — remove unhealthy providers.
148        let healthy = self.filter_by_health(&self.pricing).await;
149        // Step 2 — remove providers that exceed the latency budget.
150        let within_latency = self.filter_by_latency(healthy, req.max_latency_ms);
151        // Step 3 — remove providers below the quality floor.
152        let quality_ok = self.filter_by_quality(within_latency, req.min_quality);
153        // Step 4 — remove providers that exceed the cost ceiling.
154        let cost_ok = self.filter_by_cost(quality_ok, req);
155
156        if cost_ok.is_empty() {
157            return None;
158        }
159
160        // Select primary provider.
161        let primary = if req.prefer_quality {
162            self.best_quality(cost_ok.clone())?
163        } else {
164            self.cheapest(cost_ok.clone(), req)?
165        };
166
167        let estimated_cost = primary
168            .estimate_cost(req.estimated_input_tokens, req.estimated_output_tokens);
169
170        let reason = if req.prefer_quality {
171            format!(
172                "quality-preferred routing: {} (quality={:.2})",
173                primary.model, primary.quality_score
174            )
175        } else {
176            format!(
177                "cost-optimised routing: {} (est. ${:.6})",
178                primary.model, estimated_cost
179            )
180        };
181
182        // Build fallback list: remaining providers sorted cheapest/quality-first.
183        let mut fallbacks: Vec<&ModelPricing> = cost_ok
184            .iter()
185            .filter(|p| p.provider_id != primary.provider_id)
186            .copied()
187            .collect();
188
189        if req.prefer_quality {
190            fallbacks.sort_by(|a, b| {
191                b.quality_score
192                    .partial_cmp(&a.quality_score)
193                    .unwrap_or(std::cmp::Ordering::Equal)
194            });
195        } else {
196            fallbacks.sort_by(|a, b| {
197                let ca = a.estimate_cost(req.estimated_input_tokens, req.estimated_output_tokens);
198                let cb = b.estimate_cost(req.estimated_input_tokens, req.estimated_output_tokens);
199                ca.partial_cmp(&cb).unwrap_or(std::cmp::Ordering::Equal)
200            });
201        }
202
203        Some(RoutingDecision {
204            provider_id: primary.provider_id.clone(),
205            model: primary.model.clone(),
206            estimated_cost_usd: estimated_cost,
207            reason,
208            fallbacks: fallbacks.iter().map(|p| p.provider_id.clone()).collect(),
209        })
210    }
211
212    /// Return the built-in pricing table for common LLM models.
213    ///
214    /// Prices are in USD per 1 million tokens as of early 2025.
215    ///
216    /// # Panics
217    ///
218    /// Never panics.
219    pub fn default_pricing() -> Vec<ModelPricing> {
220        vec![
221            // Anthropic
222            ModelPricing {
223                provider_id: "anthropic".to_string(),
224                model: "claude-3-5-sonnet-20241022".to_string(),
225                input_per_million: 3.0,
226                output_per_million: 15.0,
227                quality_score: 0.92,
228                typical_p95_ms: 4500,
229            },
230            ModelPricing {
231                provider_id: "anthropic-haiku".to_string(),
232                model: "claude-3-haiku-20240307".to_string(),
233                input_per_million: 0.25,
234                output_per_million: 1.25,
235                quality_score: 0.74,
236                typical_p95_ms: 1800,
237            },
238            // OpenAI
239            ModelPricing {
240                provider_id: "openai-gpt4o".to_string(),
241                model: "gpt-4o".to_string(),
242                input_per_million: 2.5,
243                output_per_million: 10.0,
244                quality_score: 0.90,
245                typical_p95_ms: 5000,
246            },
247            ModelPricing {
248                provider_id: "openai-gpt4o-mini".to_string(),
249                model: "gpt-4o-mini".to_string(),
250                input_per_million: 0.15,
251                output_per_million: 0.60,
252                quality_score: 0.78,
253                typical_p95_ms: 2000,
254            },
255            ModelPricing {
256                provider_id: "openai-gpt35".to_string(),
257                model: "gpt-3.5-turbo".to_string(),
258                input_per_million: 0.50,
259                output_per_million: 1.50,
260                quality_score: 0.65,
261                typical_p95_ms: 2500,
262            },
263            // Meta / local
264            ModelPricing {
265                provider_id: "meta-llama".to_string(),
266                model: "llama-3.1-8b-instruct".to_string(),
267                input_per_million: 0.06,
268                output_per_million: 0.06,
269                quality_score: 0.60,
270                typical_p95_ms: 1200,
271            },
272        ]
273    }
274
275    /// Hot-reload pricing from a TOML string.
276    ///
277    /// The TOML document must contain a top-level `[[model]]` array, where each
278    /// entry has the fields: `provider_id`, `model`, `input_per_million`,
279    /// `output_per_million`, `quality_score`, `typical_p95_ms`.
280    ///
281    /// On success, replaces the current pricing table and returns the number of
282    /// models loaded.  On parse failure, the existing table is **not** modified.
283    ///
284    /// # Errors
285    ///
286    /// Returns `Err(String)` describing the parse failure.
287    ///
288    /// # Panics
289    ///
290    /// Never panics.
291    pub fn load_pricing_toml(&mut self, toml_str: &str) -> Result<usize, String> {
292        #[derive(serde::Deserialize)]
293        struct PricingFile {
294            model: Vec<ModelPricingToml>,
295        }
296
297        #[derive(serde::Deserialize)]
298        struct ModelPricingToml {
299            provider_id: String,
300            model: String,
301            input_per_million: f64,
302            output_per_million: f64,
303            quality_score: f32,
304            typical_p95_ms: u64,
305        }
306
307        let parsed: PricingFile =
308            toml::from_str(toml_str).map_err(|e| format!("TOML parse error: {e}"))?;
309
310        let new_pricing: Vec<ModelPricing> = parsed
311            .model
312            .into_iter()
313            .map(|m| ModelPricing {
314                provider_id: m.provider_id,
315                model: m.model,
316                input_per_million: m.input_per_million,
317                output_per_million: m.output_per_million,
318                quality_score: m.quality_score,
319                typical_p95_ms: m.typical_p95_ms,
320            })
321            .collect();
322
323        let count = new_pricing.len();
324        self.pricing = new_pricing;
325        Ok(count)
326    }
327
328    // -----------------------------------------------------------------------
329    // Private filter helpers
330    // -----------------------------------------------------------------------
331
332    /// Filter out providers whose health monitor marks them as unusable.
333    async fn filter_by_health<'a>(&self, pricing: &'a [ModelPricing]) -> Vec<&'a ModelPricing> {
334        let mut result = Vec::with_capacity(pricing.len());
335        for p in pricing {
336            if self.health.is_usable(&p.provider_id).await {
337                result.push(p);
338            }
339        }
340        result
341    }
342
343    /// Filter out providers whose `typical_p95_ms` exceeds `max_ms`.
344    fn filter_by_latency<'a>(
345        &self,
346        providers: Vec<&'a ModelPricing>,
347        max_ms: Option<u64>,
348    ) -> Vec<&'a ModelPricing> {
349        match max_ms {
350            None => providers,
351            Some(limit) => providers
352                .into_iter()
353                .filter(|p| p.typical_p95_ms <= limit)
354                .collect(),
355        }
356    }
357
358    /// Filter out providers whose `quality_score` is below `min_quality`.
359    fn filter_by_quality<'a>(
360        &self,
361        providers: Vec<&'a ModelPricing>,
362        min_quality: Option<f32>,
363    ) -> Vec<&'a ModelPricing> {
364        match min_quality {
365            None => providers,
366            Some(floor) => providers
367                .into_iter()
368                .filter(|p| p.quality_score >= floor)
369                .collect(),
370        }
371    }
372
373    /// Filter out providers whose estimated cost exceeds `req.max_cost_usd`.
374    fn filter_by_cost<'a>(
375        &self,
376        providers: Vec<&'a ModelPricing>,
377        req: &RoutingRequirements,
378    ) -> Vec<&'a ModelPricing> {
379        match req.max_cost_usd {
380            None => providers,
381            Some(ceiling) => providers
382                .into_iter()
383                .filter(|p| {
384                    p.estimate_cost(req.estimated_input_tokens, req.estimated_output_tokens)
385                        <= ceiling
386                })
387                .collect(),
388        }
389    }
390
391    /// Return the cheapest provider from `providers` for the given request.
392    fn cheapest<'a>(
393        &self,
394        providers: Vec<&'a ModelPricing>,
395        req: &RoutingRequirements,
396    ) -> Option<&'a ModelPricing> {
397        providers.into_iter().min_by(|a, b| {
398            let ca = a.estimate_cost(req.estimated_input_tokens, req.estimated_output_tokens);
399            let cb = b.estimate_cost(req.estimated_input_tokens, req.estimated_output_tokens);
400            ca.partial_cmp(&cb).unwrap_or(std::cmp::Ordering::Equal)
401        })
402    }
403
404    /// Return the highest-quality provider from `providers`.
405    fn best_quality<'a>(&self, providers: Vec<&'a ModelPricing>) -> Option<&'a ModelPricing> {
406        providers.into_iter().max_by(|a, b| {
407            a.quality_score
408                .partial_cmp(&b.quality_score)
409                .unwrap_or(std::cmp::Ordering::Equal)
410        })
411    }
412}
413
414// ---------------------------------------------------------------------------
415// Unit tests
416// ---------------------------------------------------------------------------
417
418#[cfg(test)]
419mod tests {
420    use super::*;
421    use crate::provider_health::ProviderHealthMonitor;
422
423    fn make_router() -> SmartRouter {
424        let health = ProviderHealthMonitor::new(20);
425        SmartRouter::new(SmartRouter::default_pricing(), health)
426    }
427
428    #[test]
429    fn test_estimate_cost_zero_tokens() {
430        let p = ModelPricing {
431            provider_id: "test".into(),
432            model: "m".into(),
433            input_per_million: 3.0,
434            output_per_million: 15.0,
435            quality_score: 0.9,
436            typical_p95_ms: 2000,
437        };
438        assert_eq!(p.estimate_cost(0, 0), 0.0);
439    }
440
441    #[test]
442    fn test_estimate_cost_one_million_tokens() {
443        let p = ModelPricing {
444            provider_id: "test".into(),
445            model: "m".into(),
446            input_per_million: 3.0,
447            output_per_million: 15.0,
448            quality_score: 0.9,
449            typical_p95_ms: 2000,
450        };
451        // 1M input + 1M output = $3 + $15 = $18
452        assert!((p.estimate_cost(1_000_000, 1_000_000) - 18.0).abs() < 1e-9);
453    }
454
455    #[test]
456    fn test_default_pricing_not_empty() {
457        let pricing = SmartRouter::default_pricing();
458        assert!(!pricing.is_empty(), "default pricing must include at least one model");
459        // All quality scores must be in [0, 1].
460        for p in &pricing {
461            assert!(
462                p.quality_score >= 0.0 && p.quality_score <= 1.0,
463                "quality_score out of range for {}",
464                p.model
465            );
466            assert!(p.input_per_million >= 0.0);
467            assert!(p.output_per_million >= 0.0);
468        }
469    }
470
471    #[tokio::test]
472    async fn test_route_no_constraints_returns_cheapest() {
473        let router = make_router();
474        let req = RoutingRequirements::default();
475        let decision = router.route(&req).await.unwrap();
476        // llama-3.1-8b is the cheapest in the default table at $0.06/$0.06.
477        assert_eq!(decision.provider_id, "meta-llama");
478    }
479
480    #[tokio::test]
481    async fn test_route_prefer_quality() {
482        let router = make_router();
483        let req = RoutingRequirements {
484            prefer_quality: true,
485            ..Default::default()
486        };
487        let decision = router.route(&req).await.unwrap();
488        // claude-3-5-sonnet has quality_score 0.92, highest in table.
489        assert_eq!(decision.provider_id, "anthropic");
490    }
491
492    #[tokio::test]
493    async fn test_route_latency_filter_excludes_slow_models() {
494        let router = make_router();
495        let req = RoutingRequirements {
496            // Only models with typical_p95_ms <= 2000 pass.
497            max_latency_ms: Some(2000),
498            ..Default::default()
499        };
500        let decision = router.route(&req).await.unwrap();
501        // Only llama (1200), haiku (1800), gpt4o-mini (2000) qualify.
502        // Cheapest of those three is llama.
503        assert_eq!(decision.provider_id, "meta-llama");
504    }
505
506    #[tokio::test]
507    async fn test_route_quality_floor_excludes_low_quality() {
508        let router = make_router();
509        let req = RoutingRequirements {
510            // Only quality >= 0.88 qualifies: sonnet (0.92) and gpt-4o (0.90).
511            min_quality: Some(0.88),
512            ..Default::default()
513        };
514        let decision = router.route(&req).await.unwrap();
515        // gpt-4o costs $2.5 input / $10 output; sonnet costs $3/$15.
516        // gpt-4o is cheaper for the default 512+256 token estimate.
517        assert_eq!(decision.provider_id, "openai-gpt4o");
518    }
519
520    #[tokio::test]
521    async fn test_route_returns_none_when_no_candidates() {
522        let router = make_router();
523        let req = RoutingRequirements {
524            // Impossible: require quality >= 0.99 AND latency <= 100 ms.
525            min_quality: Some(0.99),
526            max_latency_ms: Some(100),
527            ..Default::default()
528        };
529        assert!(router.route(&req).await.is_none());
530    }
531
532    #[tokio::test]
533    async fn test_fallbacks_ordered_cheapest_first() {
534        let router = make_router();
535        let req = RoutingRequirements::default();
536        let decision = router.route(&req).await.unwrap();
537        // Verify fallbacks exist and that no fallback is the same as primary.
538        for fb in &decision.fallbacks {
539            assert_ne!(fb, &decision.provider_id);
540        }
541        // Verify they are in ascending cost order.
542        let pricing = SmartRouter::default_pricing();
543        let cost_of = |id: &str| {
544            pricing
545                .iter()
546                .find(|p| p.provider_id == id)
547                .map(|p| p.estimate_cost(req.estimated_input_tokens, req.estimated_output_tokens))
548                .unwrap_or(f64::MAX)
549        };
550        for window in decision.fallbacks.windows(2) {
551            assert!(
552                cost_of(&window[0]) <= cost_of(&window[1]),
553                "fallbacks should be cheapest-first"
554            );
555        }
556    }
557
558    #[tokio::test]
559    async fn test_load_pricing_toml_success() {
560        let mut router = make_router();
561        let toml_str = r#"
562[[model]]
563provider_id = "test-provider"
564model = "test-model-v1"
565input_per_million = 1.0
566output_per_million = 2.0
567quality_score = 0.8
568typical_p95_ms = 1500
569"#;
570        let count = router.load_pricing_toml(toml_str).unwrap();
571        assert_eq!(count, 1);
572        let req = RoutingRequirements::default();
573        let decision = router.route(&req).await.unwrap();
574        assert_eq!(decision.provider_id, "test-provider");
575    }
576
577    #[test]
578    fn test_load_pricing_toml_invalid_returns_err() {
579        let mut router = make_router();
580        let bad_toml = "this is not valid toml ][[[";
581        assert!(router.load_pricing_toml(bad_toml).is_err());
582        // Existing pricing should be unchanged.
583        assert!(!router.pricing.is_empty());
584    }
585}