Skip to main content

lean_ctx/core/context_kernel/
usage_normalizer.rs

1//! Session-level aggregation for canonical token envelopes.
2
3use std::cmp::Reverse;
4use std::collections::HashMap;
5use std::sync::{Mutex, MutexGuard, OnceLock};
6
7use super::token_envelope::TokenEnvelope;
8
9/// Aggregated usage for one session.
10#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
11pub struct SessionUsage {
12    /// Session identifier.
13    pub session_id: String,
14    /// Per-model token totals.
15    pub per_model: HashMap<String, ModelUsageEntry>,
16    /// Per-provider token totals.
17    pub per_provider: HashMap<String, ProviderUsageEntry>,
18    /// Total envelopes recorded.
19    pub total_requests: usize,
20    /// Total tokens across all requests.
21    pub total_tokens: usize,
22    /// Total tokens saved by compression.
23    pub total_saved: usize,
24    /// Total estimated cost in USD.
25    pub total_cost_usd: f64,
26}
27
28/// Aggregated usage for one model.
29#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
30pub struct ModelUsageEntry {
31    /// Number of requests served by the model.
32    pub requests: usize,
33    /// Input tokens delivered to the model.
34    pub input_tokens: usize,
35    /// Output tokens generated by the model.
36    pub output_tokens: usize,
37    /// Input tokens served from provider cache.
38    pub cache_read_tokens: usize,
39    /// Tokens removed by compression.
40    pub tokens_saved: usize,
41    /// Estimated cost in USD.
42    pub cost_usd: f64,
43}
44
45/// Aggregated usage for one provider.
46#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
47pub struct ProviderUsageEntry {
48    /// Number of requests served by the provider.
49    pub requests: usize,
50    /// Input, output, and reasoning tokens combined.
51    pub total_tokens: usize,
52    /// Tokens removed by compression.
53    pub tokens_saved: usize,
54    /// Estimated cost in USD.
55    pub cost_usd: f64,
56}
57
58/// Compression totals and the model with the highest savings ratio.
59#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
60pub struct CompressionOverview {
61    /// Input, output, and reasoning tokens combined.
62    pub total_tokens: usize,
63    /// Tokens removed by compression.
64    pub total_saved: usize,
65    /// Saved tokens divided by original input tokens.
66    pub avg_compression_ratio: f64,
67    /// Model with the highest aggregate compression ratio.
68    pub best_model: Option<String>,
69    /// Aggregate compression ratio for `best_model`.
70    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
91/// Records one canonical envelope in the global session snapshot.
92pub 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
122/// Returns a clone of the current session usage.
123pub fn session_usage() -> SessionUsage {
124    usage_guard().clone()
125}
126
127/// Returns per-model usage sorted by descending input and output tokens.
128pub 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/// Returns per-provider usage sorted by descending total tokens.
141#[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/// Returns aggregate compression metrics for the current session.
153#[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
179/// Clears all accumulated usage while preserving the allocated global mutex.
180pub 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}