lean_ctx/core/context_kernel/
envelope_bridge.rs1use std::cmp::Reverse;
4use std::collections::HashMap;
5use std::sync::{Mutex, MutexGuard, OnceLock};
6
7use serde::Serialize;
8
9use super::token_envelope::{ProviderKind, TokenEnvelope};
10use super::{evidence_wiring, kernel_config, usage_normalizer};
11
12#[derive(Debug, Clone, Default, Serialize)]
14pub struct ProviderStat {
15 pub provider: ProviderKind,
17 pub request_count: usize,
19 pub total_input: usize,
21 pub total_output: usize,
23 pub total_cache_read: usize,
25 pub avg_input: usize,
27}
28
29#[derive(Debug, Default)]
30struct ProviderAccum {
31 count: usize,
32 sum_input: usize,
33 sum_output: usize,
34 sum_cache_read: usize,
35}
36
37static PROVIDER_STATS: OnceLock<Mutex<HashMap<ProviderKind, ProviderAccum>>> = OnceLock::new();
38
39fn stats_guard() -> MutexGuard<'static, HashMap<ProviderKind, ProviderAccum>> {
40 PROVIDER_STATS
41 .get_or_init(|| Mutex::new(HashMap::new()))
42 .lock()
43 .unwrap_or_else(std::sync::PoisonError::into_inner)
44}
45
46fn record_stats(envelope: &TokenEnvelope) {
47 let mut stats = stats_guard();
48 let accum = stats.entry(envelope.provider).or_default();
49 accum.count = accum.count.saturating_add(1);
50 accum.sum_input = accum.sum_input.saturating_add(envelope.input_tokens);
51 accum.sum_output = accum.sum_output.saturating_add(envelope.output_tokens);
52 accum.sum_cache_read = accum
53 .sum_cache_read
54 .saturating_add(envelope.cache_read_tokens);
55}
56
57pub fn record_proxy_envelope(envelope: &TokenEnvelope) {
59 if kernel_config::is_enabled() {
60 let provider = format!("{:?}", envelope.provider);
61 evidence_wiring::record_from_proxy_dispatch(
62 envelope.input_tokens,
63 envelope.output_tokens,
64 envelope.tokens_saved,
65 Some(&envelope.model),
66 Some(&provider),
67 );
68 usage_normalizer::record_envelope(envelope);
69 }
70 record_stats(envelope);
71}
72
73pub fn record_mcp_envelope(tool_name: &str, envelope: &TokenEnvelope) {
75 if kernel_config::is_enabled() {
76 evidence_wiring::record_from_tool_dispatch(
77 tool_name,
78 envelope.input_tokens,
79 envelope.output_tokens,
80 envelope.tokens_saved,
81 );
82 usage_normalizer::record_envelope(envelope);
83 }
84 record_stats(envelope);
85}
86
87#[must_use]
89pub fn provider_stats() -> Vec<ProviderStat> {
90 let mut stats = stats_guard()
91 .iter()
92 .map(|(provider, accum)| ProviderStat {
93 provider: *provider,
94 request_count: accum.count,
95 total_input: accum.sum_input,
96 total_output: accum.sum_output,
97 total_cache_read: accum.sum_cache_read,
98 avg_input: accum.sum_input.checked_div(accum.count).unwrap_or_default(),
99 })
100 .collect::<Vec<_>>();
101 stats.sort_unstable_by_key(|stat| (Reverse(stat.request_count), stat.provider as u8));
102 stats
103}
104
105pub fn reset() {
107 stats_guard().clear();
108}
109
110#[cfg(test)]
111mod tests {
112 use super::{provider_stats, record_mcp_envelope, record_proxy_envelope, reset};
113 use crate::core::context_kernel::token_envelope::{ProviderKind, TokenEnvelope};
114 use crate::core::context_kernel::{evidence_wiring, kernel_config, usage_normalizer};
115
116 fn isolated() -> std::sync::MutexGuard<'static, ()> {
117 let guard = kernel_config::KERNEL_TEST_LOCK
118 .lock()
119 .unwrap_or_else(std::sync::PoisonError::into_inner);
120 kernel_config::reset_features();
121 evidence_wiring::reset();
122 usage_normalizer::reset_usage();
123 reset();
124 guard
125 }
126
127 fn envelope(provider: ProviderKind, input: usize) -> TokenEnvelope {
128 TokenEnvelope {
129 model: "test-model".to_owned(),
130 provider,
131 input_tokens: input,
132 output_tokens: 20,
133 cache_read_tokens: 5,
134 ..TokenEnvelope::default()
135 }
136 }
137
138 #[test]
139 fn record_proxy_updates_stats() {
140 let _guard = isolated();
141 for _ in 0..3 {
142 record_proxy_envelope(&envelope(ProviderKind::OpenAi, 100));
143 }
144 assert_eq!(provider_stats()[0].request_count, 3);
145 }
146
147 #[test]
148 fn record_mcp_updates_stats() {
149 let _guard = isolated();
150 record_mcp_envelope("ctx_read", &envelope(ProviderKind::Gemini, 80));
151 let stats = provider_stats();
152 assert_eq!((stats[0].total_input, stats[0].total_output), (80, 20));
153 assert_eq!(stats[0].total_cache_read, 5);
154 }
155
156 #[test]
157 fn multi_provider_tracked() {
158 let _guard = isolated();
159 record_proxy_envelope(&envelope(ProviderKind::OpenAi, 100));
160 record_proxy_envelope(&envelope(ProviderKind::Anthropic, 200));
161 assert_eq!(provider_stats().len(), 2);
162 }
163
164 #[test]
165 fn reset_clears() {
166 let _guard = isolated();
167 record_proxy_envelope(&envelope(ProviderKind::OpenAi, 100));
168 reset();
169 assert!(provider_stats().is_empty());
170 }
171
172 #[test]
173 fn empty_stats_safe() {
174 let _guard = isolated();
175 assert!(provider_stats().is_empty());
176 }
177
178 #[test]
179 fn avg_input_correct() {
180 let _guard = isolated();
181 for input in [100, 200, 300] {
182 record_proxy_envelope(&envelope(ProviderKind::OpenAi, input));
183 }
184 assert_eq!(provider_stats()[0].avg_input, 200);
185 }
186}