Skip to main content

switchyard_libsy/algorithms/
cache_aware.rs

1// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4//! Choose among operator-approved models using request-specific serving observations.
5
6use std::{sync::Arc, time::Duration};
7
8use switchyard_protocol::{Category, Request};
9
10use crate::{Algorithm, Driver, LibsyError, Result, RoutingOutcome};
11
12/// A relative prefill-work policy. Its scores are not latency predictions.
13pub struct CacheAware {
14    costs: Vec<f64>,
15    max_age: Duration,
16}
17
18impl CacheAware {
19    /// Costs correspond to runtime targets in order; all must be finite and positive.
20    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        // Incomplete observations must not make an unobserved model look worse.
49        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}