lean_ctx/core/context_kernel/
usage_normalizer.rs1use std::cmp::Reverse;
4use std::collections::HashMap;
5use std::sync::{Mutex, MutexGuard, OnceLock};
6
7use super::token_envelope::TokenEnvelope;
8
9#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
11pub struct SessionUsage {
12 pub session_id: String,
14 pub per_model: HashMap<String, ModelUsageEntry>,
16 pub per_provider: HashMap<String, ProviderUsageEntry>,
18 pub total_requests: usize,
20 pub total_tokens: usize,
22 pub total_saved: usize,
24 pub total_cost_usd: f64,
26}
27
28#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
30pub struct ModelUsageEntry {
31 pub requests: usize,
33 pub input_tokens: usize,
35 pub output_tokens: usize,
37 pub cache_read_tokens: usize,
39 pub tokens_saved: usize,
41 pub cost_usd: f64,
43}
44
45#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
47pub struct ProviderUsageEntry {
48 pub requests: usize,
50 pub total_tokens: usize,
52 pub tokens_saved: usize,
54 pub cost_usd: f64,
56}
57
58#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
60pub struct CompressionOverview {
61 pub total_tokens: usize,
63 pub total_saved: usize,
65 pub avg_compression_ratio: f64,
67 pub best_model: Option<String>,
69 pub best_ratio: f64,
71}
72
73static NORMALIZER: OnceLock<Mutex<SessionUsage>> = OnceLock::new();
74
75fn usage_guard() -> MutexGuard<'static, SessionUsage> {
76 NORMALIZER
77 .get_or_init(|| Mutex::new(SessionUsage::default()))
78 .lock()
79 .unwrap_or_else(std::sync::PoisonError::into_inner)
80}
81
82fn compression_ratio(input_tokens: usize, tokens_saved: usize) -> f64 {
83 let original = input_tokens.saturating_add(tokens_saved);
84 if original == 0 {
85 0.0
86 } else {
87 tokens_saved as f64 / original as f64
88 }
89}
90
91pub fn record_envelope(envelope: &TokenEnvelope) {
93 let total_tokens = envelope.total_tokens();
94 let cost = envelope.cost_usd.unwrap_or_default();
95 let mut usage = usage_guard();
96
97 usage.total_requests = usage.total_requests.saturating_add(1);
98 usage.total_tokens = usage.total_tokens.saturating_add(total_tokens);
99 usage.total_saved = usage.total_saved.saturating_add(envelope.tokens_saved);
100 usage.total_cost_usd += cost;
101
102 let model = usage.per_model.entry(envelope.model.clone()).or_default();
103 model.requests = model.requests.saturating_add(1);
104 model.input_tokens = model.input_tokens.saturating_add(envelope.input_tokens);
105 model.output_tokens = model.output_tokens.saturating_add(envelope.output_tokens);
106 model.cache_read_tokens = model
107 .cache_read_tokens
108 .saturating_add(envelope.cache_read_tokens);
109 model.tokens_saved = model.tokens_saved.saturating_add(envelope.tokens_saved);
110 model.cost_usd += cost;
111
112 let provider = usage
113 .per_provider
114 .entry(format!("{:?}", envelope.provider))
115 .or_default();
116 provider.requests = provider.requests.saturating_add(1);
117 provider.total_tokens = provider.total_tokens.saturating_add(total_tokens);
118 provider.tokens_saved = provider.tokens_saved.saturating_add(envelope.tokens_saved);
119 provider.cost_usd += cost;
120}
121
122pub fn session_usage() -> SessionUsage {
124 usage_guard().clone()
125}
126
127pub fn model_breakdown() -> Vec<(String, ModelUsageEntry)> {
129 let mut entries = usage_guard()
130 .per_model
131 .clone()
132 .into_iter()
133 .collect::<Vec<_>>();
134 entries.sort_unstable_by_key(|(_, entry)| {
135 Reverse(entry.input_tokens.saturating_add(entry.output_tokens))
136 });
137 entries
138}
139
140#[must_use]
142pub fn provider_breakdown() -> Vec<(String, ProviderUsageEntry)> {
143 let mut entries = usage_guard()
144 .per_provider
145 .clone()
146 .into_iter()
147 .collect::<Vec<_>>();
148 entries.sort_unstable_by_key(|(_, entry)| Reverse(entry.total_tokens));
149 entries
150}
151
152#[must_use]
154pub fn compression_overview() -> CompressionOverview {
155 let usage = usage_guard();
156 let total_input = usage
157 .per_model
158 .values()
159 .map(|entry| entry.input_tokens)
160 .sum();
161 let best = usage.per_model.iter().fold(None, |best, (model, entry)| {
162 let ratio = compression_ratio(entry.input_tokens, entry.tokens_saved);
163 match best {
164 Some((_, best_ratio)) if ratio <= best_ratio => best,
165 _ => Some((model.clone(), ratio)),
166 }
167 });
168 let (best_model, best_ratio) = best.map_or((None, 0.0), |(model, ratio)| (Some(model), ratio));
169
170 CompressionOverview {
171 total_tokens: usage.total_tokens,
172 total_saved: usage.total_saved,
173 avg_compression_ratio: compression_ratio(total_input, usage.total_saved),
174 best_model,
175 best_ratio,
176 }
177}
178
179pub fn reset_usage() {
181 *usage_guard() = SessionUsage::default();
182}
183
184#[cfg(test)]
185mod tests {
186 use super::{compression_overview as overview, model_breakdown, provider_breakdown};
187 use super::{record_envelope, reset_usage, session_usage};
188 use crate::core::context_kernel::token_envelope::{ProviderKind, TokenEnvelope};
189 use std::sync::{Mutex, MutexGuard};
190 static TEST_LOCK: Mutex<()> = Mutex::new(());
191
192 fn setup() -> MutexGuard<'static, ()> {
193 let guard = TEST_LOCK
194 .lock()
195 .unwrap_or_else(std::sync::PoisonError::into_inner);
196 reset_usage();
197 record_envelope(&envelope("low", ProviderKind::Anthropic, 20, 0));
198 record_envelope(&envelope("high", ProviderKind::Gemini, 60, 60));
199 record_envelope(&envelope("low", ProviderKind::OpenAi, 20, 0));
200 guard
201 }
202 fn envelope(model: &str, provider: ProviderKind, input: usize, saved: usize) -> TokenEnvelope {
203 TokenEnvelope {
204 model: model.to_owned(),
205 provider,
206 input_tokens: input,
207 tokens_saved: saved,
208 ..TokenEnvelope::default()
209 }
210 }
211 #[test]
212 fn record_updates_totals() {
213 let _guard = setup();
214 let usage = session_usage();
215 assert_eq!(
216 (usage.total_requests, usage.total_tokens, usage.total_saved),
217 (3, 100, 60)
218 );
219 }
220 #[test]
221 fn per_model_breakdown() {
222 let _guard = setup();
223 let entries = model_breakdown();
224 assert_eq!((entries.len(), entries[0].0.as_str()), (2, "high"));
225 }
226 #[test]
227 fn per_provider_breakdown() {
228 let _guard = setup();
229 let entries = provider_breakdown();
230 assert_eq!((entries.len(), entries[0].0.as_str()), (3, "Gemini"));
231 assert!(entries[0].1.total_tokens > entries[1].1.total_tokens);
232 }
233 #[test]
234 fn compression_overview() {
235 let _guard = setup();
236 let summary = overview();
237 assert!((summary.avg_compression_ratio - 0.375).abs() < f64::EPSILON);
238 assert_eq!(summary.best_model.as_deref(), Some("high"));
239 assert!((summary.best_ratio - 0.5).abs() < f64::EPSILON);
240 }
241 #[test]
242 fn reset_clears() {
243 let _guard = setup();
244 reset_usage();
245 let usage = session_usage();
246 assert_eq!((usage.total_requests, usage.total_tokens), (0, 0));
247 assert!(usage.per_model.is_empty());
248 }
249}