1use std::collections::HashMap;
2use std::sync::{Arc, Mutex};
3use std::time::Instant;
4
5const STATS_FILE: &str = "mode_stats.json";
6const PREDICTOR_FLUSH_SECS: u64 = 10;
7
8static PREDICTOR_BUFFER: Mutex<Option<(Arc<ModePredictor>, Instant)>> = Mutex::new(None);
9
10#[derive(Clone, Debug, serde::Serialize, serde::Deserialize)]
12pub struct ModeOutcome {
13 pub mode: String,
14 pub tokens_in: usize,
15 pub tokens_out: usize,
16 pub density: f64,
17}
18
19impl ModeOutcome {
20 pub fn efficiency(&self) -> f64 {
22 if self.tokens_out == 0 {
23 return 0.0;
24 }
25 self.density / (self.tokens_out as f64 / self.tokens_in.max(1) as f64)
26 }
27}
28
29#[derive(Clone, Debug, Hash, Eq, PartialEq, serde::Serialize, serde::Deserialize)]
31pub struct FileSignature {
32 pub ext: String,
33 pub size_bucket: u8,
34}
35
36impl FileSignature {
37 pub fn from_path(path: &str, token_count: usize) -> Self {
39 let ext = std::path::Path::new(path)
40 .extension()
41 .and_then(|e| e.to_str())
42 .unwrap_or("")
43 .to_string();
44 let size_bucket = match token_count {
45 0..=500 => 0,
46 501..=2000 => 1,
47 2001..=5000 => 2,
48 5001..=20000 => 3,
49 _ => 4,
50 };
51 Self { ext, size_bucket }
52 }
53}
54
55#[derive(Debug, Default, Clone, serde::Serialize, serde::Deserialize)]
57pub struct ModePredictor {
58 history: HashMap<FileSignature, Vec<ModeOutcome>>,
59 project_root: Option<String>,
60}
61
62impl ModePredictor {
63 pub fn new() -> Self {
65 let mut guard = PREDICTOR_BUFFER
66 .lock()
67 .unwrap_or_else(std::sync::PoisonError::into_inner);
68 if let Some((ref predictor, _)) = *guard {
69 return Self {
70 history: predictor.history.clone(),
71 project_root: predictor.project_root.clone(),
72 };
73 }
74 let mut loaded = Self::load_from_disk().unwrap_or_default();
75 if loaded.project_root.is_none() {
76 loaded.project_root = std::env::current_dir()
77 .ok()
78 .map(|p| p.to_string_lossy().to_string());
79 }
80 *guard = Some((Arc::new(loaded.clone()), Instant::now()));
81 loaded
82 }
83
84 pub fn with_project_root(mut self, root: &str) -> Self {
85 self.project_root = Some(root.to_string());
86 self
87 }
88
89 pub fn set_project_root(&mut self, root: &str) {
90 self.project_root = Some(root.to_string());
91 }
92
93 pub fn record(&mut self, sig: FileSignature, outcome: ModeOutcome) {
95 let entries = self.history.entry(sig).or_default();
96 entries.push(outcome);
97 if entries.len() > 100 {
98 entries.drain(0..50);
99 }
100 }
101
102 pub fn predict_best_mode(&self, sig: &FileSignature) -> Option<String> {
105 let default_mode = Self::predict_from_defaults(sig);
106
107 let allow_override = |candidate: &str| -> bool {
108 let Some(def) = default_mode.as_deref() else {
109 return true;
110 };
111 if candidate == "full" {
112 return false;
113 }
114 if (def == "map" || def == "signatures")
116 && (candidate == "aggressive" || candidate == "entropy")
117 {
118 return false;
119 }
120 true
121 };
122
123 if let Some(local) = self.predict_from_local(sig)
124 && allow_override(&local)
125 {
126 return Some(local);
127 }
128 if let Some(bandit) = self.predict_from_bandit(sig)
129 && allow_override(&bandit)
130 {
131 return Some(bandit);
132 }
133 if let Some(cloud) = self.predict_from_cloud(sig)
134 && allow_override(&cloud)
135 {
136 return Some(cloud);
137 }
138 default_mode
139 }
140
141 fn predict_from_bandit(&self, sig: &FileSignature) -> Option<String> {
142 let key = format!("{}_feedback", sig.ext);
143 let store =
144 crate::core::bandit::BanditStore::load(self.project_root.as_deref().unwrap_or("."));
145 let bandit = store.bandits.get(&key)?;
146 if bandit.total_pulls < 5 {
147 return None;
148 }
149 let best_arm = bandit.arms.iter().max_by(|a, b| {
150 a.mean()
151 .partial_cmp(&b.mean())
152 .unwrap_or(std::cmp::Ordering::Equal)
153 })?;
154 let mode = match best_arm.name.as_str() {
165 "conservative" => "aggressive",
166 "balanced" => "signatures",
167 "aggressive" => "map",
168 _ => return None,
169 };
170 Some(mode.to_string())
171 }
172
173 fn predict_from_local(&self, sig: &FileSignature) -> Option<String> {
174 let entries = self.history.get(sig)?;
175 if entries.len() < 3 {
176 return None;
177 }
178
179 let mut mode_scores: HashMap<&str, (f64, usize)> = HashMap::new();
180 for entry in entries {
181 let (sum, count) = mode_scores.entry(&entry.mode).or_insert((0.0, 0));
182 *sum += entry.efficiency();
183 *count += 1;
184 }
185
186 mode_scores
187 .into_iter()
188 .max_by(|a, b| {
189 let avg_a = a.1.0 / a.1.1 as f64;
190 let avg_b = b.1.0 / b.1.1 as f64;
191 avg_a
192 .partial_cmp(&avg_b)
193 .unwrap_or(std::cmp::Ordering::Equal)
194 })
195 .map(|(mode, _)| mode.to_string())
196 }
197
198 #[allow(clippy::unused_self)]
201 fn predict_from_cloud(&self, sig: &FileSignature) -> Option<String> {
202 let data = crate::cloud_client::load_cloud_models()?;
203 let models = data["models"].as_array()?;
204
205 let ext_with_dot = format!(".{}", sig.ext);
206 let bucket_name = match sig.size_bucket {
207 0 => "0-500",
208 1 => "500-2k",
209 2 => "2k-10k",
210 _ => "10k+",
211 };
212
213 let mut best: Option<(&str, f64)> = None;
214
215 for model in models {
216 let m_ext = model["file_ext"].as_str().unwrap_or("");
217 let m_bucket = model["size_bucket"].as_str().unwrap_or("");
218 let confidence = model["confidence"].as_f64().unwrap_or(0.0);
219
220 if m_ext == ext_with_dot
221 && m_bucket == bucket_name
222 && confidence > 0.5
223 && let Some(mode) = model["recommended_mode"].as_str()
224 && best.is_none_or(|(_, c)| confidence > c)
225 {
226 best = Some((mode, confidence));
227 }
228 }
229
230 if let Some((mode, _)) = best {
231 return Some(mode.to_string());
232 }
233
234 for model in models {
235 let m_ext = model["file_ext"].as_str().unwrap_or("");
236 let confidence = model["confidence"].as_f64().unwrap_or(0.0);
237 if m_ext == ext_with_dot && confidence > 0.5 {
238 return model["recommended_mode"]
239 .as_str()
240 .map(std::string::ToString::to_string);
241 }
242 }
243
244 None
245 }
246
247 fn predict_from_defaults(sig: &FileSignature) -> Option<String> {
251 if sig.size_bucket == 0 {
252 return None;
253 }
254 if matches!(sig.ext.as_str(), "md" | "mdx" | "txt" | "rst") {
255 return None;
256 }
257
258 let mode = match (sig.ext.as_str(), sig.size_bucket) {
259 (
261 "rs" | "ts" | "tsx" | "js" | "jsx" | "py" | "go" | "java" | "c" | "cpp" | "rb"
262 | "swift" | "kt" | "cs" | "vue" | "svelte" | "gd",
263 4..,
264 ) => "signatures",
265
266 ("lock" | "json" | "yaml" | "yml" | "toml", _)
268 | (
269 "rs" | "ts" | "tsx" | "js" | "jsx" | "py" | "go" | "java" | "c" | "cpp" | "rb"
270 | "swift" | "kt" | "cs" | "vue" | "svelte" | "gd",
271 2 | 3,
272 )
273 | ("sql", 2..) => "map",
274
275 ("xml" | "csv", _) | ("css" | "scss" | "less" | "sass", 2..) | (_, 3..) => "aggressive",
277
278 _ => return None,
279 };
280 Some(mode.to_string())
281 }
282
283 pub fn save(&self) {
285 let mut guard = PREDICTOR_BUFFER
286 .lock()
287 .unwrap_or_else(std::sync::PoisonError::into_inner);
288 let should_flush = match *guard {
289 Some((_, ref last_flush)) => last_flush.elapsed().as_secs() >= PREDICTOR_FLUSH_SECS,
290 None => true,
291 };
292 *guard = Some((Arc::new(self.clone()), Instant::now()));
293 if should_flush {
294 self.save_to_disk();
295 }
296 }
297
298 fn save_to_disk(&self) {
299 let Ok(dir) = crate::core::data_dir::lean_ctx_data_dir() else {
300 return;
301 };
302 let _ = std::fs::create_dir_all(&dir);
303 let path = dir.join(STATS_FILE);
304 if let Ok(json) = serde_json::to_string_pretty(self) {
305 let tmp = dir.join(".mode_stats.tmp");
306 if std::fs::write(&tmp, &json).is_ok() {
307 let _ = std::fs::rename(&tmp, &path);
308 }
309 }
310 }
311
312 pub fn flush() {
314 let guard = PREDICTOR_BUFFER
315 .lock()
316 .unwrap_or_else(std::sync::PoisonError::into_inner);
317 if let Some((ref predictor, _)) = *guard {
318 predictor.save_to_disk();
319 }
320 }
321
322 fn load_from_disk() -> Option<Self> {
323 let path = crate::core::data_dir::lean_ctx_data_dir()
324 .ok()?
325 .join(STATS_FILE);
326 let data = std::fs::read_to_string(path).ok()?;
327 serde_json::from_str(&data).ok()
328 }
329}
330
331#[cfg(test)]
332mod tests {
333 use super::*;
334
335 #[test]
336 fn file_signature_buckets() {
337 assert_eq!(FileSignature::from_path("main.rs", 100).size_bucket, 0);
338 assert_eq!(FileSignature::from_path("main.rs", 1000).size_bucket, 1);
339 assert_eq!(FileSignature::from_path("main.rs", 3000).size_bucket, 2);
340 assert_eq!(FileSignature::from_path("main.rs", 10000).size_bucket, 3);
341 assert_eq!(FileSignature::from_path("main.rs", 50000).size_bucket, 4);
342 }
343
344 #[test]
345 fn predict_returns_none_without_history() {
346 let predictor = ModePredictor::default();
347 let sig = FileSignature::from_path("test.zzz", 500);
348 assert!(predictor.predict_from_local(&sig).is_none());
349 }
350
351 #[test]
352 fn predict_returns_none_with_too_few_entries() {
353 let mut predictor = ModePredictor::default();
354 let sig = FileSignature::from_path("test.zzz", 500);
355 predictor.record(
356 sig.clone(),
357 ModeOutcome {
358 mode: "full".to_string(),
359 tokens_in: 100,
360 tokens_out: 100,
361 density: 0.5,
362 },
363 );
364 assert!(predictor.predict_from_local(&sig).is_none());
365 }
366
367 #[test]
368 fn predict_learns_best_mode() {
369 let mut predictor = ModePredictor::default();
370 let sig = FileSignature::from_path("big.rs", 5000);
371 for _ in 0..5 {
372 predictor.record(
373 sig.clone(),
374 ModeOutcome {
375 mode: "full".to_string(),
376 tokens_in: 5000,
377 tokens_out: 5000,
378 density: 0.3,
379 },
380 );
381 predictor.record(
382 sig.clone(),
383 ModeOutcome {
384 mode: "map".to_string(),
385 tokens_in: 5000,
386 tokens_out: 800,
387 density: 0.6,
388 },
389 );
390 }
391 let best = predictor.predict_best_mode(&sig);
392 assert_eq!(best, Some("map".to_string()));
393 }
394
395 #[test]
396 fn predict_from_bandit_maps_conservative_to_high_compression() {
397 let _env = crate::core::data_dir::test_env_lock();
402 let data_dir = tempfile::tempdir().unwrap();
403 crate::test_env::set_var("LEAN_CTX_DATA_DIR", data_dir.path());
404
405 let project = tempfile::tempdir().unwrap();
406 let root = project.path().to_string_lossy().to_string();
407
408 let mut store = crate::core::bandit::BanditStore::default();
409 let bandit = store.get_or_create("rs_feedback");
410 bandit.total_pulls = 10;
411 for _ in 0..5 {
412 bandit.update("conservative", true);
413 }
414 store.save(&root).unwrap();
415
416 let mut predictor = ModePredictor::new();
417 predictor.set_project_root(&root);
418 let sig = FileSignature::from_path("big.rs", 5000);
419 assert_eq!(
420 predictor.predict_from_bandit(&sig),
421 Some("aggressive".to_string()),
422 "winning conservative arm must map to a high-compression mode, not full"
423 );
424 }
425
426 #[test]
427 fn history_caps_at_100() {
428 let mut predictor = ModePredictor::default();
429 let sig = FileSignature::from_path("test.rs", 100);
430 for _ in 0..120 {
431 predictor.record(
432 sig.clone(),
433 ModeOutcome {
434 mode: "full".to_string(),
435 tokens_in: 100,
436 tokens_out: 100,
437 density: 0.5,
438 },
439 );
440 }
441 assert!(predictor.history.get(&sig).unwrap().len() <= 100);
442 }
443
444 #[test]
445 fn defaults_return_none_for_small_files() {
446 let sig = FileSignature::from_path("small.rs", 200);
447 assert!(ModePredictor::predict_from_defaults(&sig).is_none());
448 }
449
450 #[test]
451 fn defaults_recommend_map_for_medium_code() {
452 let sig = FileSignature::from_path("medium.rs", 3000);
453 assert_eq!(
454 ModePredictor::predict_from_defaults(&sig),
455 Some("map".to_string())
456 );
457 }
458
459 #[test]
460 fn defaults_recommend_map_for_json() {
461 let sig = FileSignature::from_path("config.json", 1000);
462 assert_eq!(
463 ModePredictor::predict_from_defaults(&sig),
464 Some("map".to_string())
465 );
466 }
467
468 #[test]
469 fn defaults_recommend_signatures_for_huge_code() {
470 let sig = FileSignature::from_path("huge.ts", 25000);
471 assert_eq!(
472 ModePredictor::predict_from_defaults(&sig),
473 Some("signatures".to_string())
474 );
475 }
476
477 #[test]
478 fn defaults_recommend_aggressive_for_large_unknown() {
479 let sig = FileSignature::from_path("data.xyz", 8000);
480 assert_eq!(
481 ModePredictor::predict_from_defaults(&sig),
482 Some("aggressive".to_string())
483 );
484 }
485
486 #[test]
487 fn defaults_never_compress_markdown() {
488 for tokens in [600, 3000, 8000, 25000] {
489 let sig = FileSignature::from_path("SKILL.md", tokens);
490 assert!(
491 ModePredictor::predict_from_defaults(&sig).is_none(),
492 "SKILL.md at {tokens} tokens should get full (None), not compressed"
493 );
494 }
495 let sig = FileSignature::from_path("AGENTS.md", 5000);
496 assert!(ModePredictor::predict_from_defaults(&sig).is_none());
497 let sig = FileSignature::from_path("README.md", 12000);
498 assert!(ModePredictor::predict_from_defaults(&sig).is_none());
499 }
500
501 #[test]
502 fn mode_outcome_efficiency() {
503 let o = ModeOutcome {
504 mode: "map".to_string(),
505 tokens_in: 1000,
506 tokens_out: 200,
507 density: 0.6,
508 };
509 assert!(o.efficiency() > 0.0);
510 }
511}