Skip to main content

entrenar_shell/
state.rs

1//! Session state management for the interactive shell.
2
3use serde::{Deserialize, Serialize};
4use std::collections::HashMap;
5use std::path::PathBuf;
6
7/// Session state that persists across commands.
8#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq)]
9pub struct SessionState {
10    /// Currently loaded models
11    models: HashMap<String, LoadedModel>,
12    /// Command history
13    history: Vec<HistoryEntry>,
14    /// User preferences
15    preferences: Preferences,
16    /// Session metrics
17    metrics: SessionMetrics,
18}
19
20impl SessionState {
21    /// Create a new empty session state.
22    pub fn new() -> Self {
23        Self::default()
24    }
25
26    /// Get loaded models.
27    pub fn loaded_models(&self) -> &HashMap<String, LoadedModel> {
28        &self.models
29    }
30
31    /// Get command history.
32    pub fn history(&self) -> &[HistoryEntry] {
33        &self.history
34    }
35
36    /// Add a model to the session.
37    pub fn add_model(&mut self, name: String, model: LoadedModel) {
38        self.models.insert(name, model);
39    }
40
41    /// Remove a model from the session.
42    pub fn remove_model(&mut self, name: &str) -> Option<LoadedModel> {
43        self.models.remove(name)
44    }
45
46    /// Get a model by name.
47    pub fn get_model(&self, name: &str) -> Option<&LoadedModel> {
48        self.models.get(name)
49    }
50
51    /// Add a command to history.
52    pub fn add_to_history(&mut self, entry: HistoryEntry) {
53        self.history.push(entry);
54    }
55
56    /// Get mutable preferences.
57    pub fn preferences_mut(&mut self) -> &mut Preferences {
58        &mut self.preferences
59    }
60
61    /// Get preferences.
62    pub fn preferences(&self) -> &Preferences {
63        &self.preferences
64    }
65
66    /// Get session metrics.
67    pub fn metrics(&self) -> &SessionMetrics {
68        &self.metrics
69    }
70
71    /// Update metrics after a command.
72    pub fn record_command(&mut self, duration_ms: u64, success: bool) {
73        self.metrics.total_commands += 1;
74        if success {
75            self.metrics.successful_commands += 1;
76        }
77        self.metrics.total_duration_ms += duration_ms;
78    }
79
80    /// Save state to a file.
81    pub fn save(&self, path: &PathBuf) -> std::io::Result<()> {
82        let json = serde_json::to_string_pretty(self)?;
83        std::fs::write(path, json)
84    }
85
86    /// Load state from a file.
87    ///
88    /// #2519: a session file is UNTRUSTED INPUT. `LoadedModel` derives
89    /// `Deserialize`, so before this check a hand-written JSON file could put
90    /// any model facts it liked into the session — architecture, parameter
91    /// count, layer count, hidden dim — and every downstream command
92    /// (`inspect`, `memory`, `distill`) then presented them as measured fact.
93    /// That is the *same* fabrication #2519 closed in `fetch`, reached through
94    /// a second door the `fetch` falsifier could not see. Measured before this
95    /// change, with a crafted session naming `/nonexistent`:
96    ///
97    /// ```text
98    /// $ aprender-train-shell --session sess.json -c "distill --dry-run"
99    /// Teacher: does-not-exist/totally-fake-7b (7.0B)
100    /// Student: does-not-exist/totally-fake-1b (1.0B)
101    /// Ready to train                                     # exit 0
102    /// ```
103    ///
104    /// Loading now fails rather than admitting unprovenanced models.
105    pub fn load(path: &PathBuf) -> std::io::Result<Self> {
106        let json = std::fs::read_to_string(path)?;
107        let state: Self = serde_json::from_str(&json)
108            .map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e))?;
109        state
110            .validate_model_provenance()
111            .map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e))?;
112        Ok(state)
113    }
114
115    /// Reject models whose facts this shell could not have derived.
116    ///
117    /// Two rules, deliberately ordered so the second can be deleted on its own
118    /// the day a real loader lands, leaving the first standing:
119    ///
120    /// 1. A cached model is ON DISK. A session naming a path that does not
121    ///    exist is describing a model nobody has.
122    /// 2. Nothing in this crate can produce a `LoadedModel` at all: `fetch`
123    ///    refuses (it has no HuggingFace client, #2519) and no other
124    ///    production path calls `add_model`. So *any* model in a session file
125    ///    was typed by hand, not measured, whatever its path says.
126    ///
127    /// When a real fetch/loader is implemented, delete rule 2 and the
128    /// `no_loader_exists` half of the falsifier with it — but rule 1 stays.
129    pub fn validate_model_provenance(&self) -> Result<(), String> {
130        for (name, model) in &self.models {
131            if !model.path.exists() {
132                return Err(format!(
133                    "session claims model `{name}` ({}) is cached at {}, but nothing \
134                     is there. A model that is not on disk has no measurable \
135                     architecture, parameter count or layer count. (#2519)",
136                    model.id,
137                    model.path.display()
138                ));
139            }
140        }
141
142        if let Some((name, model)) = self.models.iter().next() {
143            return Err(format!(
144                "session carries model `{name}` ({}), but this shell has no way to \
145                 load a model: `fetch` downloads nothing and nothing else records \
146                 one. Every field of that entry — architecture `{}`, {} parameters, \
147                 {} layers, hidden_dim {} — was written by hand, not measured, so \
148                 the shell will not restate it as fact. (#2519)",
149                model.id, model.architecture, model.parameters, model.layers, model.hidden_dim
150            ));
151        }
152
153        Ok(())
154    }
155}
156
157/// A loaded model in the session.
158#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
159pub struct LoadedModel {
160    /// Model identifier (HuggingFace ID or path)
161    pub id: String,
162    /// Local path to cached model
163    pub path: PathBuf,
164    /// Model architecture
165    pub architecture: String,
166    /// Number of parameters
167    pub parameters: u64,
168    /// Number of layers
169    pub layers: u32,
170    /// Hidden dimension
171    pub hidden_dim: u32,
172    /// Role in session (teacher/student)
173    pub role: ModelRole,
174}
175
176/// Role of a model in the distillation session.
177#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
178pub enum ModelRole {
179    /// Teacher model (knowledge source)
180    Teacher,
181    /// Student model (learning target)
182    Student,
183    /// No specific role assigned
184    #[default]
185    None,
186}
187
188/// A history entry for a command.
189#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
190pub struct HistoryEntry {
191    /// The command string
192    pub command: String,
193    /// Execution timestamp (Unix seconds)
194    pub timestamp: u64,
195    /// Duration in milliseconds
196    pub duration_ms: u64,
197    /// Whether the command succeeded
198    pub success: bool,
199}
200
201impl HistoryEntry {
202    /// Create a new history entry.
203    pub fn new(command: impl Into<String>, duration_ms: u64, success: bool) -> Self {
204        Self {
205            command: command.into(),
206            timestamp: std::time::SystemTime::now()
207                .duration_since(std::time::UNIX_EPOCH)
208                .map(|d| d.as_secs())
209                .unwrap_or(0),
210            duration_ms,
211            success,
212        }
213    }
214}
215
216/// User preferences.
217#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
218pub struct Preferences {
219    /// Default output format
220    pub output_format: String,
221    /// Whether to show progress bars
222    pub show_progress: bool,
223    /// Whether to save history automatically
224    pub auto_save_history: bool,
225    /// Default batch size for operations
226    pub default_batch_size: u32,
227    /// Default sequence length
228    pub default_seq_len: usize,
229}
230
231impl Default for Preferences {
232    fn default() -> Self {
233        Self {
234            output_format: "table".to_string(),
235            show_progress: true,
236            auto_save_history: true,
237            default_batch_size: 32,
238            default_seq_len: 512,
239        }
240    }
241}
242
243/// Session-level metrics.
244#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq)]
245pub struct SessionMetrics {
246    /// Total commands executed
247    pub total_commands: u64,
248    /// Successful commands
249    pub successful_commands: u64,
250    /// Total duration in milliseconds
251    pub total_duration_ms: u64,
252}
253
254impl SessionMetrics {
255    /// Get success rate as a percentage.
256    pub fn success_rate(&self) -> f64 {
257        if self.total_commands == 0 {
258            100.0
259        } else {
260            (self.successful_commands as f64 / self.total_commands as f64) * 100.0
261        }
262    }
263
264    /// Get average command duration in milliseconds.
265    pub fn avg_duration_ms(&self) -> f64 {
266        if self.total_commands == 0 {
267            0.0
268        } else {
269            self.total_duration_ms as f64 / self.total_commands as f64
270        }
271    }
272}
273
274#[cfg(test)]
275mod tests {
276    use super::*;
277
278    #[test]
279    fn test_session_state_model_management() {
280        let mut state = SessionState::new();
281
282        let model = LoadedModel {
283            id: "test/model".to_string(),
284            path: PathBuf::from("/tmp/model"),
285            architecture: "llama".to_string(),
286            parameters: 7_000_000_000,
287            layers: 32,
288            hidden_dim: 4096,
289            role: ModelRole::Teacher,
290        };
291
292        state.add_model("teacher".to_string(), model.clone());
293        assert_eq!(state.loaded_models().len(), 1);
294        assert!(state.get_model("teacher").is_some());
295
296        state.remove_model("teacher");
297        assert!(state.get_model("teacher").is_none());
298    }
299
300    #[test]
301    fn test_session_state_history() {
302        let mut state = SessionState::new();
303
304        state.add_to_history(HistoryEntry::new("fetch model", 100, true));
305        state.add_to_history(HistoryEntry::new("inspect layers", 50, true));
306
307        assert_eq!(state.history().len(), 2);
308        assert_eq!(state.history()[0].command, "fetch model");
309    }
310
311    #[test]
312    fn test_session_metrics() {
313        let mut state = SessionState::new();
314
315        state.record_command(100, true);
316        state.record_command(200, true);
317        state.record_command(150, false);
318
319        assert_eq!(state.metrics().total_commands, 3);
320        assert_eq!(state.metrics().successful_commands, 2);
321        assert!((state.metrics().success_rate() - 66.67).abs() < 1.0);
322    }
323
324    #[test]
325    fn test_session_state_serialization_roundtrip() {
326        let mut state = SessionState::new();
327        state.add_to_history(HistoryEntry::new("test", 100, true));
328        state.preferences_mut().default_batch_size = 64;
329
330        let json = serde_json::to_string(&state).expect("JSON serialization should succeed");
331        let restored: SessionState =
332            serde_json::from_str(&json).expect("JSON deserialization should succeed");
333
334        assert_eq!(state, restored);
335    }
336
337    #[test]
338    fn test_model_role_default() {
339        assert_eq!(ModelRole::default(), ModelRole::None);
340    }
341
342    #[test]
343    fn test_preferences_default_values() {
344        let prefs = Preferences::default();
345        assert_eq!(prefs.output_format, "table");
346        assert!(prefs.show_progress);
347        assert_eq!(prefs.default_batch_size, 32);
348    }
349
350    #[test]
351    fn test_session_metrics_success_rate_zero() {
352        let metrics = SessionMetrics::default();
353        assert_eq!(metrics.success_rate(), 100.0);
354    }
355
356    #[test]
357    fn test_session_metrics_avg_duration_zero() {
358        let metrics = SessionMetrics::default();
359        assert_eq!(metrics.avg_duration_ms(), 0.0);
360    }
361
362    #[test]
363    fn test_session_metrics_avg_duration() {
364        let mut state = SessionState::new();
365        state.record_command(100, true);
366        state.record_command(200, true);
367        assert_eq!(state.metrics().avg_duration_ms(), 150.0);
368    }
369
370    #[test]
371    fn test_history_entry_new() {
372        let entry = HistoryEntry::new("test command", 50, true);
373        assert_eq!(entry.command, "test command");
374        assert_eq!(entry.duration_ms, 50);
375        assert!(entry.success);
376        assert!(entry.timestamp > 0);
377    }
378
379    #[test]
380    fn test_loaded_model_equality() {
381        let model1 = LoadedModel {
382            id: "test".to_string(),
383            path: PathBuf::from("/tmp"),
384            architecture: "llama".to_string(),
385            parameters: 7_000_000_000,
386            layers: 32,
387            hidden_dim: 4096,
388            role: ModelRole::None,
389        };
390        let model2 = model1.clone();
391        assert_eq!(model1, model2);
392    }
393
394    #[test]
395    fn test_model_role_equality() {
396        assert_eq!(ModelRole::Teacher, ModelRole::Teacher);
397        assert_ne!(ModelRole::Teacher, ModelRole::Student);
398        assert_ne!(ModelRole::Student, ModelRole::None);
399    }
400
401    #[test]
402    fn test_session_state_save_load() {
403        use tempfile::TempDir;
404
405        let temp_dir = TempDir::new().expect("temp file creation should succeed");
406        let state_path = temp_dir.path().join("state.json");
407
408        let mut state = SessionState::new();
409        state.add_to_history(HistoryEntry::new("test", 100, true));
410        state.preferences_mut().default_batch_size = 128;
411
412        state.save(&state_path).expect("save should succeed");
413        let loaded = SessionState::load(&state_path).expect("load should succeed");
414
415        assert_eq!(state, loaded);
416    }
417
418    #[test]
419    fn test_session_state_load_invalid_json() {
420        use std::io::Write;
421        use tempfile::NamedTempFile;
422
423        let mut file = NamedTempFile::new().expect("temp file creation should succeed");
424        file.write_all(b"not valid json")
425            .expect("file write should succeed");
426
427        let result = SessionState::load(&file.path().to_path_buf());
428        assert!(result.is_err());
429    }
430
431    #[test]
432    fn test_preferences_all_fields() {
433        let prefs = Preferences::default();
434        assert_eq!(prefs.output_format, "table");
435        assert!(prefs.show_progress);
436        assert!(prefs.auto_save_history);
437        assert_eq!(prefs.default_batch_size, 32);
438        assert_eq!(prefs.default_seq_len, 512);
439    }
440
441    #[test]
442    fn test_session_state_remove_nonexistent() {
443        let mut state = SessionState::new();
444        let result = state.remove_model("nonexistent");
445        assert!(result.is_none());
446    }
447
448    #[test]
449    fn test_session_state_get_nonexistent() {
450        let state = SessionState::new();
451        assert!(state.get_model("nonexistent").is_none());
452    }
453
454    #[test]
455    fn test_session_metrics_fields() {
456        let metrics = SessionMetrics {
457            total_commands: 10,
458            successful_commands: 8,
459            total_duration_ms: 1000,
460        };
461        assert_eq!(metrics.success_rate(), 80.0);
462        assert_eq!(metrics.avg_duration_ms(), 100.0);
463    }
464}