tokio_prompt_orchestrator/routing/
arbitrage.rs1use std::collections::HashMap;
60use std::sync::atomic::{AtomicU64, Ordering};
61use std::sync::Mutex;
62use std::time::Duration;
63
64#[derive(Debug, Clone)]
68pub struct ProviderProfile {
69 pub name: String,
71 pub cost_per_1k_input_tokens: f64,
73 pub cost_per_1k_output_tokens: f64,
75 pub priority: u32,
78}
79
80impl ProviderProfile {
81 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#[derive(Debug)]
90struct ProviderState {
91 profile: ProviderProfile,
92 latency_ring: Mutex<LatencyRing>,
94 requests: AtomicU64,
96 errors: AtomicU64,
98 input_tokens: AtomicU64,
100 output_tokens: AtomicU64,
102}
103
104const 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 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 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#[derive(Debug)]
164pub struct ArbitrageEngine {
165 providers: Mutex<HashMap<String, ProviderState>>,
166 selections: AtomicU64,
168 sla_misses: AtomicU64,
170}
171
172impl Default for ArbitrageEngine {
173 fn default() -> Self {
174 Self::new()
175 }
176}
177
178impl ArbitrageEngine {
179 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 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 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 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 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 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 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 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 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 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 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 pub fn total_selections(&self) -> u64 {
392 self.selections.load(Ordering::Relaxed)
393 }
394
395 pub fn total_sla_misses(&self) -> u64 {
397 self.sla_misses.load(Ordering::Relaxed)
398 }
399}
400
401#[derive(Debug, Clone)]
403pub struct ProviderSnapshot {
404 pub name: String,
406 pub cost_per_1k_input_tokens: f64,
408 pub cost_per_1k_output_tokens: f64,
410 pub priority: u32,
412 pub p50_latency_ms: Option<f64>,
414 pub p95_latency_ms: Option<f64>,
416 pub total_requests: u64,
418 pub error_rate: f64,
420 pub input_tokens: u64,
422 pub output_tokens: u64,
424}
425
426#[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)); engine.register(make_profile("openai", 0.005, 0.020)); 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 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, });
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 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}