Skip to main content

pray_core/
cli_suggest.rs

1pub const TOP_LEVEL_COMMANDS: &[&str] = &[
2    "add",
3    "apply",
4    "clean",
5    "completion",
6    "confess",
7    "drift",
8    "explain",
9    "fmt",
10    "format",
11    "help",
12    "init",
13    "install",
14    "list",
15    "login",
16    "manifest",
17    "outdated",
18    "package",
19    "plan",
20    "prayer",
21    "publish",
22    "remove",
23    "render",
24    "repo",
25    "search",
26    "serve",
27    "sync",
28    "token",
29    "tree",
30    "trust",
31    "unlock",
32    "update",
33    "upgrade",
34    "vendor",
35    "verify",
36    "version",
37    "yank",
38];
39
40pub fn unknown_command_message(command: &str) -> String {
41    let mut message = format!("unknown command: {command}");
42    if let Some(suggestion) = suggest_command(command, TOP_LEVEL_COMMANDS) {
43        message.push_str(&format!("\nDid you mean `{suggestion}`?"));
44    }
45    message.push_str("\nSee 'pray --help'.");
46    message
47}
48
49pub fn suggest_command<'a>(input: &str, candidates: &'a [&str]) -> Option<&'a str> {
50    let maximum_distance = if input.chars().count() <= 3 { 1 } else { 2 };
51    candidates
52        .iter()
53        .copied()
54        .filter(|candidate| levenshtein_distance(input, candidate) <= maximum_distance)
55        .min_by_key(|candidate| levenshtein_distance(input, candidate))
56}
57
58fn levenshtein_distance(left: &str, right: &str) -> usize {
59    let left_chars: Vec<char> = left.chars().collect();
60    let right_chars: Vec<char> = right.chars().collect();
61    let left_length = left_chars.len();
62    let right_length = right_chars.len();
63
64    if left_length == 0 {
65        return right_length;
66    }
67    if right_length == 0 {
68        return left_length;
69    }
70
71    let mut previous_row: Vec<usize> = (0..=right_length).collect();
72    let mut current_row = vec![0; right_length + 1];
73
74    for (left_index, left_character) in left_chars.iter().enumerate() {
75        current_row[0] = left_index + 1;
76        for (right_index, right_character) in right_chars.iter().enumerate() {
77            let substitution_cost = if left_character == right_character {
78                0
79            } else {
80                1
81            };
82            current_row[right_index + 1] = (previous_row[right_index + 1] + 1)
83                .min(current_row[right_index] + 1)
84                .min(previous_row[right_index] + substitution_cost);
85        }
86        std::mem::swap(&mut previous_row, &mut current_row);
87    }
88
89    previous_row[right_length]
90}
91
92#[cfg(test)]
93mod tests {
94    use super::{suggest_command, unknown_command_message, TOP_LEVEL_COMMANDS};
95
96    #[test]
97    fn suggests_install_for_common_typo() {
98        assert_eq!(
99            suggest_command("instal", TOP_LEVEL_COMMANDS),
100            Some("install")
101        );
102    }
103
104    #[test]
105    fn unknown_command_includes_suggestion() {
106        let message = unknown_command_message("instal");
107        assert!(message.contains("unknown command: instal"));
108        assert!(message.contains("Did you mean `install`?"));
109        assert!(message.contains("See 'pray --help'."));
110    }
111}