1use serde::{Deserialize, Serialize};
4use std::collections::HashMap;
5use std::path::PathBuf;
6
7#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq)]
9pub struct SessionState {
10 models: HashMap<String, LoadedModel>,
12 history: Vec<HistoryEntry>,
14 preferences: Preferences,
16 metrics: SessionMetrics,
18}
19
20impl SessionState {
21 pub fn new() -> Self {
23 Self::default()
24 }
25
26 pub fn loaded_models(&self) -> &HashMap<String, LoadedModel> {
28 &self.models
29 }
30
31 pub fn history(&self) -> &[HistoryEntry] {
33 &self.history
34 }
35
36 pub fn add_model(&mut self, name: String, model: LoadedModel) {
38 self.models.insert(name, model);
39 }
40
41 pub fn remove_model(&mut self, name: &str) -> Option<LoadedModel> {
43 self.models.remove(name)
44 }
45
46 pub fn get_model(&self, name: &str) -> Option<&LoadedModel> {
48 self.models.get(name)
49 }
50
51 pub fn add_to_history(&mut self, entry: HistoryEntry) {
53 self.history.push(entry);
54 }
55
56 pub fn preferences_mut(&mut self) -> &mut Preferences {
58 &mut self.preferences
59 }
60
61 pub fn preferences(&self) -> &Preferences {
63 &self.preferences
64 }
65
66 pub fn metrics(&self) -> &SessionMetrics {
68 &self.metrics
69 }
70
71 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 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 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#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
96pub struct LoadedModel {
97 pub id: String,
99 pub path: PathBuf,
101 pub architecture: String,
103 pub parameters: u64,
105 pub layers: u32,
107 pub hidden_dim: u32,
109 pub role: ModelRole,
111}
112
113#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
115pub enum ModelRole {
116 Teacher,
118 Student,
120 #[default]
122 None,
123}
124
125#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
127pub struct HistoryEntry {
128 pub command: String,
130 pub timestamp: u64,
132 pub duration_ms: u64,
134 pub success: bool,
136}
137
138impl HistoryEntry {
139 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#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
155pub struct Preferences {
156 pub output_format: String,
158 pub show_progress: bool,
160 pub auto_save_history: bool,
162 pub default_batch_size: u32,
164 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#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq)]
182pub struct SessionMetrics {
183 pub total_commands: u64,
185 pub successful_commands: u64,
187 pub total_duration_ms: u64,
189}
190
191impl SessionMetrics {
192 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 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}