1#[derive(Debug, Clone)]
26pub struct ModelPricing {
27 pub provider_id: String,
29 pub model: String,
31 pub input_per_million: f64,
33 pub output_per_million: f64,
35 pub quality_score: f32,
37 pub typical_p95_ms: u64,
39}
40
41impl ModelPricing {
42 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#[derive(Debug, Clone)]
61pub struct RoutingRequirements {
62 pub max_latency_ms: Option<u64>,
64 pub min_quality: Option<f32>,
66 pub max_cost_usd: Option<f64>,
68 pub estimated_input_tokens: usize,
70 pub estimated_output_tokens: usize,
72 pub prefer_quality: bool,
75}
76
77impl Default for RoutingRequirements {
78 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#[derive(Debug, Clone)]
102pub struct RoutingDecision {
103 pub provider_id: String,
105 pub model: String,
107 pub estimated_cost_usd: f64,
109 pub reason: String,
111 pub fallbacks: Vec<String>,
114}
115
116pub struct SmartRouter {
123 pricing: Vec<ModelPricing>,
124 health: crate::provider_health::ProviderHealthMonitor,
125}
126
127impl SmartRouter {
128 pub fn new(
134 pricing: Vec<ModelPricing>,
135 health: crate::provider_health::ProviderHealthMonitor,
136 ) -> Self {
137 Self { pricing, health }
138 }
139
140 pub async fn route(&self, req: &RoutingRequirements) -> Option<RoutingDecision> {
147 let healthy = self.filter_by_health(&self.pricing).await;
149 let within_latency = self.filter_by_latency(healthy, req.max_latency_ms);
151 let quality_ok = self.filter_by_quality(within_latency, req.min_quality);
153 let cost_ok = self.filter_by_cost(quality_ok, req);
155
156 if cost_ok.is_empty() {
157 return None;
158 }
159
160 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 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 pub fn default_pricing() -> Vec<ModelPricing> {
220 vec![
221 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 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 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 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 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 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 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 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 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 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#[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 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 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 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 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 max_latency_ms: Some(2000),
498 ..Default::default()
499 };
500 let decision = router.route(&req).await.unwrap();
501 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 min_quality: Some(0.88),
512 ..Default::default()
513 };
514 let decision = router.route(&req).await.unwrap();
515 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 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 for fb in &decision.fallbacks {
539 assert_ne!(fb, &decision.provider_id);
540 }
541 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 assert!(!router.pricing.is_empty());
584 }
585}