switchyard_libsy/algorithms/
cache_aware.rs1use std::{sync::Arc, time::Duration};
7
8use switchyard_protocol::{Category, Request};
9
10use crate::{Algorithm, Driver, LibsyError, Result, RoutingOutcome};
11
12pub struct CacheAware {
14 costs: Vec<f64>,
15 max_age: Duration,
16}
17
18impl CacheAware {
19 pub fn new(costs: Vec<f64>, max_age: Duration) -> Result<Self> {
21 if costs.is_empty()
22 || costs.iter().any(|c| !c.is_finite() || *c <= 0.0)
23 || max_age.is_zero()
24 {
25 return Err(LibsyError::AlgorithmError {
26 message: "cache_aware requires positive costs and signal age".into(),
27 });
28 }
29 Ok(Self { costs, max_age })
30 }
31}
32
33#[async_trait::async_trait]
34impl Algorithm for CacheAware {
35 fn name(&self) -> &str {
36 "cache_aware"
37 }
38
39 async fn route(self: Arc<Self>, driver: Driver, request: Request) -> Result<RoutingOutcome> {
40 let models = driver.models_for(&Category::Any);
41 let fallback = models.first().ok_or(LibsyError::NoTargets)?;
42 if self.costs.len() != models.len() {
43 return Err(LibsyError::AlgorithmError {
44 message: "cache_aware costs must match runtime targets".into(),
45 });
46 }
47 let observations = request.metadata.as_ref().map(|m| &m.serving_observations);
48 let scores: Option<Vec<f64>> = models
50 .iter()
51 .zip(&self.costs)
52 .map(|(model, cost)| {
53 let signal = observations?.get(model)?;
54 if signal.received_at.elapsed() > self.max_age {
55 return None;
56 }
57 let queued = signal.active_prefill_tokens?;
58 let score = (signal.effective_prefill_tokens as f64 + queued as f64) * cost;
59 score.is_finite().then_some(score)
60 })
61 .collect();
62 let selected = scores
63 .as_ref()
64 .and_then(|scores| {
65 scores
66 .iter()
67 .enumerate()
68 .min_by(|a, b| a.1.total_cmp(b.1))
69 .map(|(index, _)| &models[index])
70 })
71 .unwrap_or(fallback);
72 tracing::info!(model = %selected, used_signals = scores.is_some(), "cache-aware routing decision");
73 Ok(RoutingOutcome::route_to(
74 selected.clone(),
75 Vec::new(),
76 request,
77 ))
78 }
79}