tokio_prompt_orchestrator/
cost_estimator.rs1use std::collections::HashMap;
8
9#[derive(Debug, Clone, Copy, PartialEq, Eq)]
16pub enum TaskType {
17 Translation,
19 Summarization,
21 CodeGeneration,
23 QuestionAnswering,
25 Classification,
27 Creative,
29}
30
31#[derive(Debug, Clone)]
37pub struct TokenEstimate {
38 pub input_tokens: usize,
40 pub output_tokens: usize,
42 pub model: String,
44}
45
46#[derive(Debug, Clone)]
48pub struct CostEstimate {
49 pub estimated_cost_usd: f64,
51 pub token_estimate: TokenEstimate,
53 pub confidence: f64,
56 pub within_budget: bool,
58}
59
60#[derive(Debug, Clone)]
66pub struct BudgetConfig {
67 pub daily_limit_usd: f64,
69 pub per_request_limit_usd: f64,
71 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
85struct 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
135const 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
144pub struct CostEstimator {
153 prices: HashMap<&'static str, ModelPrice>,
154}
155
156impl CostEstimator {
157 pub fn new() -> Self {
159 Self {
160 prices: pricing_table(),
161 }
162 }
163
164 pub fn estimate_tokens(text: &str) -> usize {
170 let chars = text.chars().count();
171 (chars / 4).max(1)
172 }
173
174 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 fn compute_cost(&self, input_tokens: usize, output_tokens: usize, model: &str) -> (f64, f64) {
192 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) }
200 None => {
201 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 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 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 pub fn cheapest_model_for_budget(&self, prompt: &str, budget_usd: f64) -> Option<String> {
263 let input_tokens = Self::estimate_tokens(prompt);
264 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); }
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 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}