Skip to main content

tokio_prompt_orchestrator/
cost_optimizer.rs

1#![allow(dead_code)]
2//! # Module: Cost Optimizer
3//!
4//! ## Responsibility
5//! Analyses historical request cost patterns and produces actionable
6//! optimisation suggestions. Can optionally apply suggestions automatically
7//! when `auto_optimize = true`.
8//!
9//! ## Detections
10//! 1. **Cache candidates** — prompts whose responses are highly similar across
11//!    invocations (response fingerprint collision rate above threshold).
12//! 2. **Model overuse** — expensive models being used for structurally simple
13//!    tasks (short prompt + short response).
14//!
15//! ## Guarantees
16//! - Thread-safe (`Arc<CostOptimizer>` shareable across tasks)
17//! - Bounded memory: observation windows capped at `window_size`
18//! - Non-blocking: all methods are O(window_size), no I/O
19//! - Advisory only unless `auto_optimize = true`
20
21use std::collections::HashMap;
22use std::sync::{Arc, Mutex};
23use std::time::Instant;
24
25use serde::{Deserialize, Serialize};
26use thiserror::Error;
27use tracing::info;
28
29// ─── Errors ──────────────────────────────────────────────────────────────────
30
31/// Errors produced by the cost optimizer.
32#[derive(Debug, Error)]
33pub enum CostOptimizerError {
34    /// The internal lock was poisoned.
35    #[error("internal lock poisoned")]
36    LockPoisoned,
37    /// The supplied configuration is invalid.
38    #[error("invalid configuration: {0}")]
39    InvalidConfig(String),
40}
41
42// ─── Configuration ────────────────────────────────────────────────────────────
43
44/// Cost optimizer configuration.
45#[derive(Debug, Clone, Serialize, Deserialize)]
46pub struct CostOptimizerConfig {
47    /// When `true`, suggestions are applied automatically.
48    pub auto_optimize: bool,
49    /// Number of recent observations to retain per prompt intent.
50    pub window_size: usize,
51    /// Fraction of observations that must have matching fingerprints before
52    /// a prompt is flagged as a cache candidate (0.0–1.0).
53    pub cache_candidate_threshold: f64,
54    /// Maximum prompt+response token count (combined chars / 4) to classify
55    /// a task as "simple". Tasks below this threshold on an expensive model
56    /// are flagged for model downgrade.
57    pub simple_task_token_threshold: usize,
58    /// Model tier definitions: (model_name, cost_per_1k_tokens, tier).
59    pub model_tiers: Vec<ModelTierEntry>,
60}
61
62impl Default for CostOptimizerConfig {
63    fn default() -> Self {
64        Self {
65            auto_optimize: false,
66            window_size: 200,
67            cache_candidate_threshold: 0.70,
68            simple_task_token_threshold: 512,
69            model_tiers: vec![
70                ModelTierEntry {
71                    model: "gpt-4o".to_string(),
72                    cost_per_1k_tokens: 0.005,
73                    tier: ModelTier::Expensive,
74                },
75                ModelTierEntry {
76                    model: "gpt-4o-mini".to_string(),
77                    cost_per_1k_tokens: 0.00015,
78                    tier: ModelTier::Cheap,
79                },
80                ModelTierEntry {
81                    model: "claude-3-5-sonnet".to_string(),
82                    cost_per_1k_tokens: 0.003,
83                    tier: ModelTier::Expensive,
84                },
85                ModelTierEntry {
86                    model: "claude-3-haiku".to_string(),
87                    cost_per_1k_tokens: 0.00025,
88                    tier: ModelTier::Cheap,
89                },
90            ],
91        }
92    }
93}
94
95/// Tier classification for a model.
96#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
97pub enum ModelTier {
98    /// High-capability, high-cost model.
99    Expensive,
100    /// Lower-cost model suitable for simpler tasks.
101    Cheap,
102}
103
104/// An entry in the model tier table.
105#[derive(Debug, Clone, Serialize, Deserialize)]
106pub struct ModelTierEntry {
107    /// Model identifier as used in routing (e.g. `"gpt-4o"`).
108    pub model: String,
109    /// Cost per 1,000 tokens in USD.
110    pub cost_per_1k_tokens: f64,
111    /// Whether this model is expensive or cheap.
112    pub tier: ModelTier,
113}
114
115// ─── Observation ─────────────────────────────────────────────────────────────
116
117/// A single cost observation recorded for one inference call.
118#[derive(Debug, Clone)]
119pub struct CostObservation {
120    /// First 64 chars of the prompt (intent key).
121    pub intent: String,
122    /// Model used.
123    pub model: String,
124    /// Approximate token count (prompt_chars + response_chars) / 4.
125    pub tokens_approx: usize,
126    /// Actual cost in USD (computed from tokens × model rate).
127    pub cost_usd: f64,
128    /// A simple fingerprint of the response (first 32 chars, lowercased).
129    pub response_fingerprint: String,
130    /// When this observation was recorded.
131    pub recorded_at: Instant,
132}
133
134fn compute_fingerprint(response: &str) -> String {
135    response
136        .chars()
137        .take(32)
138        .collect::<String>()
139        .to_lowercase()
140        .trim()
141        .to_string()
142}
143
144fn derive_intent(prompt: &str) -> String {
145    prompt.chars().take(64).collect::<String>().to_lowercase()
146}
147
148// ─── Suggestions ─────────────────────────────────────────────────────────────
149
150/// The kind of cost optimisation suggested.
151#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
152pub enum SuggestionKind {
153    /// This prompt should be cached — responses are highly repetitive.
154    EnableCaching,
155    /// A cheaper model would be sufficient for this task.
156    DowngradeModel,
157}
158
159/// A single optimisation suggestion with computed ROI.
160#[derive(Debug, Clone, Serialize, Deserialize)]
161pub struct OptimizationSuggestion {
162    /// The prompt intent this suggestion applies to.
163    pub intent: String,
164    /// What kind of optimisation is recommended.
165    pub kind: SuggestionKind,
166    /// Human-readable description of the suggestion.
167    pub description: String,
168    /// Estimated monthly cost savings in USD.
169    pub estimated_monthly_savings_usd: f64,
170    /// The model currently in use (for DowngradeModel suggestions).
171    pub current_model: Option<String>,
172    /// The recommended cheaper model (for DowngradeModel suggestions).
173    pub suggested_model: Option<String>,
174    /// Whether this suggestion was automatically applied.
175    pub auto_applied: bool,
176}
177
178// ─── Auto-apply actions ──────────────────────────────────────────────────────
179
180/// An automatically applied optimisation action.
181#[derive(Debug, Clone)]
182pub struct AutoApplyAction {
183    /// The suggestion that triggered this action.
184    pub suggestion: OptimizationSuggestion,
185    /// When the action was applied.
186    pub applied_at: Instant,
187}
188
189// ─── Per-intent stats ────────────────────────────────────────────────────────
190
191#[derive(Debug)]
192struct IntentWindow {
193    observations: std::collections::VecDeque<CostObservation>,
194    max_size: usize,
195}
196
197impl IntentWindow {
198    fn new(max_size: usize) -> Self {
199        Self {
200            observations: std::collections::VecDeque::new(),
201            max_size: max_size.max(1),
202        }
203    }
204
205    fn push(&mut self, obs: CostObservation) {
206        if self.observations.len() >= self.max_size {
207            self.observations.pop_front();
208        }
209        self.observations.push_back(obs);
210    }
211
212    fn len(&self) -> usize {
213        self.observations.len()
214    }
215
216    /// Fraction of observations whose fingerprint matches the most common one.
217    fn fingerprint_collision_rate(&self) -> f64 {
218        if self.observations.is_empty() {
219            return 0.0;
220        }
221        let mut counts: HashMap<&str, usize> = HashMap::new();
222        for obs in &self.observations {
223            *counts.entry(obs.response_fingerprint.as_str()).or_insert(0) += 1;
224        }
225        let max_count = counts.values().copied().max().unwrap_or(0);
226        max_count as f64 / self.observations.len() as f64
227    }
228
229    /// Most recently used model.
230    fn latest_model(&self) -> Option<&str> {
231        self.observations.back().map(|o| o.model.as_str())
232    }
233
234    /// Average token count.
235    fn avg_tokens(&self) -> f64 {
236        if self.observations.is_empty() {
237            return 0.0;
238        }
239        let sum: usize = self.observations.iter().map(|o| o.tokens_approx).sum();
240        sum as f64 / self.observations.len() as f64
241    }
242
243    /// Average cost per call.
244    fn avg_cost(&self) -> f64 {
245        if self.observations.is_empty() {
246            return 0.0;
247        }
248        let sum: f64 = self.observations.iter().map(|o| o.cost_usd).sum();
249        sum / self.observations.len() as f64
250    }
251
252    /// Estimated calls per month (extrapolated from observation timestamps).
253    fn estimated_calls_per_month(&self) -> f64 {
254        if self.observations.len() < 2 {
255            return 0.0;
256        }
257        let first = self.observations.front().map(|o| o.recorded_at);
258        let last = self.observations.back().map(|o| o.recorded_at);
259        if let (Some(first), Some(last)) = (first, last) {
260            let elapsed = last.duration_since(first);
261            if elapsed.as_secs() == 0 {
262                return 0.0;
263            }
264            let calls_per_sec =
265                (self.observations.len() - 1) as f64 / elapsed.as_secs_f64();
266            calls_per_sec * 86_400.0 * 30.0
267        } else {
268            0.0
269        }
270    }
271}
272
273// ─── Cost optimizer ──────────────────────────────────────────────────────────
274
275/// Analyses cost patterns and generates (or applies) optimisation suggestions.
276pub struct CostOptimizer {
277    config: CostOptimizerConfig,
278    inner: Mutex<CostOptimizerInner>,
279}
280
281#[derive(Debug)]
282struct CostOptimizerInner {
283    windows: HashMap<String, IntentWindow>, // intent → window
284    auto_applied: Vec<AutoApplyAction>,
285    overridden_models: HashMap<String, String>, // intent → forced cheaper model
286}
287
288impl CostOptimizer {
289    /// Create a new optimizer with the given configuration.
290    ///
291    /// # Errors
292    /// Returns [`CostOptimizerError::InvalidConfig`] if `window_size` is zero
293    /// or `cache_candidate_threshold` is outside [0, 1].
294    pub fn new(config: CostOptimizerConfig) -> Result<Arc<Self>, CostOptimizerError> {
295        if config.window_size == 0 {
296            return Err(CostOptimizerError::InvalidConfig(
297                "window_size must be > 0".to_string(),
298            ));
299        }
300        if !(0.0..=1.0).contains(&config.cache_candidate_threshold) {
301            return Err(CostOptimizerError::InvalidConfig(
302                "cache_candidate_threshold must be in [0, 1]".to_string(),
303            ));
304        }
305        Ok(Arc::new(Self {
306            config,
307            inner: Mutex::new(CostOptimizerInner {
308                windows: HashMap::new(),
309                auto_applied: Vec::new(),
310                overridden_models: HashMap::new(),
311            }),
312        }))
313    }
314
315    /// Record a completed inference call.
316    pub fn record(
317        &self,
318        prompt: &str,
319        model: &str,
320        response: &str,
321    ) -> Result<(), CostOptimizerError> {
322        let intent = derive_intent(prompt);
323        let tokens_approx = (prompt.len() + response.len()) / 4;
324        let cost_usd = self.compute_cost(model, tokens_approx);
325        let obs = CostObservation {
326            intent: intent.clone(),
327            model: model.to_string(),
328            tokens_approx,
329            cost_usd,
330            response_fingerprint: compute_fingerprint(response),
331            recorded_at: Instant::now(),
332        };
333
334        let mut inner = self.inner.lock().map_err(|_| CostOptimizerError::LockPoisoned)?;
335        let window = inner
336            .windows
337            .entry(intent)
338            .or_insert_with(|| IntentWindow::new(self.config.window_size));
339        window.push(obs);
340        Ok(())
341    }
342
343    /// Return a list of optimisation suggestions based on observed patterns.
344    pub fn suggestions(&self) -> Result<Vec<OptimizationSuggestion>, CostOptimizerError> {
345        let inner = self.inner.lock().map_err(|_| CostOptimizerError::LockPoisoned)?;
346        let mut suggestions = Vec::new();
347
348        for (intent, window) in &inner.windows {
349            // Require at least 5 observations before generating suggestions
350            if window.len() < 5 {
351                continue;
352            }
353
354            // Cache candidate detection
355            let collision_rate = window.fingerprint_collision_rate();
356            if collision_rate >= self.config.cache_candidate_threshold {
357                let avg_cost = window.avg_cost();
358                let calls_per_month = window.estimated_calls_per_month();
359                let savings = avg_cost * calls_per_month * collision_rate;
360                suggestions.push(OptimizationSuggestion {
361                    intent: intent.clone(),
362                    kind: SuggestionKind::EnableCaching,
363                    description: format!(
364                        "Prompt '{intent:.40}…' produces identical responses {:.0}% of the time. \
365                         Enable result caching to save ~${savings:.2}/month.",
366                        collision_rate * 100.0
367                    ),
368                    estimated_monthly_savings_usd: savings,
369                    current_model: window.latest_model().map(str::to_string),
370                    suggested_model: None,
371                    auto_applied: false,
372                });
373            }
374
375            // Model downgrade detection
376            if let Some(model) = window.latest_model() {
377                if let Some(tier_entry) = self.config.model_tiers.iter().find(|t| t.model == model) {
378                    if tier_entry.tier == ModelTier::Expensive {
379                        let avg_tokens = window.avg_tokens();
380                        if avg_tokens < self.config.simple_task_token_threshold as f64 {
381                            // Find cheapest alternative
382                            if let Some(cheap) = self
383                                .config
384                                .model_tiers
385                                .iter()
386                                .filter(|t| t.tier == ModelTier::Cheap)
387                                .min_by(|a, b| {
388                                    a.cost_per_1k_tokens
389                                        .partial_cmp(&b.cost_per_1k_tokens)
390                                        .unwrap_or(std::cmp::Ordering::Equal)
391                                })
392                            {
393                                let current_cost_per_call =
394                                    avg_tokens / 1000.0 * tier_entry.cost_per_1k_tokens;
395                                let cheap_cost_per_call =
396                                    avg_tokens / 1000.0 * cheap.cost_per_1k_tokens;
397                                let calls_per_month = window.estimated_calls_per_month();
398                                let savings = (current_cost_per_call - cheap_cost_per_call)
399                                    * calls_per_month;
400                                suggestions.push(OptimizationSuggestion {
401                                    intent: intent.clone(),
402                                    kind: SuggestionKind::DowngradeModel,
403                                    description: format!(
404                                        "Simple task (avg {avg_tokens:.0} tokens) is using expensive model \
405                                         '{model}'. Switch to '{}' to save ~${savings:.2}/month.",
406                                        cheap.model
407                                    ),
408                                    estimated_monthly_savings_usd: savings,
409                                    current_model: Some(model.to_string()),
410                                    suggested_model: Some(cheap.model.clone()),
411                                    auto_applied: false,
412                                });
413                            }
414                        }
415                    }
416                }
417            }
418        }
419
420        // Sort by highest savings first
421        suggestions.sort_by(|a, b| {
422            b.estimated_monthly_savings_usd
423                .partial_cmp(&a.estimated_monthly_savings_usd)
424                .unwrap_or(std::cmp::Ordering::Equal)
425        });
426
427        Ok(suggestions)
428    }
429
430    /// Apply all pending suggestions (only when `auto_optimize = true`).
431    ///
432    /// Returns the list of actions taken.
433    pub fn auto_apply(&self) -> Result<Vec<AutoApplyAction>, CostOptimizerError> {
434        if !self.config.auto_optimize {
435            return Ok(vec![]);
436        }
437        let pending = self.suggestions()?;
438        let mut inner = self.inner.lock().map_err(|_| CostOptimizerError::LockPoisoned)?;
439        let mut applied = Vec::new();
440
441        for mut suggestion in pending {
442            match suggestion.kind {
443                SuggestionKind::EnableCaching => {
444                    // In a real system this would toggle a cache flag; here we log it
445                    info!(
446                        intent = %suggestion.intent,
447                        savings_usd = suggestion.estimated_monthly_savings_usd,
448                        "auto-optimizer: enabling caching for intent"
449                    );
450                    suggestion.auto_applied = true;
451                    let action = AutoApplyAction {
452                        suggestion: suggestion.clone(),
453                        applied_at: Instant::now(),
454                    };
455                    inner.auto_applied.push(action.clone());
456                    applied.push(action);
457                }
458                SuggestionKind::DowngradeModel => {
459                    if let Some(ref cheap_model) = suggestion.suggested_model {
460                        info!(
461                            intent = %suggestion.intent,
462                            from = suggestion.current_model.as_deref().unwrap_or("?"),
463                            to = %cheap_model,
464                            savings_usd = suggestion.estimated_monthly_savings_usd,
465                            "auto-optimizer: overriding model for intent"
466                        );
467                        inner
468                            .overridden_models
469                            .insert(suggestion.intent.clone(), cheap_model.clone());
470                        suggestion.auto_applied = true;
471                        let action = AutoApplyAction {
472                            suggestion: suggestion.clone(),
473                            applied_at: Instant::now(),
474                        };
475                        inner.auto_applied.push(action.clone());
476                        applied.push(action);
477                    }
478                }
479            }
480        }
481
482        Ok(applied)
483    }
484
485    /// Return the model override for an intent (applied by auto-optimizer), if any.
486    #[must_use]
487    pub fn model_override(&self, prompt: &str) -> Option<String> {
488        let intent = derive_intent(prompt);
489        let Ok(inner) = self.inner.lock() else {
490            return None;
491        };
492        inner.overridden_models.get(&intent).cloned()
493    }
494
495    /// Return total cost recorded across all observations.
496    pub fn total_cost_usd(&self) -> Result<f64, CostOptimizerError> {
497        let inner = self.inner.lock().map_err(|_| CostOptimizerError::LockPoisoned)?;
498        let total: f64 = inner
499            .windows
500            .values()
501            .flat_map(|w| w.observations.iter().map(|o| o.cost_usd))
502            .sum();
503        Ok(total)
504    }
505
506    /// Number of distinct intents tracked.
507    pub fn intent_count(&self) -> Result<usize, CostOptimizerError> {
508        let inner = self.inner.lock().map_err(|_| CostOptimizerError::LockPoisoned)?;
509        Ok(inner.windows.len())
510    }
511
512    fn compute_cost(&self, model: &str, tokens_approx: usize) -> f64 {
513        let rate = self
514            .config
515            .model_tiers
516            .iter()
517            .find(|t| t.model == model)
518            .map(|t| t.cost_per_1k_tokens)
519            .unwrap_or(0.002); // fallback mid-tier rate
520        tokens_approx as f64 / 1000.0 * rate
521    }
522}
523
524// ─── Tests ───────────────────────────────────────────────────────────────────
525
526#[cfg(test)]
527mod tests {
528    use super::*;
529
530    fn make_optimizer(auto: bool) -> Arc<CostOptimizer> {
531        CostOptimizer::new(CostOptimizerConfig {
532            auto_optimize: auto,
533            window_size: 20,
534            cache_candidate_threshold: 0.6,
535            simple_task_token_threshold: 200,
536            ..Default::default()
537        })
538        .expect("optimizer should construct")
539    }
540
541    #[test]
542    fn records_observations() {
543        let opt = make_optimizer(false);
544        opt.record("Tell me a joke", "gpt-4o", "Why did the chicken...").unwrap();
545        opt.record("Tell me a joke", "gpt-4o", "Why did the chicken...").unwrap();
546        assert_eq!(opt.intent_count().unwrap(), 1);
547        assert!(opt.total_cost_usd().unwrap() > 0.0);
548    }
549
550    #[test]
551    fn detects_cache_candidate() {
552        let opt = make_optimizer(false);
553        for _ in 0..8 {
554            opt.record("Repeat after me: hello", "gpt-4o", "hello").unwrap();
555        }
556        let suggestions = opt.suggestions().unwrap();
557        let cache_sugg = suggestions
558            .iter()
559            .find(|s| s.kind == SuggestionKind::EnableCaching);
560        assert!(cache_sugg.is_some(), "should detect cache candidate");
561    }
562
563    #[test]
564    fn detects_model_downgrade() {
565        let opt = make_optimizer(false);
566        // Very short prompt+response → simple task
567        for _ in 0..8 {
568            opt.record("Hi", "gpt-4o", "Hello").unwrap();
569        }
570        let suggestions = opt.suggestions().unwrap();
571        let downgrade = suggestions
572            .iter()
573            .find(|s| s.kind == SuggestionKind::DowngradeModel);
574        assert!(downgrade.is_some(), "should suggest model downgrade");
575        assert!(downgrade.unwrap().suggested_model.is_some());
576    }
577
578    #[test]
579    fn auto_apply_writes_override() {
580        let opt = make_optimizer(true);
581        for _ in 0..8 {
582            opt.record("Hi", "gpt-4o", "Hello").unwrap();
583        }
584        let applied = opt.auto_apply().unwrap();
585        assert!(!applied.is_empty());
586        // At least one downgrade action should exist
587        let has_override = opt.model_override("Hi").is_some();
588        assert!(has_override, "model override should be set after auto-apply");
589    }
590
591    #[test]
592    fn invalid_config_rejected() {
593        let result = CostOptimizer::new(CostOptimizerConfig {
594            window_size: 0,
595            ..Default::default()
596        });
597        assert!(result.is_err());
598    }
599}