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> {
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 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#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
159pub struct LoadedModel {
160 pub id: String,
162 pub path: PathBuf,
164 pub architecture: String,
166 pub parameters: u64,
168 pub layers: u32,
170 pub hidden_dim: u32,
172 pub role: ModelRole,
174}
175
176#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
178pub enum ModelRole {
179 Teacher,
181 Student,
183 #[default]
185 None,
186}
187
188#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
190pub struct HistoryEntry {
191 pub command: String,
193 pub timestamp: u64,
195 pub duration_ms: u64,
197 pub success: bool,
199}
200
201impl HistoryEntry {
202 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#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
218pub struct Preferences {
219 pub output_format: String,
221 pub show_progress: bool,
223 pub auto_save_history: bool,
225 pub default_batch_size: u32,
227 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#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq)]
245pub struct SessionMetrics {
246 pub total_commands: u64,
248 pub successful_commands: u64,
250 pub total_duration_ms: u64,
252}
253
254impl SessionMetrics {
255 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 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}