Skip to main content

tokio_prompt_orchestrator/
cost_estimator.rs

1//! Pre-flight cost estimation for LLM requests.
2//!
3//! Provides token-count heuristics, per-model pricing, and budget-checking
4//! before a request is dispatched so that callers can decide whether to
5//! proceed, choose a cheaper model, or reject the request entirely.
6
7use std::collections::HashMap;
8
9// ---------------------------------------------------------------------------
10// Task type
11// ---------------------------------------------------------------------------
12
13/// Characterises the kind of work being requested, which influences the
14/// expected output token count.
15#[derive(Debug, Clone, Copy, PartialEq, Eq)]
16pub enum TaskType {
17    /// Translate text from one language to another.
18    Translation,
19    /// Condense a longer document into a short summary.
20    Summarization,
21    /// Generate source code from a natural-language description.
22    CodeGeneration,
23    /// Answer a factual or reasoning question.
24    QuestionAnswering,
25    /// Categorise input into one of several labels.
26    Classification,
27    /// Open-ended creative writing.
28    Creative,
29}
30
31// ---------------------------------------------------------------------------
32// Token / cost value types
33// ---------------------------------------------------------------------------
34
35/// Input/output token count estimate for a request.
36#[derive(Debug, Clone)]
37pub struct TokenEstimate {
38    /// Estimated number of input (prompt) tokens.
39    pub input_tokens: usize,
40    /// Estimated number of output (completion) tokens.
41    pub output_tokens: usize,
42    /// Model the estimate is scoped to.
43    pub model: String,
44}
45
46/// Full cost estimate including budget compliance.
47#[derive(Debug, Clone)]
48pub struct CostEstimate {
49    /// Estimated cost in US dollars.
50    pub estimated_cost_usd: f64,
51    /// Underlying token counts.
52    pub token_estimate: TokenEstimate,
53    /// Confidence in the estimate (0.0–1.0).  Lower when the model is
54    /// unknown or the text is very short.
55    pub confidence: f64,
56    /// Whether the estimate fits within the per-request budget limit.
57    pub within_budget: bool,
58}
59
60// ---------------------------------------------------------------------------
61// Budget configuration
62// ---------------------------------------------------------------------------
63
64/// Spending limits used to check whether a request is within budget.
65#[derive(Debug, Clone)]
66pub struct BudgetConfig {
67    /// Maximum cumulative spend per calendar day (USD).
68    pub daily_limit_usd: f64,
69    /// Maximum cost for a single request (USD).
70    pub per_request_limit_usd: f64,
71    /// Maximum cumulative spend per calendar month (USD).
72    pub monthly_limit_usd: f64,
73}
74
75impl Default for BudgetConfig {
76    fn default() -> Self {
77        Self {
78            daily_limit_usd: 10.0,
79            per_request_limit_usd: 0.50,
80            monthly_limit_usd: 200.0,
81        }
82    }
83}
84
85// ---------------------------------------------------------------------------
86// Internal pricing table
87// ---------------------------------------------------------------------------
88
89/// Per-1 M token pricing (input, output) in USD.
90struct ModelPrice {
91    input_per_1m: f64,
92    output_per_1m: f64,
93}
94
95fn pricing_table() -> HashMap<&'static str, ModelPrice> {
96    let mut m = HashMap::new();
97    m.insert(
98        "gpt-4o",
99        ModelPrice {
100            input_per_1m: 5.00,
101            output_per_1m: 15.00,
102        },
103    );
104    m.insert(
105        "gpt-4o-mini",
106        ModelPrice {
107            input_per_1m: 0.15,
108            output_per_1m: 0.60,
109        },
110    );
111    m.insert(
112        "claude-3-5-sonnet",
113        ModelPrice {
114            input_per_1m: 3.00,
115            output_per_1m: 15.00,
116        },
117    );
118    m.insert(
119        "claude-3-haiku",
120        ModelPrice {
121            input_per_1m: 0.25,
122            output_per_1m: 1.25,
123        },
124    );
125    m.insert(
126        "gemini-1.5-pro",
127        ModelPrice {
128            input_per_1m: 3.50,
129            output_per_1m: 10.50,
130        },
131    );
132    m
133}
134
135/// Canonical model names in ascending cost order (input price).
136const MODELS_BY_COST: &[&str] = &[
137    "gpt-4o-mini",
138    "claude-3-haiku",
139    "gemini-1.5-pro",
140    "claude-3-5-sonnet",
141    "gpt-4o",
142];
143
144// ---------------------------------------------------------------------------
145// CostEstimator
146// ---------------------------------------------------------------------------
147
148/// Pre-flight cost estimation engine.
149///
150/// Thread-safe — all methods take `&self` and the pricing table is built once
151/// on construction.
152pub struct CostEstimator {
153    prices: HashMap<&'static str, ModelPrice>,
154}
155
156impl CostEstimator {
157    /// Create a new estimator with the built-in pricing table.
158    pub fn new() -> Self {
159        Self {
160            prices: pricing_table(),
161        }
162    }
163
164    // -----------------------------------------------------------------------
165    // Token counting heuristics
166    // -----------------------------------------------------------------------
167
168    /// Estimate the number of tokens in `text` using the ~4 chars/token rule.
169    pub fn estimate_tokens(text: &str) -> usize {
170        let chars = text.chars().count();
171        (chars / 4).max(1)
172    }
173
174    /// Estimate the expected number of output tokens given the input count and
175    /// the nature of the task.
176    pub fn estimate_output_tokens(input_tokens: usize, task_type: TaskType) -> usize {
177        match task_type {
178            TaskType::Classification => 10.min(input_tokens),
179            TaskType::Summarization => (input_tokens / 4).max(50),
180            TaskType::QuestionAnswering => (input_tokens / 3).clamp(30, 500),
181            TaskType::Translation => input_tokens,
182            TaskType::CodeGeneration => (input_tokens * 2).max(100),
183            TaskType::Creative => (input_tokens * 3).max(200),
184        }
185    }
186
187    // -----------------------------------------------------------------------
188    // Cost computation helpers
189    // -----------------------------------------------------------------------
190
191    fn compute_cost(&self, input_tokens: usize, output_tokens: usize, model: &str) -> (f64, f64) {
192        // Look up by exact string match against the static key.
193        let price = self.prices.iter().find(|(k, _)| **k == model);
194        match price {
195            Some((_, p)) => {
196                let cost = (input_tokens as f64 / 1_000_000.0) * p.input_per_1m
197                    + (output_tokens as f64 / 1_000_000.0) * p.output_per_1m;
198                (cost, 0.90) // confidence high for known model
199            }
200            None => {
201                // Unknown model: fall back to gpt-4o pricing, lower confidence.
202                // `prices` is public, so gpt-4o may have been removed; estimate zero then.
203                let cost = self.prices.get("gpt-4o").map_or(0.0, |fallback| {
204                    (input_tokens as f64 / 1_000_000.0) * fallback.input_per_1m
205                        + (output_tokens as f64 / 1_000_000.0) * fallback.output_per_1m
206                });
207                (cost, 0.40)
208            }
209        }
210    }
211
212    // -----------------------------------------------------------------------
213    // Public API
214    // -----------------------------------------------------------------------
215
216    /// Estimate the cost of a single prompt/model/task combination.
217    pub fn estimate_cost(
218        &self,
219        prompt: &str,
220        model: &str,
221        task_type: TaskType,
222        budget: &BudgetConfig,
223    ) -> CostEstimate {
224        let input_tokens = Self::estimate_tokens(prompt);
225        let output_tokens = Self::estimate_output_tokens(input_tokens, task_type);
226        let (cost, confidence) = self.compute_cost(input_tokens, output_tokens, model);
227
228        CostEstimate {
229            estimated_cost_usd: cost,
230            token_estimate: TokenEstimate {
231                input_tokens,
232                output_tokens,
233                model: model.to_string(),
234            },
235            confidence,
236            within_budget: cost <= budget.per_request_limit_usd,
237        }
238    }
239
240    /// Estimate costs for a batch of prompts with a shared model and task type.
241    ///
242    /// Each prompt is estimated independently; no budget is enforced here since
243    /// the caller typically aggregates totals and applies limits externally.
244    pub fn batch_estimate(
245        &self,
246        prompts: &[&str],
247        model: &str,
248        task_type: TaskType,
249    ) -> Vec<CostEstimate> {
250        let budget = BudgetConfig {
251            per_request_limit_usd: f64::MAX,
252            ..BudgetConfig::default()
253        };
254        prompts
255            .iter()
256            .map(|p| self.estimate_cost(p, model, task_type, &budget))
257            .collect()
258    }
259
260    /// Return the name of the cheapest known model whose per-request cost for
261    /// `prompt` fits within `budget_usd`, or `None` if none qualifies.
262    pub fn cheapest_model_for_budget(&self, prompt: &str, budget_usd: f64) -> Option<String> {
263        let input_tokens = Self::estimate_tokens(prompt);
264        // Use QuestionAnswering as a neutral default for model selection.
265        let output_tokens = Self::estimate_output_tokens(input_tokens, TaskType::QuestionAnswering);
266
267        for model_name in MODELS_BY_COST {
268            let (cost, _) = self.compute_cost(input_tokens, output_tokens, model_name);
269            if cost <= budget_usd {
270                return Some(model_name.to_string());
271            }
272        }
273        None
274    }
275}
276
277impl Default for CostEstimator {
278    fn default() -> Self {
279        Self::new()
280    }
281}
282
283#[cfg(test)]
284mod tests {
285    use super::*;
286
287    #[test]
288    fn test_estimate_tokens() {
289        assert_eq!(CostEstimator::estimate_tokens("hello world"), 2);
290        assert_eq!(CostEstimator::estimate_tokens(""), 1); // minimum 1
291    }
292
293    #[test]
294    fn test_estimate_cost_known_model() {
295        let est = CostEstimator::new();
296        let budget = BudgetConfig::default();
297        let result = est.estimate_cost("Some prompt text here.", "gpt-4o-mini", TaskType::Classification, &budget);
298        assert!(result.estimated_cost_usd >= 0.0);
299        assert!(result.confidence > 0.5);
300    }
301
302    #[test]
303    fn test_estimate_cost_unknown_model() {
304        let est = CostEstimator::new();
305        let budget = BudgetConfig::default();
306        let result = est.estimate_cost("Some text", "unknown-model-xyz", TaskType::Summarization, &budget);
307        assert!(result.confidence < 0.5);
308    }
309
310    #[test]
311    fn test_batch_estimate_length() {
312        let est = CostEstimator::new();
313        let prompts = vec!["hello", "world", "foo"];
314        let results = est.batch_estimate(&prompts, "gpt-4o", TaskType::QuestionAnswering);
315        assert_eq!(results.len(), 3);
316    }
317
318    #[test]
319    fn test_cheapest_model() {
320        let est = CostEstimator::new();
321        // Any small prompt should have a cheapest model.
322        let model = est.cheapest_model_for_budget("hello", 1.0);
323        assert!(model.is_some());
324    }
325
326    #[test]
327    fn test_within_budget_flag() {
328        let est = CostEstimator::new();
329        let tight = BudgetConfig {
330            per_request_limit_usd: 0.0,
331            ..BudgetConfig::default()
332        };
333        let result = est.estimate_cost("test", "gpt-4o", TaskType::CodeGeneration, &tight);
334        assert!(!result.within_budget);
335    }
336}