Skip to main content

lean_ctx/core/context_kernel/
token_envelope.rs

1//! Provider-neutral token usage representation.
2
3use serde::{Deserialize, Serialize};
4
5/// Provider responsible for serving a model request.
6#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Hash, Serialize, Deserialize)]
7pub enum ProviderKind {
8    /// OpenAI API.
9    OpenAi,
10    /// Anthropic API.
11    Anthropic,
12    /// Google Gemini API.
13    Gemini,
14    /// OpenRouter gateway.
15    OpenRouter,
16    /// Amazon Bedrock.
17    Bedrock,
18    /// Azure OpenAI Service.
19    Azure,
20    /// Locally hosted model.
21    Local,
22    /// Unknown or unavailable provider.
23    #[default]
24    Unknown,
25}
26
27/// Provider-neutral token usage envelope.
28#[derive(Debug, Clone, Default, Serialize, Deserialize)]
29pub struct TokenEnvelope {
30    /// Canonical model identifier.
31    pub model: String,
32    /// Provider that served the request.
33    pub provider: ProviderKind,
34    /// Total input tokens (prompt + context).
35    pub input_tokens: usize,
36    /// Total output tokens (completion).
37    pub output_tokens: usize,
38    /// Cache read tokens (provider-specific).
39    pub cache_read_tokens: usize,
40    /// Cache write tokens.
41    pub cache_write_tokens: usize,
42    /// Reasoning/thinking tokens (if supported).
43    pub reasoning_tokens: usize,
44    /// Estimated cost in USD (if available).
45    pub cost_usd: Option<f64>,
46    /// Tokens saved by compression.
47    pub tokens_saved: usize,
48    /// Whether this was a retry.
49    pub is_retry: bool,
50}
51
52/// Converts proxy request usage into a canonical token envelope.
53#[must_use]
54pub fn from_proxy_data(data: &super::proxy_bridge::ProxyRequestData) -> TokenEnvelope {
55    TokenEnvelope {
56        model: data.model.clone().unwrap_or_default(),
57        provider: data
58            .provider
59            .as_deref()
60            .map_or(ProviderKind::Unknown, parse_provider),
61        input_tokens: data.input_tokens,
62        output_tokens: data.output_tokens,
63        reasoning_tokens: data.reasoning_tokens,
64        tokens_saved: data.tokens_saved,
65        is_retry: data.is_retry,
66        ..TokenEnvelope::default()
67    }
68}
69
70/// Converts MCP call usage into a canonical token envelope.
71#[must_use]
72pub fn from_mcp_call(data: &super::mcp_bridge::McpCallData) -> TokenEnvelope {
73    TokenEnvelope {
74        provider: ProviderKind::Unknown,
75        input_tokens: data.input_tokens,
76        output_tokens: data.output_tokens,
77        is_retry: data.is_retry,
78        ..TokenEnvelope::default()
79    }
80}
81
82/// Parses a provider label into its canonical provider kind.
83#[must_use]
84pub fn parse_provider(label: &str) -> ProviderKind {
85    match label.trim().to_ascii_lowercase().as_str() {
86        "openai" => ProviderKind::OpenAi,
87        "anthropic" => ProviderKind::Anthropic,
88        "gemini" | "google" => ProviderKind::Gemini,
89        "openrouter" => ProviderKind::OpenRouter,
90        "bedrock" => ProviderKind::Bedrock,
91        "azure" | "azure_openai" => ProviderKind::Azure,
92        "local" => ProviderKind::Local,
93        _ => ProviderKind::Unknown,
94    }
95}
96
97impl TokenEnvelope {
98    /// Returns input, output, and reasoning tokens combined.
99    #[must_use]
100    pub fn total_tokens(&self) -> usize {
101        self.input_tokens
102            .saturating_add(self.output_tokens)
103            .saturating_add(self.reasoning_tokens)
104    }
105
106    /// Returns total tokens excluding tokens served from the provider cache.
107    #[must_use]
108    pub fn effective_tokens(&self) -> usize {
109        self.total_tokens().saturating_sub(self.cache_read_tokens)
110    }
111
112    /// Returns the fraction of original input eliminated by compression.
113    #[must_use]
114    pub fn compression_ratio(&self) -> f64 {
115        let original_input = self.input_tokens.saturating_add(self.tokens_saved);
116        if original_input == 0 {
117            0.0
118        } else {
119            self.tokens_saved as f64 / original_input as f64
120        }
121    }
122
123    /// Returns whether any input tokens were served from a provider cache.
124    #[must_use]
125    pub const fn is_cached(&self) -> bool {
126        self.cache_read_tokens > 0
127    }
128
129    /// Aggregates envelopes into a single canonical usage summary.
130    #[must_use]
131    pub fn merge(envelopes: &[Self]) -> Self {
132        let Some(first) = envelopes.first() else {
133            return Self::default();
134        };
135
136        let same_model = envelopes
137            .iter()
138            .all(|envelope| envelope.model == first.model);
139        let same_provider = envelopes
140            .iter()
141            .all(|envelope| envelope.provider == first.provider);
142        let sum = |field: fn(&Self) -> usize| {
143            envelopes.iter().fold(0usize, |total, envelope| {
144                total.saturating_add(field(envelope))
145            })
146        };
147
148        Self {
149            model: if same_model {
150                first.model.clone()
151            } else {
152                String::new()
153            },
154            provider: if same_provider {
155                first.provider
156            } else {
157                ProviderKind::Unknown
158            },
159            input_tokens: sum(|envelope| envelope.input_tokens),
160            output_tokens: sum(|envelope| envelope.output_tokens),
161            cache_read_tokens: sum(|envelope| envelope.cache_read_tokens),
162            cache_write_tokens: sum(|envelope| envelope.cache_write_tokens),
163            reasoning_tokens: sum(|envelope| envelope.reasoning_tokens),
164            cost_usd: envelopes
165                .iter()
166                .filter_map(|envelope| envelope.cost_usd)
167                .reduce(|total, cost| total + cost),
168            tokens_saved: sum(|envelope| envelope.tokens_saved),
169            is_retry: envelopes.iter().any(|envelope| envelope.is_retry),
170        }
171    }
172}
173
174#[cfg(test)]
175mod tests {
176    use super::{ProviderKind, TokenEnvelope, from_mcp_call, from_proxy_data, parse_provider};
177    use crate::core::context_kernel::mcp_bridge::McpCallData;
178    use crate::core::context_kernel::proxy_bridge::ProxyRequestData;
179
180    #[test]
181    fn from_proxy_openai() {
182        let envelope = from_proxy_data(&ProxyRequestData {
183            provider: Some("OpenAI".to_owned()),
184            model: Some("gpt-5".to_owned()),
185            input_tokens: 100,
186            output_tokens: 20,
187            reasoning_tokens: 5,
188            tokens_saved: 30,
189            is_retry: true,
190            ..ProxyRequestData::default()
191        });
192
193        assert_eq!(envelope.provider, ProviderKind::OpenAi);
194        assert_eq!(envelope.model, "gpt-5");
195        assert_eq!(envelope.total_tokens(), 125);
196        assert_eq!(envelope.tokens_saved, 30);
197        assert!(envelope.is_retry);
198    }
199
200    #[test]
201    fn from_proxy_anthropic() {
202        let envelope = from_proxy_data(&ProxyRequestData {
203            provider: Some("Anthropic".to_owned()),
204            ..ProxyRequestData::default()
205        });
206
207        assert_eq!(envelope.provider, ProviderKind::Anthropic);
208    }
209
210    #[test]
211    fn from_mcp_call_maps_usage() {
212        let envelope = from_mcp_call(&McpCallData {
213            input_tokens: 80,
214            output_tokens: 12,
215            is_retry: true,
216            ..McpCallData::default()
217        });
218
219        assert_eq!(envelope.provider, ProviderKind::Unknown);
220        assert_eq!(envelope.input_tokens, 80);
221        assert_eq!(envelope.output_tokens, 12);
222        assert!(envelope.is_retry);
223    }
224
225    #[test]
226    fn total_tokens_sum() {
227        let envelope = TokenEnvelope {
228            input_tokens: 100,
229            output_tokens: 20,
230            reasoning_tokens: 7,
231            ..TokenEnvelope::default()
232        };
233
234        assert_eq!(envelope.total_tokens(), 127);
235    }
236
237    #[test]
238    fn effective_excludes_cache() {
239        let envelope = TokenEnvelope {
240            input_tokens: 100,
241            output_tokens: 20,
242            reasoning_tokens: 7,
243            cache_read_tokens: 40,
244            ..TokenEnvelope::default()
245        };
246
247        assert_eq!(envelope.effective_tokens(), 87);
248        assert!(envelope.is_cached());
249    }
250
251    #[test]
252    fn compression_ratio_correct() {
253        let envelope = TokenEnvelope {
254            input_tokens: 1_000,
255            tokens_saved: 300,
256            ..TokenEnvelope::default()
257        };
258
259        assert!((envelope.compression_ratio() - 0.230_769).abs() < 0.000_001);
260    }
261
262    #[test]
263    fn merge_aggregates() {
264        let envelopes = (1..=3)
265            .map(|multiplier| TokenEnvelope {
266                model: "gpt-5".to_owned(),
267                provider: ProviderKind::OpenAi,
268                input_tokens: 10 * multiplier,
269                output_tokens: 2 * multiplier,
270                cache_read_tokens: multiplier,
271                cache_write_tokens: multiplier,
272                reasoning_tokens: multiplier,
273                cost_usd: Some(0.01 * multiplier as f64),
274                tokens_saved: 3 * multiplier,
275                is_retry: multiplier == 3,
276            })
277            .collect::<Vec<_>>();
278
279        let merged = TokenEnvelope::merge(&envelopes);
280        assert_eq!(merged.model, "gpt-5");
281        assert_eq!(merged.provider, ProviderKind::OpenAi);
282        assert_eq!(merged.input_tokens, 60);
283        assert_eq!(merged.output_tokens, 12);
284        assert_eq!(merged.cache_read_tokens, 6);
285        assert_eq!(merged.cache_write_tokens, 6);
286        assert_eq!(merged.reasoning_tokens, 6);
287        assert!((merged.cost_usd.unwrap_or_default() - 0.06).abs() < f64::EPSILON);
288        assert_eq!(merged.tokens_saved, 18);
289        assert!(merged.is_retry);
290    }
291
292    #[test]
293    fn parse_case_insensitive() {
294        for label in ["openai", "OPENAI", "OpenAI"] {
295            assert_eq!(parse_provider(label), ProviderKind::OpenAi);
296        }
297    }
298
299    #[test]
300    fn parse_provider_aliases_and_unknown() {
301        assert_eq!(parse_provider("google"), ProviderKind::Gemini);
302        assert_eq!(parse_provider("openrouter"), ProviderKind::OpenRouter);
303        assert_eq!(parse_provider("local"), ProviderKind::Local);
304        assert_eq!(parse_provider("other"), ProviderKind::Unknown);
305    }
306}