1use crate::state::{HistoryEntry, LoadedModel, ModelRole, SessionState};
4use entrenar_common::{EntrenarError, Result};
5
6#[derive(Debug, Clone, PartialEq)]
8pub enum Command {
9 Fetch { model_id: String, role: ModelRole },
11 Inspect { target: InspectTarget },
13 Memory {
15 batch_size: Option<u32>,
16 seq_len: Option<usize>,
17 },
18 Set { key: String, value: String },
20 Distill { dry_run: bool },
22 Export { format: String, path: String },
24 History,
26 Help { topic: Option<String> },
28 Clear,
30 Quit,
32 Unknown { input: String },
34}
35
36#[derive(Debug, Clone, PartialEq)]
38pub enum InspectTarget {
39 Layers,
41 Memory,
43 All,
45 Model(String),
47}
48
49pub fn parse(input: &str) -> Result<Command> {
51 let input = input.trim();
52 if input.is_empty() {
53 return Ok(Command::Unknown {
54 input: String::new(),
55 });
56 }
57
58 let parts: Vec<&str> = input.split_whitespace().collect();
59 let cmd = parts[0].to_lowercase();
60 let args = &parts[1..];
61
62 match cmd.as_str() {
63 "fetch" | "download" => parse_fetch(args),
64 "inspect" | "show" => parse_inspect(args),
65 "memory" | "mem" => parse_memory(args),
66 "set" => parse_set(args),
67 "distill" | "train" => parse_distill(args),
68 "export" | "save" => parse_export(args),
69 "history" | "hist" => Ok(Command::History),
70 "help" | "?" => parse_help(args),
71 "clear" | "cls" => Ok(Command::Clear),
72 "quit" | "exit" | "q" => Ok(Command::Quit),
73 _ => Ok(Command::Unknown {
74 input: input.to_string(),
75 }),
76 }
77}
78
79fn parse_fetch(args: &[&str]) -> Result<Command> {
80 if args.is_empty() {
81 return Err(EntrenarError::ConfigValue {
82 field: "model_id".into(),
83 message: "No model ID provided".into(),
84 suggestion: "Usage: fetch <model_id> [--teacher|--student]".into(),
85 });
86 }
87
88 let model_id = args[0].to_string();
89 let role = if args.contains(&"--teacher") {
90 ModelRole::Teacher
91 } else if args.contains(&"--student") {
92 ModelRole::Student
93 } else {
94 ModelRole::None
95 };
96
97 Ok(Command::Fetch { model_id, role })
98}
99
100fn parse_inspect(args: &[&str]) -> Result<Command> {
101 let target = if args.is_empty() {
102 InspectTarget::All
103 } else {
104 match args[0].to_lowercase().as_str() {
105 "layers" | "layer" => InspectTarget::Layers,
106 "memory" | "mem" => InspectTarget::Memory,
107 "all" => InspectTarget::All,
108 name => InspectTarget::Model(name.to_string()),
109 }
110 };
111
112 Ok(Command::Inspect { target })
113}
114
115fn parse_memory(args: &[&str]) -> Result<Command> {
116 let mut batch_size = None;
117 let mut seq_len = None;
118
119 let mut i = 0;
120 while i < args.len() {
121 match args[i] {
122 "--batch" | "-b" if i + 1 < args.len() => {
123 batch_size = args[i + 1].parse().ok();
124 i += 2;
125 }
126 "--seq" | "-s" if i + 1 < args.len() => {
127 seq_len = args[i + 1].parse().ok();
128 i += 2;
129 }
130 _ => i += 1,
131 }
132 }
133
134 Ok(Command::Memory {
135 batch_size,
136 seq_len,
137 })
138}
139
140fn parse_set(args: &[&str]) -> Result<Command> {
141 if args.len() < 2 {
142 return Err(EntrenarError::ConfigValue {
143 field: "set".into(),
144 message: "Not enough arguments".into(),
145 suggestion: "Usage: set <key> <value>".into(),
146 });
147 }
148
149 Ok(Command::Set {
150 key: args[0].to_string(),
151 value: args[1..].join(" "),
152 })
153}
154
155fn parse_distill(args: &[&str]) -> Result<Command> {
156 let dry_run = args.contains(&"--dry-run") || args.contains(&"-n");
157 Ok(Command::Distill { dry_run })
158}
159
160fn parse_export(args: &[&str]) -> Result<Command> {
161 if args.len() < 2 {
162 return Err(EntrenarError::ConfigValue {
163 field: "export".into(),
164 message: "Not enough arguments".into(),
165 suggestion: "Usage: export <format> <path>".into(),
166 });
167 }
168
169 Ok(Command::Export {
170 format: args[0].to_string(),
171 path: args[1].to_string(),
172 })
173}
174
175fn parse_help(args: &[&str]) -> Result<Command> {
176 let topic = args.first().map(ToString::to_string);
177 Ok(Command::Help { topic })
178}
179
180pub fn execute(cmd: &Command, state: &mut SessionState) -> Result<String> {
182 let start = std::time::Instant::now();
183
184 let result = match cmd {
185 Command::Fetch { model_id, role } => execute_fetch(model_id, *role, state),
186 Command::Inspect { target } => execute_inspect(target, state),
187 Command::Memory {
188 batch_size,
189 seq_len,
190 } => execute_memory(*batch_size, *seq_len, state),
191 Command::Set { key, value } => execute_set(key, value, state),
192 Command::Distill { dry_run } => execute_distill(*dry_run, state),
193 Command::Export { format, path } => execute_export(format, path, state),
194 Command::History => execute_history(state),
195 Command::Help { topic } => execute_help(topic.as_deref()),
196 Command::Clear => Ok(String::new()),
197 Command::Quit => Ok("Goodbye!".to_string()),
198 Command::Unknown { input } => {
199 if input.is_empty() {
200 Ok(String::new())
201 } else {
202 Err(EntrenarError::ConfigValue {
203 field: "command".into(),
204 message: format!("Unknown command: {input}"),
205 suggestion: "Type 'help' for available commands".into(),
206 })
207 }
208 }
209 };
210
211 let duration_ms = start.elapsed().as_millis() as u64;
212 let success = result.is_ok();
213
214 if !matches!(
216 cmd,
217 Command::Help { .. }
218 | Command::History
219 | Command::Clear
220 | Command::Quit
221 | Command::Unknown { .. }
222 ) {
223 let cmd_str = format!("{cmd:?}");
224 state.add_to_history(HistoryEntry::new(cmd_str, duration_ms, success));
225 state.record_command(duration_ms, success);
226 }
227
228 result
229}
230
231fn execute_fetch(model_id: &str, role: ModelRole, state: &mut SessionState) -> Result<String> {
232 let model = LoadedModel {
234 id: model_id.to_string(),
235 path: std::path::PathBuf::from(format!("/tmp/models/{}", model_id.replace('/', "_"))),
236 architecture: detect_architecture(model_id),
237 parameters: estimate_params(model_id),
238 layers: estimate_layers(model_id),
239 hidden_dim: 4096,
240 role,
241 };
242
243 let name = if role == ModelRole::Teacher {
244 "teacher"
245 } else if role == ModelRole::Student {
246 "student"
247 } else {
248 model_id.split('/').next_back().unwrap_or(model_id)
249 };
250
251 state.add_model(name.to_string(), model.clone());
252
253 Ok(format!(
254 "✓ Fetched {}\n Architecture: {}\n Parameters: {:.1}B\n Layers: {}",
255 model_id,
256 model.architecture,
257 model.parameters as f64 / 1e9,
258 model.layers
259 ))
260}
261
262fn execute_inspect(target: &InspectTarget, state: &SessionState) -> Result<String> {
263 match target {
264 InspectTarget::All => {
265 if state.loaded_models().is_empty() {
266 return Ok("No models loaded. Use 'fetch <model_id>' to load a model.".to_string());
267 }
268
269 let mut output = String::from("Loaded Models:\n");
270 for (name, model) in state.loaded_models() {
271 output.push_str(&format!(
272 " {} ({}): {:.1}B params, {} layers\n",
273 name,
274 model.id,
275 model.parameters as f64 / 1e9,
276 model.layers
277 ));
278 }
279 Ok(output)
280 }
281 InspectTarget::Layers => {
282 let mut output = String::from("Layer Analysis:\n");
283 for (name, model) in state.loaded_models() {
284 output.push_str(&format!(
285 " {}: {} layers, hidden_dim={}\n",
286 name, model.layers, model.hidden_dim
287 ));
288 }
289 Ok(output)
290 }
291 InspectTarget::Memory => execute_memory(None, None, state),
292 InspectTarget::Model(name) => {
293 if let Some(model) = state.get_model(name) {
294 Ok(format!(
295 "Model: {}\n ID: {}\n Path: {}\n Architecture: {}\n Parameters: {:.1}B\n Layers: {}\n Hidden Dim: {}",
296 name, model.id, model.path.display(), model.architecture,
297 model.parameters as f64 / 1e9, model.layers, model.hidden_dim
298 ))
299 } else {
300 Err(EntrenarError::ModelNotFound { path: name.into() })
301 }
302 }
303 }
304}
305
306fn execute_memory(
307 batch_size: Option<u32>,
308 seq_len: Option<usize>,
309 state: &SessionState,
310) -> Result<String> {
311 let batch = batch_size.unwrap_or(state.preferences().default_batch_size);
312 let seq = seq_len.unwrap_or(state.preferences().default_seq_len);
313
314 let total_params: u64 = state.loaded_models().values().map(|m| m.parameters).sum();
315 let model_mem = total_params * 2; let activation_mem = u64::from(batch) * (seq as u64) * 4096 * 32 * 2;
317 let total = model_mem + activation_mem;
318
319 Ok(format!(
320 "Memory Estimate (batch={}, seq={}):\n Model: {:.1} GB\n Activations: {:.1} GB\n Total: {:.1} GB",
321 batch, seq,
322 model_mem as f64 / 1e9,
323 activation_mem as f64 / 1e9,
324 total as f64 / 1e9
325 ))
326}
327
328fn execute_set(key: &str, value: &str, state: &mut SessionState) -> Result<String> {
329 match key {
330 "batch_size" | "batch" => {
331 let v: u32 = value.parse().map_err(|_| EntrenarError::ConfigValue {
332 field: "batch_size".into(),
333 message: "Invalid number".into(),
334 suggestion: "Use a positive integer".into(),
335 })?;
336 state.preferences_mut().default_batch_size = v;
337 Ok(format!("Set batch_size = {v}"))
338 }
339 "seq_len" | "seq" => {
340 let v: usize = value.parse().map_err(|_| EntrenarError::ConfigValue {
341 field: "seq_len".into(),
342 message: "Invalid number".into(),
343 suggestion: "Use a positive integer".into(),
344 })?;
345 state.preferences_mut().default_seq_len = v;
346 Ok(format!("Set seq_len = {v}"))
347 }
348 _ => Err(EntrenarError::ConfigValue {
349 field: key.into(),
350 message: "Unknown setting".into(),
351 suggestion: "Available settings: batch_size, seq_len".into(),
352 }),
353 }
354}
355
356fn execute_distill(dry_run: bool, state: &SessionState) -> Result<String> {
357 let teacher = state
358 .loaded_models()
359 .values()
360 .find(|m| m.role == ModelRole::Teacher);
361 let student = state
362 .loaded_models()
363 .values()
364 .find(|m| m.role == ModelRole::Student);
365
366 if teacher.is_none() {
367 return Err(EntrenarError::ConfigValue {
368 field: "teacher".into(),
369 message: "No teacher model loaded".into(),
370 suggestion: "Use 'fetch <model_id> --teacher' to load a teacher model".into(),
371 });
372 }
373
374 if student.is_none() {
375 return Err(EntrenarError::ConfigValue {
376 field: "student".into(),
377 message: "No student model loaded".into(),
378 suggestion: "Use 'fetch <model_id> --student' to load a student model".into(),
379 });
380 }
381
382 if dry_run {
383 let teacher = teacher.expect("teacher validated non-None above");
384 let student = student.expect("student validated non-None above");
385 Ok(format!(
386 "Dry run configuration:\n Teacher: {} ({:.1}B)\n Student: {} ({:.1}B)\n Ready to train",
387 teacher.id, teacher.parameters as f64 / 1e9,
388 student.id, student.parameters as f64 / 1e9
389 ))
390 } else {
391 Ok("Training started... (simulated)".to_string())
392 }
393}
394
395fn execute_export(format: &str, path: &str, _state: &SessionState) -> Result<String> {
396 Ok(format!("Exported to {path} in {format} format"))
397}
398
399fn execute_history(state: &SessionState) -> Result<String> {
400 if state.history().is_empty() {
401 return Ok("No command history.".to_string());
402 }
403
404 let mut output = String::from("Command History:\n");
405 for (i, entry) in state.history().iter().enumerate() {
406 let status = if entry.success { "✓" } else { "✗" };
407 output.push_str(&format!(
408 " {}. {} {} ({}ms)\n",
409 i + 1,
410 status,
411 entry.command,
412 entry.duration_ms
413 ));
414 }
415 Ok(output)
416}
417
418fn execute_help(topic: Option<&str>) -> Result<String> {
419 match topic {
420 Some("fetch") => Ok(
421 "fetch <model_id> [--teacher|--student]\n Download a model from HuggingFace"
422 .to_string(),
423 ),
424 Some("inspect") => {
425 Ok("inspect [layers|memory|all|<model>]\n Inspect loaded models".to_string())
426 }
427 Some("memory") => {
428 Ok("memory [--batch <n>] [--seq <n>]\n Estimate memory requirements".to_string())
429 }
430 Some("distill") => Ok("distill [--dry-run]\n Run distillation training".to_string()),
431 _ => Ok("Available commands:
432 fetch <model> Download model from HuggingFace
433 inspect [target] Inspect loaded models
434 memory Estimate memory requirements
435 set <key> <value> Configure settings
436 distill Run distillation
437 export <fmt> <path> Export model
438 history Show command history
439 help [topic] Show help
440 quit Exit shell"
441 .to_string()),
442 }
443}
444
445const ARCH_PATTERNS: &[(&[&str], &str)] = &[
451 (&["qwen"], "qwen"),
452 (&["phi"], "phi"),
453 (&["falcon"], "falcon"),
454 (&["mistral", "mixtral"], "mistral"),
455 (&["llama"], "llama"),
456 (&["bert"], "bert"),
457 (&["gpt"], "gpt"),
458];
459
460fn detect_architecture(model_id: &str) -> String {
461 let lower = model_id.to_lowercase();
462 for (patterns, arch) in ARCH_PATTERNS {
463 if patterns.iter().any(|p| lower.contains(p)) {
464 return (*arch).to_string();
465 }
466 }
467 eprintln!(
468 "Warning: could not detect architecture from model ID '{model_id}', \
469 defaulting to 'unknown' (use tensor-based detection for accuracy)"
470 );
471 "unknown".to_string()
472}
473
474fn estimate_params(model_id: &str) -> u64 {
475 let lower = model_id.to_lowercase();
476 if lower.contains("70b") {
477 70_000_000_000
478 } else if lower.contains("13b") {
479 13_000_000_000
480 } else if lower.contains("7b") {
481 7_000_000_000
482 } else if lower.contains("1.1b") || lower.contains("1b") {
483 1_100_000_000
484 } else if lower.contains("base") {
485 350_000_000
486 } else {
487 1_000_000_000
488 }
489}
490
491fn estimate_layers(model_id: &str) -> u32 {
492 let lower = model_id.to_lowercase();
493 if lower.contains("70b") {
494 80
495 } else if lower.contains("13b") {
496 40
497 } else if lower.contains("7b") {
498 32
499 } else if lower.contains("base") {
500 12
501 } else {
502 24
503 }
504}
505
506#[cfg(test)]
507mod tests {
508 use super::*;
509
510 #[test]
511 fn test_parse_fetch() {
512 let cmd = parse("fetch meta-llama/Llama-2-7b --teacher").expect("parsing should succeed");
513 assert!(matches!(
514 cmd,
515 Command::Fetch {
516 role: ModelRole::Teacher,
517 ..
518 }
519 ));
520
521 let cmd =
522 parse("fetch TinyLlama/TinyLlama-1.1B --student").expect("parsing should succeed");
523 assert!(matches!(
524 cmd,
525 Command::Fetch {
526 role: ModelRole::Student,
527 ..
528 }
529 ));
530 }
531
532 #[test]
533 fn test_parse_inspect() {
534 assert!(matches!(
535 parse("inspect").expect("parsing should succeed"),
536 Command::Inspect {
537 target: InspectTarget::All
538 }
539 ));
540 assert!(matches!(
541 parse("inspect layers").expect("parsing should succeed"),
542 Command::Inspect {
543 target: InspectTarget::Layers
544 }
545 ));
546 }
547
548 #[test]
549 fn test_parse_memory() {
550 let cmd = parse("memory --batch 64 --seq 1024").expect("parsing should succeed");
551 if let Command::Memory {
552 batch_size,
553 seq_len,
554 } = cmd
555 {
556 assert_eq!(batch_size, Some(64));
557 assert_eq!(seq_len, Some(1024));
558 } else {
559 panic!("Expected Memory command");
560 }
561 }
562
563 #[test]
564 fn test_parse_quit_variants() {
565 assert!(matches!(
566 parse("quit").expect("parsing should succeed"),
567 Command::Quit
568 ));
569 assert!(matches!(
570 parse("exit").expect("parsing should succeed"),
571 Command::Quit
572 ));
573 assert!(matches!(
574 parse("q").expect("parsing should succeed"),
575 Command::Quit
576 ));
577 }
578
579 #[test]
580 fn test_execute_fetch() {
581 let mut state = SessionState::new();
582 let result = execute_fetch("meta-llama/Llama-2-7b", ModelRole::Teacher, &mut state);
583
584 assert!(result.is_ok());
585 assert!(state.get_model("teacher").is_some());
586 }
587
588 #[test]
589 fn test_execute_set() {
590 let mut state = SessionState::new();
591
592 execute_set("batch_size", "64", &mut state).expect("operation should succeed");
593 assert_eq!(state.preferences().default_batch_size, 64);
594
595 execute_set("seq_len", "1024", &mut state).expect("operation should succeed");
596 assert_eq!(state.preferences().default_seq_len, 1024);
597 }
598
599 #[test]
600 fn test_unknown_command() {
601 let cmd = parse("foobar").expect("parsing should succeed");
602 assert!(matches!(cmd, Command::Unknown { .. }));
603 }
604
605 #[test]
606 fn test_parse_fetch_missing_model() {
607 let result = parse("fetch");
608 assert!(result.is_err());
609 }
610
611 #[test]
612 fn test_parse_set_not_enough_args() {
613 let result = parse("set batch_size");
614 assert!(result.is_err());
615 }
616
617 #[test]
618 fn test_parse_export_not_enough_args() {
619 let result = parse("export safetensors");
620 assert!(result.is_err());
621 }
622
623 #[test]
624 fn test_parse_export_valid() {
625 let cmd = parse("export safetensors /tmp/model.st").expect("parsing should succeed");
626 if let Command::Export { format, path } = cmd {
627 assert_eq!(format, "safetensors");
628 assert_eq!(path, "/tmp/model.st");
629 } else {
630 panic!("Expected Export command");
631 }
632 }
633
634 #[test]
635 fn test_parse_help_with_topic() {
636 let cmd = parse("help fetch").expect("parsing should succeed");
637 if let Command::Help { topic } = cmd {
638 assert_eq!(topic, Some("fetch".to_string()));
639 } else {
640 panic!("Expected Help command");
641 }
642 }
643
644 #[test]
645 fn test_parse_distill_dry_run() {
646 let cmd = parse("distill --dry-run").expect("parsing should succeed");
647 if let Command::Distill { dry_run } = cmd {
648 assert!(dry_run);
649 } else {
650 panic!("Expected Distill command");
651 }
652 }
653
654 #[test]
655 fn test_parse_distill_short_flag() {
656 let cmd = parse("distill -n").expect("parsing should succeed");
657 if let Command::Distill { dry_run } = cmd {
658 assert!(dry_run);
659 } else {
660 panic!("Expected Distill command");
661 }
662 }
663
664 #[test]
665 fn test_parse_inspect_model() {
666 let cmd = parse("inspect teacher").expect("parsing should succeed");
667 if let Command::Inspect { target } = cmd {
668 assert_eq!(target, InspectTarget::Model("teacher".to_string()));
669 } else {
670 panic!("Expected Inspect command");
671 }
672 }
673
674 #[test]
675 fn test_parse_inspect_memory() {
676 let cmd = parse("inspect memory").expect("parsing should succeed");
677 assert!(matches!(
678 cmd,
679 Command::Inspect {
680 target: InspectTarget::Memory
681 }
682 ));
683 }
684
685 #[test]
686 fn test_parse_command_aliases() {
687 assert!(matches!(
689 parse("download model").expect("parsing should succeed"),
690 Command::Fetch { .. }
691 ));
692 assert!(matches!(
694 parse("show layers").expect("parsing should succeed"),
695 Command::Inspect { .. }
696 ));
697 assert!(matches!(
699 parse("mem").expect("parsing should succeed"),
700 Command::Memory { .. }
701 ));
702 assert!(matches!(
704 parse("train").expect("parsing should succeed"),
705 Command::Distill { .. }
706 ));
707 assert!(matches!(
709 parse("save gguf /tmp/out").expect("parsing should succeed"),
710 Command::Export { .. }
711 ));
712 assert!(matches!(
714 parse("cls").expect("parsing should succeed"),
715 Command::Clear
716 ));
717 assert!(matches!(
719 parse("?").expect("parsing should succeed"),
720 Command::Help { .. }
721 ));
722 assert!(matches!(
724 parse("hist").expect("parsing should succeed"),
725 Command::History
726 ));
727 }
728
729 #[test]
730 fn test_execute_inspect_no_models() {
731 let state = SessionState::new();
732 let result = execute_inspect(&InspectTarget::All, &state);
733 assert!(result
734 .expect("load should succeed")
735 .contains("No models loaded"));
736 }
737
738 #[test]
739 fn test_execute_inspect_layers() {
740 let mut state = SessionState::new();
741 let model = LoadedModel {
742 id: "test".to_string(),
743 path: std::path::PathBuf::from("/tmp"),
744 architecture: "llama".to_string(),
745 parameters: 7_000_000_000,
746 layers: 32,
747 hidden_dim: 4096,
748 role: ModelRole::None,
749 };
750 state.add_model("test".to_string(), model);
751
752 let result =
753 execute_inspect(&InspectTarget::Layers, &state).expect("operation should succeed");
754 assert!(result.contains("32 layers"));
755 }
756
757 #[test]
758 fn test_execute_inspect_model_not_found() {
759 let state = SessionState::new();
760 let result = execute_inspect(&InspectTarget::Model("unknown".to_string()), &state);
761 assert!(result.is_err());
762 }
763
764 #[test]
765 fn test_execute_history_empty() {
766 let state = SessionState::new();
767 let result = execute_history(&state).expect("operation should succeed");
768 assert!(result.contains("No command history"));
769 }
770
771 #[test]
772 fn test_execute_help_topics() {
773 let fetch_help = execute_help(Some("fetch")).expect("operation should succeed");
774 assert!(fetch_help.contains("Download"));
775
776 let inspect_help = execute_help(Some("inspect")).expect("operation should succeed");
777 assert!(inspect_help.contains("Inspect"));
778
779 let memory_help = execute_help(Some("memory")).expect("operation should succeed");
780 assert!(memory_help.contains("memory"));
781
782 let distill_help = execute_help(Some("distill")).expect("operation should succeed");
783 assert!(distill_help.contains("distill"));
784
785 let general_help = execute_help(None).expect("operation should succeed");
786 assert!(general_help.contains("Available commands"));
787 }
788
789 #[test]
790 fn test_detect_architecture_variants() {
791 assert_eq!(detect_architecture("meta-llama/Llama-2-7b"), "llama");
792 assert_eq!(detect_architecture("bert-base-uncased"), "bert");
793 assert_eq!(detect_architecture("openai-gpt"), "gpt");
794 assert_eq!(detect_architecture("mistralai/Mistral-7B"), "mistral");
795 assert_eq!(detect_architecture("Qwen/Qwen2.5-Coder-0.5B"), "qwen");
796 assert_eq!(detect_architecture("microsoft/phi-2"), "phi");
797 assert_eq!(detect_architecture("tiiuae/falcon-7b"), "falcon");
798 assert_eq!(detect_architecture("mistralai/Mixtral-8x7B"), "mistral");
799 assert_eq!(detect_architecture("custom-model"), "unknown");
800 }
801
802 #[test]
807 fn test_falsify_n04_mistral_before_llama_in_ambiguous_id() {
808 assert_eq!(
812 detect_architecture("mistral-llama-variant"),
813 "mistral",
814 "Mistral should take priority over LLaMA in ambiguous IDs"
815 );
816 }
817
818 #[test]
819 fn test_falsify_n04_unknown_is_explicit() {
820 let result = detect_architecture("totally-custom-model-v3");
823 assert_eq!(result, "unknown");
824 assert!(!result.is_empty());
825 }
826
827 #[test]
828 fn test_estimate_params_variants() {
829 assert_eq!(estimate_params("model-70b"), 70_000_000_000);
830 assert_eq!(estimate_params("model-13b"), 13_000_000_000);
831 assert_eq!(estimate_params("model-7b"), 7_000_000_000);
832 assert_eq!(estimate_params("model-1.1b"), 1_100_000_000);
833 assert_eq!(estimate_params("bert-base"), 350_000_000);
834 }
835
836 #[test]
837 fn test_estimate_layers_variants() {
838 assert_eq!(estimate_layers("model-70b"), 80);
839 assert_eq!(estimate_layers("model-13b"), 40);
840 assert_eq!(estimate_layers("model-7b"), 32);
841 assert_eq!(estimate_layers("bert-base"), 12);
842 }
843
844 #[test]
845 fn test_execute_set_invalid_number() {
846 let mut state = SessionState::new();
847 let result = execute_set("batch_size", "not_a_number", &mut state);
848 assert!(result.is_err());
849 }
850
851 #[test]
852 fn test_execute_set_unknown_key() {
853 let mut state = SessionState::new();
854 let result = execute_set("unknown_setting", "value", &mut state);
855 assert!(result.is_err());
856 }
857
858 #[test]
859 fn test_execute_distill_no_teacher() {
860 let state = SessionState::new();
861 let result = execute_distill(true, &state);
862 assert!(result.is_err());
863 }
864
865 #[test]
866 fn test_execute_distill_no_student() {
867 let mut state = SessionState::new();
868 let model = LoadedModel {
869 id: "teacher".to_string(),
870 path: std::path::PathBuf::from("/tmp"),
871 architecture: "llama".to_string(),
872 parameters: 7_000_000_000,
873 layers: 32,
874 hidden_dim: 4096,
875 role: ModelRole::Teacher,
876 };
877 state.add_model("teacher".to_string(), model);
878
879 let result = execute_distill(true, &state);
880 assert!(result.is_err());
881 }
882
883 #[test]
884 fn test_execute_distill_success() {
885 let mut state = SessionState::new();
886
887 let teacher = LoadedModel {
888 id: "teacher/model".to_string(),
889 path: std::path::PathBuf::from("/tmp/t"),
890 architecture: "llama".to_string(),
891 parameters: 7_000_000_000,
892 layers: 32,
893 hidden_dim: 4096,
894 role: ModelRole::Teacher,
895 };
896 state.add_model("teacher".to_string(), teacher);
897
898 let student = LoadedModel {
899 id: "student/model".to_string(),
900 path: std::path::PathBuf::from("/tmp/s"),
901 architecture: "llama".to_string(),
902 parameters: 1_000_000_000,
903 layers: 12,
904 hidden_dim: 2048,
905 role: ModelRole::Student,
906 };
907 state.add_model("student".to_string(), student);
908
909 let result = execute_distill(true, &state).expect("operation should succeed");
910 assert!(result.contains("Dry run"));
911 }
912
913 #[test]
914 fn test_execute_export() {
915 let state = SessionState::new();
916 let result = execute_export("safetensors", "/tmp/model.st", &state)
917 .expect("operation should succeed");
918 assert!(result.contains("Exported"));
919 }
920
921 #[test]
922 fn test_parse_empty_input() {
923 let cmd = parse("").expect("parsing should succeed");
924 assert!(matches!(cmd, Command::Unknown { .. }));
925 }
926
927 #[test]
928 fn test_execute_memory_with_args() {
929 let state = SessionState::new();
930 let result =
931 execute_memory(Some(64), Some(1024), &state).expect("operation should succeed");
932 assert!(result.contains("batch=64"));
933 assert!(result.contains("seq=1024"));
934 }
935
936 #[test]
937 fn test_command_enum_equality() {
938 assert_eq!(Command::Quit, Command::Quit);
939 assert_eq!(Command::Clear, Command::Clear);
940 assert_eq!(Command::History, Command::History);
941 }
942}