1use crate::state::{HistoryEntry, 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, .. } => execute_fetch(model_id),
189 Command::Inspect { target } => execute_inspect(target, state),
190 Command::Memory {
191 batch_size,
192 seq_len,
193 } => execute_memory(*batch_size, *seq_len, state),
194 Command::Set { key, value } => execute_set(key, value, state),
195 Command::Distill { dry_run } => execute_distill(*dry_run, state),
196 Command::Export { format, path } => execute_export(format, path, state),
197 Command::History => execute_history(state),
198 Command::Help { topic } => execute_help(topic.as_deref()),
199 Command::Clear => Ok(String::new()),
200 Command::Quit => Ok("Goodbye!".to_string()),
201 Command::Unknown { input } => {
202 if input.is_empty() {
203 Ok(String::new())
204 } else {
205 Err(EntrenarError::ConfigValue {
206 field: "command".into(),
207 message: format!("Unknown command: {input}"),
208 suggestion: "Type 'help' for available commands".into(),
209 })
210 }
211 }
212 };
213
214 let duration_ms = start.elapsed().as_millis() as u64;
215 let success = result.is_ok();
216
217 if !matches!(
219 cmd,
220 Command::Help { .. }
221 | Command::History
222 | Command::Clear
223 | Command::Quit
224 | Command::Unknown { .. }
225 ) {
226 let cmd_str = format!("{cmd:?}");
227 state.add_to_history(HistoryEntry::new(cmd_str, duration_ms, success));
228 state.record_command(duration_ms, success);
229 }
230
231 result
232}
233
234fn execute_fetch(model_id: &str) -> Result<String> {
235 Err(EntrenarError::ConfigValue {
263 field: "fetch".into(),
264 message: format!(
265 "cannot fetch `{model_id}`: this shell has no HuggingFace client, so it \
266 downloads nothing. It previously reported success for any string at all, \
267 with a parameter count and layer count string-matched out of the model \
268 ID itself"
269 ),
270 suggestion: "Download with `apr pull <model_id>` or `apr import hf://<model_id>`, \
271 then read the real file with `apr inspect` / `apr tensors`. \
272 Tracked in #2519."
273 .into(),
274 })
275}
276
277fn execute_inspect(target: &InspectTarget, state: &SessionState) -> Result<String> {
278 match target {
279 InspectTarget::All => {
280 if state.loaded_models().is_empty() {
281 return Ok("No models loaded. Use 'fetch <model_id>' to load a model.".to_string());
282 }
283
284 let mut output = String::from("Loaded Models:\n");
285 for (name, model) in state.loaded_models() {
286 output.push_str(&format!(
287 " {} ({}): {:.1}B params, {} layers\n",
288 name,
289 model.id,
290 model.parameters as f64 / 1e9,
291 model.layers
292 ));
293 }
294 Ok(output)
295 }
296 InspectTarget::Layers => {
297 let mut output = String::from("Layer Analysis:\n");
298 for (name, model) in state.loaded_models() {
299 output.push_str(&format!(
300 " {}: {} layers, hidden_dim={}\n",
301 name, model.layers, model.hidden_dim
302 ));
303 }
304 Ok(output)
305 }
306 InspectTarget::Memory => execute_memory(None, None, state),
307 InspectTarget::Model(name) => {
308 if let Some(model) = state.get_model(name) {
309 Ok(format!(
310 "Model: {}\n ID: {}\n Path: {}\n Architecture: {}\n Parameters: {:.1}B\n Layers: {}\n Hidden Dim: {}",
311 name, model.id, model.path.display(), model.architecture,
312 model.parameters as f64 / 1e9, model.layers, model.hidden_dim
313 ))
314 } else {
315 Err(EntrenarError::ModelNotFound { path: name.into() })
316 }
317 }
318 }
319}
320
321fn execute_memory(
322 batch_size: Option<u32>,
323 seq_len: Option<usize>,
324 state: &SessionState,
325) -> Result<String> {
326 let batch = batch_size.unwrap_or(state.preferences().default_batch_size);
327 let seq = seq_len.unwrap_or(state.preferences().default_seq_len);
328
329 let total_params: u64 = state.loaded_models().values().map(|m| m.parameters).sum();
330 let model_mem = total_params * 2; let activation_mem = u64::from(batch) * (seq as u64) * 4096 * 32 * 2;
332 let total = model_mem + activation_mem;
333
334 Ok(format!(
335 "Memory Estimate (batch={}, seq={}):\n Model: {:.1} GB\n Activations: {:.1} GB\n Total: {:.1} GB",
336 batch, seq,
337 model_mem as f64 / 1e9,
338 activation_mem as f64 / 1e9,
339 total as f64 / 1e9
340 ))
341}
342
343fn execute_set(key: &str, value: &str, state: &mut SessionState) -> Result<String> {
344 match key {
345 "batch_size" | "batch" => {
346 let v: u32 = value.parse().map_err(|_| EntrenarError::ConfigValue {
347 field: "batch_size".into(),
348 message: "Invalid number".into(),
349 suggestion: "Use a positive integer".into(),
350 })?;
351 state.preferences_mut().default_batch_size = v;
352 Ok(format!("Set batch_size = {v}"))
353 }
354 "seq_len" | "seq" => {
355 let v: usize = value.parse().map_err(|_| EntrenarError::ConfigValue {
356 field: "seq_len".into(),
357 message: "Invalid number".into(),
358 suggestion: "Use a positive integer".into(),
359 })?;
360 state.preferences_mut().default_seq_len = v;
361 Ok(format!("Set seq_len = {v}"))
362 }
363 _ => Err(EntrenarError::ConfigValue {
364 field: key.into(),
365 message: "Unknown setting".into(),
366 suggestion: "Available settings: batch_size, seq_len".into(),
367 }),
368 }
369}
370
371fn execute_distill(dry_run: bool, state: &SessionState) -> Result<String> {
372 let teacher = state
373 .loaded_models()
374 .values()
375 .find(|m| m.role == ModelRole::Teacher);
376 let student = state
377 .loaded_models()
378 .values()
379 .find(|m| m.role == ModelRole::Student);
380
381 if teacher.is_none() {
382 return Err(EntrenarError::ConfigValue {
383 field: "teacher".into(),
384 message: "No teacher model loaded".into(),
385 suggestion: "Use 'fetch <model_id> --teacher' to load a teacher model".into(),
386 });
387 }
388
389 if student.is_none() {
390 return Err(EntrenarError::ConfigValue {
391 field: "student".into(),
392 message: "No student model loaded".into(),
393 suggestion: "Use 'fetch <model_id> --student' to load a student model".into(),
394 });
395 }
396
397 if dry_run {
398 let teacher = teacher.expect("teacher validated non-None above");
399 let student = student.expect("student validated non-None above");
400 Ok(format!(
401 "Dry run configuration:\n Teacher: {} ({:.1}B)\n Student: {} ({:.1}B)\n Ready to train",
402 teacher.id, teacher.parameters as f64 / 1e9,
403 student.id, student.parameters as f64 / 1e9
404 ))
405 } else {
406 Err(EntrenarError::ConfigValue {
415 field: "distill".into(),
416 message: format!(
417 "cannot distill `{}` into `{}`: this shell has no training loop, \
422 no optimizer and no dataset. It previously reported that \
423 training had begun, and exited 0 without training anything",
424 teacher.map_or("<teacher>", |t| t.id.as_str()),
425 student.map_or("<student>", |s| s.id.as_str()),
426 ),
427 suggestion: "Run the real trainer: `apr distill` / `apr finetune`. \
428 `distill --dry-run` here still prints the configuration. \
429 Tracked in #2519."
430 .into(),
431 })
432 }
433}
434
435fn execute_export(format: &str, path: &str, _state: &SessionState) -> Result<String> {
436 Err(EntrenarError::ConfigValue {
444 field: "export".into(),
445 message: format!("export is not implemented in this shell ({format} -> {path})"),
446 suggestion: "Use `apr export` for real format conversion.".into(),
447 })
448}
449
450fn execute_history(state: &SessionState) -> Result<String> {
451 if state.history().is_empty() {
452 return Ok("No command history.".to_string());
453 }
454
455 let mut output = String::from("Command History:\n");
456 for (i, entry) in state.history().iter().enumerate() {
457 let status = if entry.success { "✓" } else { "✗" };
458 output.push_str(&format!(
459 " {}. {} {} ({}ms)\n",
460 i + 1,
461 status,
462 entry.command,
463 entry.duration_ms
464 ));
465 }
466 Ok(output)
467}
468
469fn execute_help(topic: Option<&str>) -> Result<String> {
470 match topic {
471 Some("fetch") => Ok(
472 "fetch <model_id> [--teacher|--student]\n Download a model from HuggingFace"
473 .to_string(),
474 ),
475 Some("inspect") => {
476 Ok("inspect [layers|memory|all|<model>]\n Inspect loaded models".to_string())
477 }
478 Some("memory") => {
479 Ok("memory [--batch <n>] [--seq <n>]\n Estimate memory requirements".to_string())
480 }
481 Some("distill") => Ok("distill [--dry-run]\n Run distillation training".to_string()),
482 _ => Ok("Available commands:
483 fetch <model> Download model from HuggingFace
484 inspect [target] Inspect loaded models
485 memory Estimate memory requirements
486 set <key> <value> Configure settings
487 distill Run distillation
488 export <fmt> <path> Export model
489 history Show command history
490 clear Clear the screen
491 help [topic] Show help
492 quit Exit shell"
493 .to_string()),
494 }
495}
496
497#[cfg(test)]
508const ARCH_PATTERNS: &[(&[&str], &str)] = &[
509 (&["qwen"], "qwen"),
510 (&["phi"], "phi"),
511 (&["falcon"], "falcon"),
512 (&["mistral", "mixtral"], "mistral"),
513 (&["llama"], "llama"),
514 (&["bert"], "bert"),
515 (&["gpt"], "gpt"),
516];
517
518#[cfg(test)]
519fn detect_architecture(model_id: &str) -> String {
520 let lower = model_id.to_lowercase();
521 for (patterns, arch) in ARCH_PATTERNS {
522 if patterns.iter().any(|p| lower.contains(p)) {
523 return (*arch).to_string();
524 }
525 }
526 eprintln!(
527 "Warning: could not detect architecture from model ID '{model_id}', \
528 defaulting to 'unknown' (use tensor-based detection for accuracy)"
529 );
530 "unknown".to_string()
531}
532
533#[cfg(test)]
534fn estimate_params(model_id: &str) -> u64 {
535 let lower = model_id.to_lowercase();
536 if lower.contains("70b") {
537 70_000_000_000
538 } else if lower.contains("13b") {
539 13_000_000_000
540 } else if lower.contains("7b") {
541 7_000_000_000
542 } else if lower.contains("1.1b") || lower.contains("1b") {
543 1_100_000_000
544 } else if lower.contains("base") {
545 350_000_000
546 } else {
547 1_000_000_000
548 }
549}
550
551#[cfg(test)]
552fn estimate_layers(model_id: &str) -> u32 {
553 let lower = model_id.to_lowercase();
554 if lower.contains("70b") {
555 80
556 } else if lower.contains("13b") {
557 40
558 } else if lower.contains("7b") {
559 32
560 } else if lower.contains("base") {
561 12
562 } else {
563 24
564 }
565}
566
567#[cfg(test)]
568mod tests {
569 use super::*;
570 use crate::state::LoadedModel;
572
573 #[test]
574 fn test_parse_fetch() {
575 let cmd = parse("fetch meta-llama/Llama-2-7b --teacher").expect("parsing should succeed");
576 assert!(matches!(
577 cmd,
578 Command::Fetch {
579 role: ModelRole::Teacher,
580 ..
581 }
582 ));
583
584 let cmd =
585 parse("fetch TinyLlama/TinyLlama-1.1B --student").expect("parsing should succeed");
586 assert!(matches!(
587 cmd,
588 Command::Fetch {
589 role: ModelRole::Student,
590 ..
591 }
592 ));
593 }
594
595 #[test]
596 fn test_parse_inspect() {
597 assert!(matches!(
598 parse("inspect").expect("parsing should succeed"),
599 Command::Inspect {
600 target: InspectTarget::All
601 }
602 ));
603 assert!(matches!(
604 parse("inspect layers").expect("parsing should succeed"),
605 Command::Inspect {
606 target: InspectTarget::Layers
607 }
608 ));
609 }
610
611 #[test]
612 fn test_parse_memory() {
613 let cmd = parse("memory --batch 64 --seq 1024").expect("parsing should succeed");
614 if let Command::Memory {
615 batch_size,
616 seq_len,
617 } = cmd
618 {
619 assert_eq!(batch_size, Some(64));
620 assert_eq!(seq_len, Some(1024));
621 } else {
622 panic!("Expected Memory command");
623 }
624 }
625
626 #[test]
627 fn test_parse_quit_variants() {
628 assert!(matches!(
629 parse("quit").expect("parsing should succeed"),
630 Command::Quit
631 ));
632 assert!(matches!(
633 parse("exit").expect("parsing should succeed"),
634 Command::Quit
635 ));
636 assert!(matches!(
637 parse("q").expect("parsing should succeed"),
638 Command::Quit
639 ));
640 }
641
642 #[test]
646 fn test_execute_fetch_refuses_and_loads_nothing() {
647 let mut state = SessionState::new();
648 let err = execute(
649 &Command::Fetch {
650 model_id: "meta-llama/Llama-2-7b".to_string(),
651 role: ModelRole::Teacher,
652 },
653 &mut state,
654 )
655 .expect_err("fetch must not claim to have downloaded a model");
656
657 assert!(format!("{err}").contains("no HuggingFace client"));
658 assert!(state.get_model("teacher").is_none());
659 assert!(state.loaded_models().is_empty());
660 }
661
662 #[test]
663 fn test_execute_set() {
664 let mut state = SessionState::new();
665
666 execute_set("batch_size", "64", &mut state).expect("operation should succeed");
667 assert_eq!(state.preferences().default_batch_size, 64);
668
669 execute_set("seq_len", "1024", &mut state).expect("operation should succeed");
670 assert_eq!(state.preferences().default_seq_len, 1024);
671 }
672
673 #[test]
674 fn test_unknown_command() {
675 let cmd = parse("foobar").expect("parsing should succeed");
676 assert!(matches!(cmd, Command::Unknown { .. }));
677 }
678
679 #[test]
680 fn test_parse_fetch_missing_model() {
681 let result = parse("fetch");
682 assert!(result.is_err());
683 }
684
685 #[test]
686 fn test_parse_set_not_enough_args() {
687 let result = parse("set batch_size");
688 assert!(result.is_err());
689 }
690
691 #[test]
692 fn test_parse_export_not_enough_args() {
693 let result = parse("export safetensors");
694 assert!(result.is_err());
695 }
696
697 #[test]
698 fn test_parse_export_valid() {
699 let cmd = parse("export safetensors /tmp/model.st").expect("parsing should succeed");
700 if let Command::Export { format, path } = cmd {
701 assert_eq!(format, "safetensors");
702 assert_eq!(path, "/tmp/model.st");
703 } else {
704 panic!("Expected Export command");
705 }
706 }
707
708 #[test]
709 fn test_parse_help_with_topic() {
710 let cmd = parse("help fetch").expect("parsing should succeed");
711 if let Command::Help { topic } = cmd {
712 assert_eq!(topic, Some("fetch".to_string()));
713 } else {
714 panic!("Expected Help command");
715 }
716 }
717
718 #[test]
719 fn test_parse_distill_dry_run() {
720 let cmd = parse("distill --dry-run").expect("parsing should succeed");
721 if let Command::Distill { dry_run } = cmd {
722 assert!(dry_run);
723 } else {
724 panic!("Expected Distill command");
725 }
726 }
727
728 #[test]
729 fn test_parse_distill_short_flag() {
730 let cmd = parse("distill -n").expect("parsing should succeed");
731 if let Command::Distill { dry_run } = cmd {
732 assert!(dry_run);
733 } else {
734 panic!("Expected Distill command");
735 }
736 }
737
738 #[test]
739 fn test_parse_inspect_model() {
740 let cmd = parse("inspect teacher").expect("parsing should succeed");
741 if let Command::Inspect { target } = cmd {
742 assert_eq!(target, InspectTarget::Model("teacher".to_string()));
743 } else {
744 panic!("Expected Inspect command");
745 }
746 }
747
748 #[test]
749 fn test_parse_inspect_memory() {
750 let cmd = parse("inspect memory").expect("parsing should succeed");
751 assert!(matches!(
752 cmd,
753 Command::Inspect {
754 target: InspectTarget::Memory
755 }
756 ));
757 }
758
759 #[test]
760 fn test_parse_command_aliases() {
761 assert!(matches!(
763 parse("download model").expect("parsing should succeed"),
764 Command::Fetch { .. }
765 ));
766 assert!(matches!(
768 parse("show layers").expect("parsing should succeed"),
769 Command::Inspect { .. }
770 ));
771 assert!(matches!(
773 parse("mem").expect("parsing should succeed"),
774 Command::Memory { .. }
775 ));
776 assert!(matches!(
778 parse("train").expect("parsing should succeed"),
779 Command::Distill { .. }
780 ));
781 assert!(matches!(
783 parse("save gguf /tmp/out").expect("parsing should succeed"),
784 Command::Export { .. }
785 ));
786 assert!(matches!(
788 parse("cls").expect("parsing should succeed"),
789 Command::Clear
790 ));
791 assert!(matches!(
793 parse("?").expect("parsing should succeed"),
794 Command::Help { .. }
795 ));
796 assert!(matches!(
798 parse("hist").expect("parsing should succeed"),
799 Command::History
800 ));
801 }
802
803 #[test]
804 fn test_execute_inspect_no_models() {
805 let state = SessionState::new();
806 let result = execute_inspect(&InspectTarget::All, &state);
807 assert!(result
808 .expect("load should succeed")
809 .contains("No models loaded"));
810 }
811
812 #[test]
813 fn test_execute_inspect_layers() {
814 let mut state = SessionState::new();
815 let model = LoadedModel {
816 id: "test".to_string(),
817 path: std::path::PathBuf::from("/tmp"),
818 architecture: "llama".to_string(),
819 parameters: 7_000_000_000,
820 layers: 32,
821 hidden_dim: 4096,
822 role: ModelRole::None,
823 };
824 state.add_model("test".to_string(), model);
825
826 let result =
827 execute_inspect(&InspectTarget::Layers, &state).expect("operation should succeed");
828 assert!(result.contains("32 layers"));
829 }
830
831 #[test]
832 fn test_execute_inspect_model_not_found() {
833 let state = SessionState::new();
834 let result = execute_inspect(&InspectTarget::Model("unknown".to_string()), &state);
835 assert!(result.is_err());
836 }
837
838 #[test]
839 fn test_execute_history_empty() {
840 let state = SessionState::new();
841 let result = execute_history(&state).expect("operation should succeed");
842 assert!(result.contains("No command history"));
843 }
844
845 #[test]
846 fn test_execute_help_topics() {
847 let fetch_help = execute_help(Some("fetch")).expect("operation should succeed");
848 assert!(fetch_help.contains("Download"));
849
850 let inspect_help = execute_help(Some("inspect")).expect("operation should succeed");
851 assert!(inspect_help.contains("Inspect"));
852
853 let memory_help = execute_help(Some("memory")).expect("operation should succeed");
854 assert!(memory_help.contains("memory"));
855
856 let distill_help = execute_help(Some("distill")).expect("operation should succeed");
857 assert!(distill_help.contains("distill"));
858
859 let general_help = execute_help(None).expect("operation should succeed");
860 assert!(general_help.contains("Available commands"));
861 }
862
863 #[test]
866 fn test_general_help_lists_every_parsed_command() {
867 let general_help = execute_help(None).expect("operation should succeed");
868 let listed: Vec<&str> = general_help
869 .lines()
870 .skip(1)
871 .filter_map(|l| l.split_whitespace().next())
872 .collect();
873 for name in [
874 "fetch", "inspect", "memory", "set", "distill", "export", "history", "help", "clear",
875 "quit",
876 ] {
877 assert!(
878 !matches!(parse(name), Ok(Command::Unknown { .. })),
879 "{name} must be a parsed command"
880 );
881 assert!(listed.contains(&name), "general help omits `{name}`");
882 }
883 assert_eq!(listed.len(), 10, "help lists {listed:?}");
884 }
885
886 #[test]
887 fn test_detect_architecture_variants() {
888 assert_eq!(detect_architecture("meta-llama/Llama-2-7b"), "llama");
889 assert_eq!(detect_architecture("bert-base-uncased"), "bert");
890 assert_eq!(detect_architecture("openai-gpt"), "gpt");
891 assert_eq!(detect_architecture("mistralai/Mistral-7B"), "mistral");
892 assert_eq!(detect_architecture("Qwen/Qwen2.5-Coder-0.5B"), "qwen");
893 assert_eq!(detect_architecture("microsoft/phi-2"), "phi");
894 assert_eq!(detect_architecture("tiiuae/falcon-7b"), "falcon");
895 assert_eq!(detect_architecture("mistralai/Mixtral-8x7B"), "mistral");
896 assert_eq!(detect_architecture("custom-model"), "unknown");
897 }
898
899 #[test]
904 fn test_falsify_n04_mistral_before_llama_in_ambiguous_id() {
905 assert_eq!(
909 detect_architecture("mistral-llama-variant"),
910 "mistral",
911 "Mistral should take priority over LLaMA in ambiguous IDs"
912 );
913 }
914
915 #[test]
916 fn test_falsify_n04_unknown_is_explicit() {
917 let result = detect_architecture("totally-custom-model-v3");
920 assert_eq!(result, "unknown");
921 assert!(!result.is_empty());
922 }
923
924 #[test]
925 fn test_estimate_params_variants() {
926 assert_eq!(estimate_params("model-70b"), 70_000_000_000);
927 assert_eq!(estimate_params("model-13b"), 13_000_000_000);
928 assert_eq!(estimate_params("model-7b"), 7_000_000_000);
929 assert_eq!(estimate_params("model-1.1b"), 1_100_000_000);
930 assert_eq!(estimate_params("bert-base"), 350_000_000);
931 }
932
933 #[test]
934 fn test_estimate_layers_variants() {
935 assert_eq!(estimate_layers("model-70b"), 80);
936 assert_eq!(estimate_layers("model-13b"), 40);
937 assert_eq!(estimate_layers("model-7b"), 32);
938 assert_eq!(estimate_layers("bert-base"), 12);
939 }
940
941 #[test]
942 fn test_execute_set_invalid_number() {
943 let mut state = SessionState::new();
944 let result = execute_set("batch_size", "not_a_number", &mut state);
945 assert!(result.is_err());
946 }
947
948 #[test]
949 fn test_execute_set_unknown_key() {
950 let mut state = SessionState::new();
951 let result = execute_set("unknown_setting", "value", &mut state);
952 assert!(result.is_err());
953 }
954
955 #[test]
956 fn test_execute_distill_no_teacher() {
957 let state = SessionState::new();
958 let result = execute_distill(true, &state);
959 assert!(result.is_err());
960 }
961
962 #[test]
963 fn test_execute_distill_no_student() {
964 let mut state = SessionState::new();
965 let model = LoadedModel {
966 id: "teacher".to_string(),
967 path: std::path::PathBuf::from("/tmp"),
968 architecture: "llama".to_string(),
969 parameters: 7_000_000_000,
970 layers: 32,
971 hidden_dim: 4096,
972 role: ModelRole::Teacher,
973 };
974 state.add_model("teacher".to_string(), model);
975
976 let result = execute_distill(true, &state);
977 assert!(result.is_err());
978 }
979
980 #[test]
981 fn test_execute_distill_success() {
982 let mut state = SessionState::new();
983
984 let teacher = LoadedModel {
985 id: "teacher/model".to_string(),
986 path: std::path::PathBuf::from("/tmp/t"),
987 architecture: "llama".to_string(),
988 parameters: 7_000_000_000,
989 layers: 32,
990 hidden_dim: 4096,
991 role: ModelRole::Teacher,
992 };
993 state.add_model("teacher".to_string(), teacher);
994
995 let student = LoadedModel {
996 id: "student/model".to_string(),
997 path: std::path::PathBuf::from("/tmp/s"),
998 architecture: "llama".to_string(),
999 parameters: 1_000_000_000,
1000 layers: 12,
1001 hidden_dim: 2048,
1002 role: ModelRole::Student,
1003 };
1004 state.add_model("student".to_string(), student);
1005
1006 let result = execute_distill(true, &state).expect("operation should succeed");
1007 assert!(result.contains("Dry run"));
1008 }
1009
1010 #[test]
1011 fn test_execute_export_refuses_rather_than_fabricating() {
1012 let state = SessionState::new();
1013 let result = execute_export("safetensors", "/tmp/model.st", &state);
1014 assert!(
1019 result.is_err(),
1020 "export must not report success without exporting"
1021 );
1022 let msg = format!("{}", result.unwrap_err());
1023 assert!(
1024 !msg.contains("Exported to"),
1025 "the error must not read like a success: {msg}"
1026 );
1027 assert!(msg.contains("not implemented"), "must say why: {msg}");
1028 }
1029
1030 #[test]
1031 fn test_execute_export_leaves_no_file_behind() {
1032 let state = SessionState::new();
1035 let path = std::env::temp_dir().join("apr-2519-export-probe.st");
1036 let _ = std::fs::remove_file(&path);
1037 let _ = execute_export(
1038 "safetensors",
1039 path.to_str().expect("utf-8 temp path"),
1040 &state,
1041 );
1042 assert!(!path.exists(), "export refused but still created {path:?}");
1043 }
1044
1045 #[test]
1046 fn test_parse_empty_input() {
1047 let cmd = parse("").expect("parsing should succeed");
1048 assert!(matches!(cmd, Command::Unknown { .. }));
1049 }
1050
1051 #[test]
1052 fn test_execute_memory_with_args() {
1053 let state = SessionState::new();
1054 let result =
1055 execute_memory(Some(64), Some(1024), &state).expect("operation should succeed");
1056 assert!(result.contains("batch=64"));
1057 assert!(result.contains("seq=1024"));
1058 }
1059
1060 #[test]
1061 fn test_command_enum_equality() {
1062 assert_eq!(Command::Quit, Command::Quit);
1063 assert_eq!(Command::Clear, Command::Clear);
1064 assert_eq!(Command::History, Command::History);
1065 }
1066}