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    pub fn load(path: &PathBuf) -> std::io::Result<Self> {
88        let json = std::fs::read_to_string(path)?;
89        serde_json::from_str(&json)
90            .map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e))
91    }
92}
93
94/// A loaded model in the session.
95#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
96pub struct LoadedModel {
97    /// Model identifier (HuggingFace ID or path)
98    pub id: String,
99    /// Local path to cached model
100    pub path: PathBuf,
101    /// Model architecture
102    pub architecture: String,
103    /// Number of parameters
104    pub parameters: u64,
105    /// Number of layers
106    pub layers: u32,
107    /// Hidden dimension
108    pub hidden_dim: u32,
109    /// Role in session (teacher/student)
110    pub role: ModelRole,
111}
112
113/// Role of a model in the distillation session.
114#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
115pub enum ModelRole {
116    /// Teacher model (knowledge source)
117    Teacher,
118    /// Student model (learning target)
119    Student,
120    /// No specific role assigned
121    #[default]
122    None,
123}
124
125/// A history entry for a command.
126#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
127pub struct HistoryEntry {
128    /// The command string
129    pub command: String,
130    /// Execution timestamp (Unix seconds)
131    pub timestamp: u64,
132    /// Duration in milliseconds
133    pub duration_ms: u64,
134    /// Whether the command succeeded
135    pub success: bool,
136}
137
138impl HistoryEntry {
139    /// Create a new history entry.
140    pub fn new(command: impl Into<String>, duration_ms: u64, success: bool) -> Self {
141        Self {
142            command: command.into(),
143            timestamp: std::time::SystemTime::now()
144                .duration_since(std::time::UNIX_EPOCH)
145                .map(|d| d.as_secs())
146                .unwrap_or(0),
147            duration_ms,
148            success,
149        }
150    }
151}
152
153/// User preferences.
154#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
155pub struct Preferences {
156    /// Default output format
157    pub output_format: String,
158    /// Whether to show progress bars
159    pub show_progress: bool,
160    /// Whether to save history automatically
161    pub auto_save_history: bool,
162    /// Default batch size for operations
163    pub default_batch_size: u32,
164    /// Default sequence length
165    pub default_seq_len: usize,
166}
167
168impl Default for Preferences {
169    fn default() -> Self {
170        Self {
171            output_format: "table".to_string(),
172            show_progress: true,
173            auto_save_history: true,
174            default_batch_size: 32,
175            default_seq_len: 512,
176        }
177    }
178}
179
180/// Session-level metrics.
181#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq)]
182pub struct SessionMetrics {
183    /// Total commands executed
184    pub total_commands: u64,
185    /// Successful commands
186    pub successful_commands: u64,
187    /// Total duration in milliseconds
188    pub total_duration_ms: u64,
189}
190
191impl SessionMetrics {
192    /// Get success rate as a percentage.
193    pub fn success_rate(&self) -> f64 {
194        if self.total_commands == 0 {
195            100.0
196        } else {
197            (self.successful_commands as f64 / self.total_commands as f64) * 100.0
198        }
199    }
200
201    /// Get average command duration in milliseconds.
202    pub fn avg_duration_ms(&self) -> f64 {
203        if self.total_commands == 0 {
204            0.0
205        } else {
206            self.total_duration_ms as f64 / self.total_commands as f64
207        }
208    }
209}
210
211#[cfg(test)]
212mod tests {
213    use super::*;
214
215    #[test]
216    fn test_session_state_model_management() {
217        let mut state = SessionState::new();
218
219        let model = LoadedModel {
220            id: "test/model".to_string(),
221            path: PathBuf::from("/tmp/model"),
222            architecture: "llama".to_string(),
223            parameters: 7_000_000_000,
224            layers: 32,
225            hidden_dim: 4096,
226            role: ModelRole::Teacher,
227        };
228
229        state.add_model("teacher".to_string(), model.clone());
230        assert_eq!(state.loaded_models().len(), 1);
231        assert!(state.get_model("teacher").is_some());
232
233        state.remove_model("teacher");
234        assert!(state.get_model("teacher").is_none());
235    }
236
237    #[test]
238    fn test_session_state_history() {
239        let mut state = SessionState::new();
240
241        state.add_to_history(HistoryEntry::new("fetch model", 100, true));
242        state.add_to_history(HistoryEntry::new("inspect layers", 50, true));
243
244        assert_eq!(state.history().len(), 2);
245        assert_eq!(state.history()[0].command, "fetch model");
246    }
247
248    #[test]
249    fn test_session_metrics() {
250        let mut state = SessionState::new();
251
252        state.record_command(100, true);
253        state.record_command(200, true);
254        state.record_command(150, false);
255
256        assert_eq!(state.metrics().total_commands, 3);
257        assert_eq!(state.metrics().successful_commands, 2);
258        assert!((state.metrics().success_rate() - 66.67).abs() < 1.0);
259    }
260
261    #[test]
262    fn test_session_state_serialization_roundtrip() {
263        let mut state = SessionState::new();
264        state.add_to_history(HistoryEntry::new("test", 100, true));
265        state.preferences_mut().default_batch_size = 64;
266
267        let json = serde_json::to_string(&state).expect("JSON serialization should succeed");
268        let restored: SessionState =
269            serde_json::from_str(&json).expect("JSON deserialization should succeed");
270
271        assert_eq!(state, restored);
272    }
273
274    #[test]
275    fn test_model_role_default() {
276        assert_eq!(ModelRole::default(), ModelRole::None);
277    }
278
279    #[test]
280    fn test_preferences_default_values() {
281        let prefs = Preferences::default();
282        assert_eq!(prefs.output_format, "table");
283        assert!(prefs.show_progress);
284        assert_eq!(prefs.default_batch_size, 32);
285    }
286
287    #[test]
288    fn test_session_metrics_success_rate_zero() {
289        let metrics = SessionMetrics::default();
290        assert_eq!(metrics.success_rate(), 100.0);
291    }
292
293    #[test]
294    fn test_session_metrics_avg_duration_zero() {
295        let metrics = SessionMetrics::default();
296        assert_eq!(metrics.avg_duration_ms(), 0.0);
297    }
298
299    #[test]
300    fn test_session_metrics_avg_duration() {
301        let mut state = SessionState::new();
302        state.record_command(100, true);
303        state.record_command(200, true);
304        assert_eq!(state.metrics().avg_duration_ms(), 150.0);
305    }
306
307    #[test]
308    fn test_history_entry_new() {
309        let entry = HistoryEntry::new("test command", 50, true);
310        assert_eq!(entry.command, "test command");
311        assert_eq!(entry.duration_ms, 50);
312        assert!(entry.success);
313        assert!(entry.timestamp > 0);
314    }
315
316    #[test]
317    fn test_loaded_model_equality() {
318        let model1 = LoadedModel {
319            id: "test".to_string(),
320            path: PathBuf::from("/tmp"),
321            architecture: "llama".to_string(),
322            parameters: 7_000_000_000,
323            layers: 32,
324            hidden_dim: 4096,
325            role: ModelRole::None,
326        };
327        let model2 = model1.clone();
328        assert_eq!(model1, model2);
329    }
330
331    #[test]
332    fn test_model_role_equality() {
333        assert_eq!(ModelRole::Teacher, ModelRole::Teacher);
334        assert_ne!(ModelRole::Teacher, ModelRole::Student);
335        assert_ne!(ModelRole::Student, ModelRole::None);
336    }
337
338    #[test]
339    fn test_session_state_save_load() {
340        use tempfile::TempDir;
341
342        let temp_dir = TempDir::new().expect("temp file creation should succeed");
343        let state_path = temp_dir.path().join("state.json");
344
345        let mut state = SessionState::new();
346        state.add_to_history(HistoryEntry::new("test", 100, true));
347        state.preferences_mut().default_batch_size = 128;
348
349        state.save(&state_path).expect("save should succeed");
350        let loaded = SessionState::load(&state_path).expect("load should succeed");
351
352        assert_eq!(state, loaded);
353    }
354
355    #[test]
356    fn test_session_state_load_invalid_json() {
357        use std::io::Write;
358        use tempfile::NamedTempFile;
359
360        let mut file = NamedTempFile::new().expect("temp file creation should succeed");
361        file.write_all(b"not valid json")
362            .expect("file write should succeed");
363
364        let result = SessionState::load(&file.path().to_path_buf());
365        assert!(result.is_err());
366    }
367
368    #[test]
369    fn test_preferences_all_fields() {
370        let prefs = Preferences::default();
371        assert_eq!(prefs.output_format, "table");
372        assert!(prefs.show_progress);
373        assert!(prefs.auto_save_history);
374        assert_eq!(prefs.default_batch_size, 32);
375        assert_eq!(prefs.default_seq_len, 512);
376    }
377
378    #[test]
379    fn test_session_state_remove_nonexistent() {
380        let mut state = SessionState::new();
381        let result = state.remove_model("nonexistent");
382        assert!(result.is_none());
383    }
384
385    #[test]
386    fn test_session_state_get_nonexistent() {
387        let state = SessionState::new();
388        assert!(state.get_model("nonexistent").is_none());
389    }
390
391    #[test]
392    fn test_session_metrics_fields() {
393        let metrics = SessionMetrics {
394            total_commands: 10,
395            successful_commands: 8,
396            total_duration_ms: 1000,
397        };
398        assert_eq!(metrics.success_rate(), 80.0);
399        assert_eq!(metrics.avg_duration_ms(), 100.0);
400    }
401}