Skip to main content

entrenar_shell/
commands.rs

1//! Command parsing and execution for the REPL.
2
3use crate::state::{HistoryEntry, LoadedModel, ModelRole, SessionState};
4use entrenar_common::{EntrenarError, Result};
5
6/// A parsed command.
7#[derive(Debug, Clone, PartialEq)]
8pub enum Command {
9    /// Fetch a model from HuggingFace
10    Fetch { model_id: String, role: ModelRole },
11    /// Inspect a loaded model
12    Inspect { target: InspectTarget },
13    /// Estimate memory requirements
14    Memory {
15        batch_size: Option<u32>,
16        seq_len: Option<usize>,
17    },
18    /// Set configuration values
19    Set { key: String, value: String },
20    /// Run distillation
21    Distill { dry_run: bool },
22    /// Export model to file
23    Export { format: String, path: String },
24    /// Show command history
25    History,
26    /// Show help
27    Help { topic: Option<String> },
28    /// Clear screen
29    Clear,
30    /// Quit the shell
31    Quit,
32    /// Unknown command
33    Unknown { input: String },
34}
35
36/// Target for inspect command.
37#[derive(Debug, Clone, PartialEq)]
38pub enum InspectTarget {
39    /// Inspect layer structure
40    Layers,
41    /// Inspect memory usage
42    Memory,
43    /// Inspect all info
44    All,
45    /// Inspect specific model by name
46    Model(String),
47}
48
49/// Parse a command string into a Command.
50pub 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
180/// Execute a command and update state.
181pub 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    // Record in history (except for help/history/clear/quit)
215    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    // Simulate model fetching
233    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; // FP16
316    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
445/// N-04 (Meyer DbC): Best-effort architecture guess from model ID string.
446/// This is informational only (shell display) — actual architecture detection
447/// for inference uses tensor-name-based `ArchitectureDetector::detect()`.
448/// Order matters: more specific patterns must come before generic ones
449/// (e.g., "mistral" before "llama" since Mistral inherits LLaMA naming).
450const 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        // download = fetch
688        assert!(matches!(
689            parse("download model").expect("parsing should succeed"),
690            Command::Fetch { .. }
691        ));
692        // show = inspect
693        assert!(matches!(
694            parse("show layers").expect("parsing should succeed"),
695            Command::Inspect { .. }
696        ));
697        // mem = memory
698        assert!(matches!(
699            parse("mem").expect("parsing should succeed"),
700            Command::Memory { .. }
701        ));
702        // train = distill
703        assert!(matches!(
704            parse("train").expect("parsing should succeed"),
705            Command::Distill { .. }
706        ));
707        // save = export (needs args)
708        assert!(matches!(
709            parse("save gguf /tmp/out").expect("parsing should succeed"),
710            Command::Export { .. }
711        ));
712        // cls = clear
713        assert!(matches!(
714            parse("cls").expect("parsing should succeed"),
715            Command::Clear
716        ));
717        // ? = help
718        assert!(matches!(
719            parse("?").expect("parsing should succeed"),
720            Command::Help { .. }
721        ));
722        // hist = history
723        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    // =========================================================================
803    // FALSIFY tests — contract violation sweep (N-04)
804    // =========================================================================
805
806    #[test]
807    fn test_falsify_n04_mistral_before_llama_in_ambiguous_id() {
808        // N-04: "mistral" must be checked before "llama" since some Mistral
809        // model IDs contain both substrings. Mistral-specific detection
810        // should take priority.
811        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        // N-04: Unknown architecture must be an explicit "unknown" string,
821        // never an empty string or a silent default to a known architecture.
822        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}