Skip to main content

tokio_prompt_orchestrator/routing/
arbitrage.rs

1//! # Provider Arbitrage Engine
2//!
3//! Routes inference requests to the **cheapest provider that historically meets
4//! a caller-specified latency SLA**.
5//!
6//! ## Problem
7//!
8//! Multiple providers (Anthropic, OpenAI, vLLM, llama.cpp) have different
9//! per-token prices and different latency characteristics.  A naïve router
10//! always uses the cheapest provider, but that provider may be slower or less
11//! reliable.  A naïve latency router always uses the fastest provider, but
12//! wastes money.
13//!
14//! The arbitrage engine tracks a rolling P95 latency window for each registered
15//! provider and, given a latency budget, selects whichever eligible provider
16//! (i.e., whose P95 is below the budget) has the lowest per-token cost.
17//!
18//! ## Guarantees
19//!
20//! - Thread-safe: all hot-path state uses atomics; the latency ring buffer uses
21//!   a `Mutex` but the lock is held for microseconds.
22//! - Non-blocking: `select_provider` never performs I/O.
23//! - Graceful degradation: if no provider meets the SLA, the one with the
24//!   lowest P95 latency is returned (best-effort).
25//!
26//! ## Example
27//!
28//! ```rust
29//! use tokio_prompt_orchestrator::routing::arbitrage::{ArbitrageEngine, ProviderProfile};
30//! use std::time::Duration;
31//!
32//! let engine = ArbitrageEngine::new();
33//!
34//! engine.register(ProviderProfile {
35//!     name: "anthropic".to_string(),
36//!     cost_per_1k_input_tokens: 0.003,
37//!     cost_per_1k_output_tokens: 0.015,
38//!     priority: 0,
39//! });
40//! engine.register(ProviderProfile {
41//!     name: "openai".to_string(),
42//!     cost_per_1k_input_tokens: 0.005,
43//!     cost_per_1k_output_tokens: 0.015,
44//!     priority: 0,
45//! });
46//!
47//! // Record observed latencies
48//! engine.record_latency("anthropic", Duration::from_millis(320));
49//! engine.record_latency("openai", Duration::from_millis(180));
50//!
51//! // With a 500ms SLA budget, pick the cheapest provider that historically
52//! // completes within 500ms.
53//! let sla = Duration::from_millis(500);
54//! let winner = engine.select_provider(Some(sla));
55//! // anthropic is cheaper AND within the 500ms SLA → selected
56//! assert_eq!(winner.as_ref().map(|p| p.name.as_str()), Some("anthropic"));
57//! ```
58
59use std::collections::HashMap;
60use std::sync::atomic::{AtomicU64, Ordering};
61use std::sync::Mutex;
62use std::time::Duration;
63
64// ── Types ───────────────────────────────────────────────────────────────────
65
66/// Static profile for a single inference provider.
67#[derive(Debug, Clone)]
68pub struct ProviderProfile {
69    /// Unique name identifying this provider (e.g., `"anthropic"`, `"openai"`).
70    pub name: String,
71    /// Cost in USD per 1 000 input tokens.
72    pub cost_per_1k_input_tokens: f64,
73    /// Cost in USD per 1 000 output tokens.
74    pub cost_per_1k_output_tokens: f64,
75    /// Lower priority value = preferred when cost is equal.
76    /// Use to break ties (e.g., prefer local vLLM over cloud at the same cost).
77    pub priority: u32,
78}
79
80impl ProviderProfile {
81    /// Estimated total cost in USD for a request with the given token counts.
82    pub fn estimate_cost(&self, input_tokens: u64, output_tokens: u64) -> f64 {
83        (input_tokens as f64 / 1000.0) * self.cost_per_1k_input_tokens
84            + (output_tokens as f64 / 1000.0) * self.cost_per_1k_output_tokens
85    }
86}
87
88/// Runtime state tracked per provider.
89#[derive(Debug)]
90struct ProviderState {
91    profile: ProviderProfile,
92    /// Rolling ring buffer of observed latencies in microseconds.
93    latency_ring: Mutex<LatencyRing>,
94    /// Total requests routed to this provider.
95    requests: AtomicU64,
96    /// Total errors from this provider.
97    errors: AtomicU64,
98    /// Total input tokens sent to this provider.
99    input_tokens: AtomicU64,
100    /// Total output tokens received from this provider.
101    output_tokens: AtomicU64,
102}
103
104/// Fixed-capacity ring buffer for latency samples (microseconds).
105const RING_CAPACITY: usize = 128;
106
107#[derive(Debug)]
108struct LatencyRing {
109    samples: [u64; RING_CAPACITY],
110    head: usize,
111    len: usize,
112}
113
114impl LatencyRing {
115    fn new() -> Self {
116        Self {
117            samples: [0u64; RING_CAPACITY],
118            head: 0,
119            len: 0,
120        }
121    }
122
123    fn push(&mut self, micros: u64) {
124        self.samples[self.head] = micros;
125        self.head = (self.head + 1) % RING_CAPACITY;
126        if self.len < RING_CAPACITY {
127            self.len += 1;
128        }
129    }
130
131    /// Returns the P95 latency in microseconds.  Returns `u64::MAX` if no
132    /// samples have been recorded yet.
133    fn p95_micros(&self) -> u64 {
134        if self.len == 0 {
135            return u64::MAX;
136        }
137        let mut sorted: Vec<u64> = self.samples[..self.len].to_vec();
138        sorted.sort_unstable();
139        let idx = ((self.len as f64 * 0.95) as usize).min(self.len - 1);
140        sorted[idx]
141    }
142
143    /// Returns the P50 (median) latency in microseconds.
144    fn p50_micros(&self) -> u64 {
145        if self.len == 0 {
146            return u64::MAX;
147        }
148        let mut sorted: Vec<u64> = self.samples[..self.len].to_vec();
149        sorted.sort_unstable();
150        sorted[self.len / 2]
151    }
152}
153
154// ── Engine ──────────────────────────────────────────────────────────────────
155
156/// The provider arbitrage engine.
157///
158/// Thread-safe; wrap in `Arc` to share across pipeline stages.
159///
160/// # Panics
161///
162/// No method panics.
163#[derive(Debug)]
164pub struct ArbitrageEngine {
165    providers: Mutex<HashMap<String, ProviderState>>,
166    /// Total `select_provider` calls.
167    selections: AtomicU64,
168    /// Calls where all providers exceeded the SLA (fell back to fastest).
169    sla_misses: AtomicU64,
170}
171
172impl Default for ArbitrageEngine {
173    fn default() -> Self {
174        Self::new()
175    }
176}
177
178impl ArbitrageEngine {
179    /// Create a new, empty arbitrage engine.
180    pub fn new() -> Self {
181        Self {
182            providers: Mutex::new(HashMap::new()),
183            selections: AtomicU64::new(0),
184            sla_misses: AtomicU64::new(0),
185        }
186    }
187
188    /// Register a provider.  If a provider with the same name already exists,
189    /// its profile is updated (existing runtime state is preserved).
190    ///
191    /// # Panics
192    ///
193    /// Does not panic.
194    pub fn register(&self, profile: ProviderProfile) {
195        let mut map = self.providers.lock().unwrap_or_else(|e| e.into_inner());
196        map.entry(profile.name.clone())
197            .and_modify(|s| s.profile = profile.clone())
198            .or_insert_with(|| ProviderState {
199                profile,
200                latency_ring: Mutex::new(LatencyRing::new()),
201                requests: AtomicU64::new(0),
202                errors: AtomicU64::new(0),
203                input_tokens: AtomicU64::new(0),
204                output_tokens: AtomicU64::new(0),
205            });
206    }
207
208    /// Record an observed end-to-end latency for a provider.
209    ///
210    /// Call this after every inference call completes (success or failure).
211    ///
212    /// # Panics
213    ///
214    /// Does not panic.
215    pub fn record_latency(&self, provider: &str, latency: Duration) {
216        let map = self.providers.lock().unwrap_or_else(|e| e.into_inner());
217        if let Some(state) = map.get(provider) {
218            let micros = latency.as_micros().min(u64::MAX as u128) as u64;
219            let mut ring = state
220                .latency_ring
221                .lock()
222                .unwrap_or_else(|e| e.into_inner());
223            ring.push(micros);
224        }
225    }
226
227    /// Record a successful inference call for accounting purposes.
228    ///
229    /// # Panics
230    ///
231    /// Does not panic.
232    pub fn record_success(
233        &self,
234        provider: &str,
235        input_tokens: u64,
236        output_tokens: u64,
237        latency: Duration,
238    ) {
239        self.record_latency(provider, latency);
240        let map = self.providers.lock().unwrap_or_else(|e| e.into_inner());
241        if let Some(state) = map.get(provider) {
242            state.requests.fetch_add(1, Ordering::Relaxed);
243            state.input_tokens.fetch_add(input_tokens, Ordering::Relaxed);
244            state
245                .output_tokens
246                .fetch_add(output_tokens, Ordering::Relaxed);
247        }
248    }
249
250    /// Record a failed inference call for a provider.
251    ///
252    /// # Panics
253    ///
254    /// Does not panic.
255    pub fn record_error(&self, provider: &str, latency: Duration) {
256        self.record_latency(provider, latency);
257        let map = self.providers.lock().unwrap_or_else(|e| e.into_inner());
258        if let Some(state) = map.get(provider) {
259            state.requests.fetch_add(1, Ordering::Relaxed);
260            state.errors.fetch_add(1, Ordering::Relaxed);
261        }
262    }
263
264    /// Select the cheapest provider whose P95 latency is at or below
265    /// `sla_budget`.
266    ///
267    /// ## Selection algorithm
268    ///
269    /// 1. Filter providers whose P95 latency ≤ `sla_budget`.
270    /// 2. Among those, pick the one with the lowest cost (input + output rate).
271    ///    Ties are broken by `priority` (lower = preferred), then name
272    ///    (alphabetical, for determinism).
273    /// 3. If no provider meets the SLA (or `sla_budget` is `None`), fall back
274    ///    to the provider with the lowest P95 latency.
275    /// 4. If no providers are registered, returns `None`.
276    ///
277    /// # Returns
278    ///
279    /// A cloned [`ProviderProfile`] for the selected provider, or `None` if
280    /// no providers are registered.
281    ///
282    /// # Panics
283    ///
284    /// Does not panic.
285    pub fn select_provider(&self, sla_budget: Option<Duration>) -> Option<ProviderProfile> {
286        self.selections.fetch_add(1, Ordering::Relaxed);
287
288        let map = self.providers.lock().unwrap_or_else(|e| e.into_inner());
289        if map.is_empty() {
290            return None;
291        }
292
293        let sla_micros: Option<u64> = sla_budget.map(|d| {
294            d.as_micros().min(u64::MAX as u128) as u64
295        });
296
297        // Compute P95 for each provider
298        let candidates: Vec<(&ProviderState, u64)> = map
299            .values()
300            .map(|state| {
301                let p95 = {
302                    let ring = state.latency_ring.lock().unwrap_or_else(|e| e.into_inner());
303                    ring.p95_micros()
304                };
305                (state, p95)
306            })
307            .collect();
308
309        // Try to find providers within SLA
310        let within_sla: Vec<_> = if let Some(budget) = sla_micros {
311            candidates
312                .iter()
313                .filter(|(_, p95)| *p95 <= budget)
314                .collect()
315        } else {
316            candidates.iter().collect()
317        };
318
319        let selected = if within_sla.is_empty() {
320            // SLA miss — fall back to fastest
321            self.sla_misses.fetch_add(1, Ordering::Relaxed);
322            candidates
323                .iter()
324                .min_by(|(_, a_p95), (_, b_p95)| a_p95.cmp(b_p95))
325                .map(|(state, _)| state)
326        } else {
327            // Among SLA-meeting providers, pick cheapest (then priority, then name)
328            within_sla
329                .iter()
330                .min_by(|(a_state, _), (b_state, _)| {
331                    let a_cost = a_state.profile.cost_per_1k_input_tokens
332                        + a_state.profile.cost_per_1k_output_tokens;
333                    let b_cost = b_state.profile.cost_per_1k_input_tokens
334                        + b_state.profile.cost_per_1k_output_tokens;
335                    a_cost
336                        .partial_cmp(&b_cost)
337                        .unwrap_or(std::cmp::Ordering::Equal)
338                        .then(a_state.profile.priority.cmp(&b_state.profile.priority))
339                        .then(a_state.profile.name.cmp(&b_state.profile.name))
340                })
341                .map(|(state, _)| state)
342        };
343
344        selected.map(|s| s.profile.clone())
345    }
346
347    /// Return a point-in-time snapshot of all registered providers.
348    ///
349    /// # Panics
350    ///
351    /// Does not panic.
352    pub fn snapshot(&self) -> Vec<ProviderSnapshot> {
353        let map = self.providers.lock().unwrap_or_else(|e| e.into_inner());
354        map.values()
355            .map(|state| {
356                let (p50, p95) = {
357                    let ring = state.latency_ring.lock().unwrap_or_else(|e| e.into_inner());
358                    (ring.p50_micros(), ring.p95_micros())
359                };
360                let requests = state.requests.load(Ordering::Relaxed);
361                let errors = state.errors.load(Ordering::Relaxed);
362                ProviderSnapshot {
363                    name: state.profile.name.clone(),
364                    cost_per_1k_input_tokens: state.profile.cost_per_1k_input_tokens,
365                    cost_per_1k_output_tokens: state.profile.cost_per_1k_output_tokens,
366                    priority: state.profile.priority,
367                    p50_latency_ms: if p50 == u64::MAX {
368                        None
369                    } else {
370                        Some(p50 as f64 / 1000.0)
371                    },
372                    p95_latency_ms: if p95 == u64::MAX {
373                        None
374                    } else {
375                        Some(p95 as f64 / 1000.0)
376                    },
377                    total_requests: requests,
378                    error_rate: if requests > 0 {
379                        errors as f64 / requests as f64
380                    } else {
381                        0.0
382                    },
383                    input_tokens: state.input_tokens.load(Ordering::Relaxed),
384                    output_tokens: state.output_tokens.load(Ordering::Relaxed),
385                }
386            })
387            .collect()
388    }
389
390    /// Total `select_provider` calls since engine creation.
391    pub fn total_selections(&self) -> u64 {
392        self.selections.load(Ordering::Relaxed)
393    }
394
395    /// Total selections that fell back to fastest (SLA could not be met).
396    pub fn total_sla_misses(&self) -> u64 {
397        self.sla_misses.load(Ordering::Relaxed)
398    }
399}
400
401/// Point-in-time snapshot of a single provider's state.
402#[derive(Debug, Clone)]
403pub struct ProviderSnapshot {
404    /// Provider name.
405    pub name: String,
406    /// Input token cost per 1K tokens (USD).
407    pub cost_per_1k_input_tokens: f64,
408    /// Output token cost per 1K tokens (USD).
409    pub cost_per_1k_output_tokens: f64,
410    /// Priority for tie-breaking (lower = preferred).
411    pub priority: u32,
412    /// P50 (median) latency in milliseconds, or `None` if no samples yet.
413    pub p50_latency_ms: Option<f64>,
414    /// P95 latency in milliseconds, or `None` if no samples yet.
415    pub p95_latency_ms: Option<f64>,
416    /// Total requests routed to this provider.
417    pub total_requests: u64,
418    /// Error rate (0.0 – 1.0).
419    pub error_rate: f64,
420    /// Total input tokens sent.
421    pub input_tokens: u64,
422    /// Total output tokens received.
423    pub output_tokens: u64,
424}
425
426// ── Tests ───────────────────────────────────────────────────────────────────
427
428#[cfg(test)]
429mod tests {
430    use super::*;
431
432    fn make_profile(name: &str, input_cost: f64, output_cost: f64) -> ProviderProfile {
433        ProviderProfile {
434            name: name.to_string(),
435            cost_per_1k_input_tokens: input_cost,
436            cost_per_1k_output_tokens: output_cost,
437            priority: 0,
438        }
439    }
440
441    #[test]
442    fn empty_engine_returns_none() {
443        let engine = ArbitrageEngine::new();
444        assert!(engine.select_provider(None).is_none());
445    }
446
447    #[test]
448    fn single_provider_always_selected() {
449        let engine = ArbitrageEngine::new();
450        engine.register(make_profile("anthropic", 0.003, 0.015));
451        engine.record_latency("anthropic", Duration::from_millis(300));
452        let result = engine.select_provider(Some(Duration::from_secs(1)));
453        assert_eq!(result.map(|p| p.name), Some("anthropic".to_string()));
454    }
455
456    #[test]
457    fn cheaper_provider_wins_when_both_meet_sla() {
458        let engine = ArbitrageEngine::new();
459        engine.register(make_profile("anthropic", 0.003, 0.015)); // cheaper
460        engine.register(make_profile("openai", 0.005, 0.020));    // more expensive
461        // Both have P95 = 200ms which is below 500ms SLA
462        engine.record_latency("anthropic", Duration::from_millis(200));
463        engine.record_latency("openai", Duration::from_millis(150));
464        let result = engine.select_provider(Some(Duration::from_millis(500)));
465        assert_eq!(result.map(|p| p.name), Some("anthropic".to_string()));
466    }
467
468    #[test]
469    fn sla_miss_falls_back_to_fastest() {
470        let engine = ArbitrageEngine::new();
471        engine.register(make_profile("slow_cheap", 0.001, 0.003));
472        engine.register(make_profile("fast_expensive", 0.010, 0.030));
473        for _ in 0..10 {
474            engine.record_latency("slow_cheap", Duration::from_millis(800));
475            engine.record_latency("fast_expensive", Duration::from_millis(150));
476        }
477        // SLA = 100ms; both providers exceed it, so falls back to fastest (fast_expensive)
478        let result = engine.select_provider(Some(Duration::from_millis(100)));
479        assert_eq!(result.map(|p| p.name), Some("fast_expensive".to_string()));
480        assert_eq!(engine.total_sla_misses(), 1);
481    }
482
483    #[test]
484    fn no_sla_budget_picks_cheapest() {
485        let engine = ArbitrageEngine::new();
486        engine.register(make_profile("cheap", 0.001, 0.002));
487        engine.register(make_profile("expensive", 0.010, 0.020));
488        let result = engine.select_provider(None);
489        assert_eq!(result.map(|p| p.name), Some("cheap".to_string()));
490    }
491
492    #[test]
493    fn priority_breaks_cost_tie() {
494        let engine = ArbitrageEngine::new();
495        engine.register(ProviderProfile {
496            name: "local_vllm".to_string(),
497            cost_per_1k_input_tokens: 0.0,
498            cost_per_1k_output_tokens: 0.0,
499            priority: 0, // preferred
500        });
501        engine.register(ProviderProfile {
502            name: "other_local".to_string(),
503            cost_per_1k_input_tokens: 0.0,
504            cost_per_1k_output_tokens: 0.0,
505            priority: 1,
506        });
507        let result = engine.select_provider(None);
508        assert_eq!(result.map(|p| p.name), Some("local_vllm".to_string()));
509    }
510
511    #[test]
512    fn record_success_updates_accounting() {
513        let engine = ArbitrageEngine::new();
514        engine.register(make_profile("a", 0.003, 0.015));
515        engine.record_success("a", 500, 200, Duration::from_millis(250));
516        let snap: Vec<_> = engine.snapshot();
517        let a = snap.iter().find(|s| s.name == "a").unwrap();
518        assert_eq!(a.total_requests, 1);
519        assert!((a.error_rate).abs() < f64::EPSILON);
520        assert_eq!(a.input_tokens, 500);
521        assert_eq!(a.output_tokens, 200);
522    }
523
524    #[test]
525    fn record_error_increments_error_rate() {
526        let engine = ArbitrageEngine::new();
527        engine.register(make_profile("b", 0.003, 0.015));
528        engine.record_success("b", 100, 50, Duration::from_millis(100));
529        engine.record_error("b", Duration::from_millis(100));
530        let snap: Vec<_> = engine.snapshot();
531        let b = snap.iter().find(|s| s.name == "b").unwrap();
532        assert_eq!(b.total_requests, 2);
533        assert!((b.error_rate - 0.5).abs() < 1e-6);
534    }
535
536    #[test]
537    fn provider_estimate_cost() {
538        let p = make_profile("x", 0.003, 0.015);
539        let cost = p.estimate_cost(1000, 500);
540        // 1000/1000 * 0.003 + 500/1000 * 0.015 = 0.003 + 0.0075 = 0.0105
541        assert!((cost - 0.0105).abs() < 1e-9);
542    }
543
544    #[test]
545    fn snapshot_no_latency_returns_none_for_p95() {
546        let engine = ArbitrageEngine::new();
547        engine.register(make_profile("c", 0.003, 0.015));
548        let snap = engine.snapshot();
549        let c = snap.iter().find(|s| s.name == "c").unwrap();
550        assert!(c.p95_latency_ms.is_none());
551        assert!(c.p50_latency_ms.is_none());
552    }
553}