Skip to main content

lean_ctx/core/
mode_predictor.rs

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/// Observed outcome of a read mode: tokens in/out and information density.
11#[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    /// Computes an efficiency score: density / compression ratio.
21    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/// File identity for mode prediction: extension + token-count size bucket.
30#[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    /// Creates a file signature from its path and token count.
38    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/// Learns the best read mode per file signature from historical outcomes.
56#[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    /// Loads or creates the predictor, using an in-memory buffer for caching.
64    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    /// Records a mode outcome for a file signature (capped at 100 entries).
94    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    /// Returns the best mode based on historical efficiency.
103    /// Chain: local history -> cloud adaptive models -> built-in defaults.
104    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            // For code-structured defaults, never override to lossy modes.
115            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        // Arm semantics are defined by the trainer (`feedback::update_bandit`),
155        // which buckets each outcome by the entropy threshold actually used: a
156        // HIGH threshold (>= 1.0, the *most* compression) trains `conservative`,
157        // a LOW threshold (< 0.7, the *least*) trains `aggressive`. So a winning
158        // `conservative` arm means "high compression has been succeeding" and must
159        // map to a high-compression mode. The previous `conservative => "full"`
160        // inverted this: it disabled compression precisely when aggressive
161        // compression was working (GL #622). Spans heaviest → lightest structural
162        // compression; `full` is intentionally not a learned suggestion (forced
163        // full reads are handled by `should_force_full`).
164        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    /// Loads cloud adaptive models (synced from LeanCTX Cloud).
199    /// Models are cached locally and auto-updated for cloud users.
200    #[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    /// Built-in defaults for common file types and sizes.
248    /// Ensures reasonable compression even without local history or cloud models.
249    /// Respects Kolmogorov-Gate: files with K>0.7 skip aggressive modes.
250    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            // Large code files: signatures only
260            (
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            // Code 2k-10k, SQL, lock, config/data: structured map
267            ("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            // CSS, XML/CSV, and large unknown files: aggressive
276            ("xml" | "csv", _) | ("css" | "scss" | "less" | "sass", 2..) | (_, 3..) => "aggressive",
277
278            _ => return None,
279        };
280        Some(mode.to_string())
281    }
282
283    /// Saves to the in-memory buffer and flushes to disk if the interval elapsed.
284    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    /// Forces an immediate write of the buffered predictor state to disk.
313    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        // GL #622: `feedback::update_bandit` rewards the `conservative` arm on
398        // HIGH-compression success, so a winning `conservative` arm must resolve
399        // to a high-compression mode (not `full`). Guards against re-inverting the
400        // arm→mode mapping.
401        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}