Skip to main content

pray_core/
cli_suggest.rs

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