Skip to main content

cli_rs/
parser.rs

1use std::{env, fmt::Write};
2
3use crate::{
4    cli_error::{CliError, CliResult},
5    command::{CompletionMode, ParserInfo},
6    flag::Flag,
7    input::{Input, InputType},
8};
9
10use colored::*;
11
12impl<C> Cmd for C where C: ParserInfo {}
13
14pub struct CompOut {
15    pub name: String,
16    pub desc: Option<String>,
17}
18
19fn version_flag() -> Flag<'static, bool> {
20    Flag::bool("version").description("display CLI version")
21}
22
23fn help_flag() -> Flag<'static, bool> {
24    Flag::bool("help").description("view help")
25}
26
27pub trait Cmd: ParserInfo {
28    fn gen_help(&mut self) -> CliError {
29        let cmd_path = self.docs().cmd_path();
30        let mut help_message = String::new();
31
32        write!(help_message, "{}", cmd_path.bold().green()).unwrap();
33
34        if let Some(description) = &self.docs().description {
35            write!(help_message, " - {description}").unwrap();
36        }
37
38        writeln!(help_message, "\n").unwrap();
39
40        let mut version = version_flag();
41        let mut help = help_flag();
42        let mut built_in: Vec<&mut dyn Input> = vec![&mut help];
43        if self.docs().version.is_some() {
44            built_in.push(&mut version);
45        }
46        let subcommands = self.subcommand_docs();
47        writeln!(help_message, "{}", "USAGE:".bold().yellow()).unwrap();
48        if subcommands.is_empty() {
49            let usage = format!("{cmd_path} [options], <args>").bold();
50            writeln!(help_message, "\t{usage}").unwrap();
51            writeln!(help_message, "\n{}", "FLAGS:".yellow().bold()).unwrap();
52
53            let width = self
54                .symbols()
55                .into_iter()
56                .map(|s| s.display_name().len() + 3)
57                .chain([10]) // --version
58                .max()
59                .unwrap();
60
61            for symbol in self.symbols().iter().chain(built_in.iter()) {
62                if symbol.type_name() == InputType::Flag {
63                    write!(help_message, "\t--{:width$}", symbol.display_name().bold()).unwrap();
64                    if let Some(desc) = symbol.description() {
65                        write!(help_message, " {desc}").unwrap();
66                    }
67                    writeln!(help_message).unwrap();
68                }
69            }
70            writeln!(help_message, "\n{}", "ARGS:".yellow().bold()).unwrap();
71            for symbol in self.symbols() {
72                if symbol.type_name() == InputType::Arg {
73                    write!(help_message, "\t{:width$}", symbol.display_name().bold()).unwrap();
74                    if let Some(desc) = symbol.description() {
75                        write!(help_message, " {desc}").unwrap();
76                    }
77                    writeln!(help_message).unwrap();
78                }
79            }
80        } else {
81            let usage = format! {"{cmd_path} <subcommand>"}.bold();
82            writeln!(help_message, "\t{usage}").unwrap();
83            writeln!(help_message, "\n{}", "SUBCOMMANDS:".yellow().bold()).unwrap();
84            let sub_width = subcommands.iter().map(|s| s.name.len()).max().unwrap();
85            for subcommand in subcommands {
86                write!(help_message, "\t{:sub_width$}", subcommand.name.bold()).unwrap();
87                if let Some(description) = subcommand.description {
88                    write!(help_message, " {description}").unwrap();
89                }
90
91                writeln!(help_message).unwrap();
92            }
93
94            writeln!(help_message, "\n{}", "FLAGS:".yellow().bold()).unwrap();
95
96            let flag_width = built_in
97                .iter()
98                .map(|s| s.display_name().len() + 3)
99                .max()
100                .unwrap();
101
102            for symbol in built_in {
103                if symbol.type_name() == InputType::Flag {
104                    write!(
105                        help_message,
106                        "\t--{:flag_width$}",
107                        symbol.display_name().bold()
108                    )
109                    .unwrap();
110                    if let Some(desc) = symbol.description() {
111                        write!(help_message, " {desc}").unwrap();
112                    }
113                    writeln!(help_message).unwrap();
114                }
115            }
116        }
117
118        CliError::from(help_message)
119    }
120
121    // split this out into a trait that is pub, make the rest not pub
122    fn parse(&mut self) -> CliResult<()> {
123        #[cfg(windows)]
124        let _ = colored::control::set_virtual_terminal(true);
125
126        let args: Vec<String> = env::args().collect();
127        // cmd complete shell word_idx [input]
128        if args.len() >= 5 && args[1] == "complete" {
129            let name = self.docs().name.to_string();
130            // going to need some serious tests here
131            let shell = args[2].parse::<CompletionMode>().unwrap();
132            let prompt: Vec<String> = if shell == CompletionMode::Fish {
133                let prompt = &args[4];
134                prompt.split(' ').map(|s| s.to_string()).collect()
135            } else {
136                let idx: usize = args[3].parse().unwrap();
137                let prompt = &args[4];
138                prompt
139                    .split(' ')
140                    .map(|s| s.to_string())
141                    .take(idx + 1)
142                    .collect()
143            };
144
145            let mut last_command_location = 0;
146
147            for (i, token) in prompt.iter().enumerate().rev() {
148                if token == &name {
149                    last_command_location = i;
150                    break;
151                }
152            }
153
154            let prompt = &prompt[last_command_location..];
155
156            match self.complete_args(&prompt[1..]) {
157                Ok(outputs) => {
158                    match shell {
159                        CompletionMode::Bash => {
160                            for out in outputs {
161                                println!("{}", out.name);
162                            }
163                        }
164                        CompletionMode::Fish => {
165                            for out in outputs {
166                                if let Some(desc) = out.desc {
167                                    println!("{}\t{}", out.name, desc);
168                                } else {
169                                    println!("{}", out.name);
170                                }
171                            }
172                        }
173                        CompletionMode::Zsh => {
174                            let comps = outputs
175                                .into_iter()
176                                .map(|out| {
177                                    if let Some(desc) = out.desc {
178                                        let desc = desc.replace('\'', "");
179                                        let desc = desc.replace('"', "");
180                                        format!("'{}:{}'", out.name, desc)
181                                    } else {
182                                        format!("'{}'", out.name)
183                                    }
184                                })
185                                .collect::<Vec<String>>()
186                                .join(" ");
187
188                            println!("_describe '{name}' \"({comps})\"");
189                        }
190                    };
191
192                    return Ok(());
193                }
194                Err(error) => {
195                    return Err(error);
196                }
197            }
198        }
199
200        self.parse_args(&args[1..])
201    }
202
203    fn complete_args(&mut self, tokens: &[String]) -> CliResult<Vec<CompOut>> {
204        let mut completions = vec![];
205        if tokens.is_empty() {
206            return Ok(completions);
207        }
208
209        let subcommands = self.subcommand_docs();
210
211        // recurse into subcommand?
212        if !subcommands.is_empty() && !tokens.is_empty() {
213            let token = &tokens[0];
214            let mut subcommand_index = None;
215            for (idx, subcommand) in subcommands.iter().enumerate() {
216                if &subcommand.name == token {
217                    subcommand_index = Some(idx);
218                }
219            }
220
221            // todo check this
222            if let Some(index) = subcommand_index {
223                return self.complete_subcommand(index, &tokens[1..]);
224            }
225
226            // print subcommands that begin with the token
227            if tokens.len() == 1 && !tokens[0].starts_with('-') {
228                for sub in subcommands {
229                    if sub.name.starts_with(token) {
230                        let name = &sub.name;
231                        let desc = &sub.description;
232                        completions.push(CompOut {
233                            name: name.to_string(),
234                            desc: desc.to_owned(),
235                        })
236                    }
237                }
238
239                return Ok(completions);
240            }
241        }
242
243        let has_version = self.docs().version.is_some();
244        let mut symbols = self.symbols();
245
246        let mut positional_args_so_far = 0;
247        if tokens.len() > 1 {
248            for token in &tokens[0..tokens.len() - 1] {
249                if !token.starts_with('-') {
250                    // in a future where we manage errors more properly, this section could be
251                    // closer to how the parser works, eliminating consumed symbols and helping
252                    // the end user not see completions for flags they've already typed. Presently
253                    // that code would start outputting errors.
254                    positional_args_so_far += 1;
255                }
256            }
257        }
258
259        let token = &tokens[tokens.len() - 1];
260        if let Some(mut completion_token) = token.strip_prefix('-') {
261            if let Some(second_dash_removed) = completion_token.strip_prefix('-') {
262                completion_token = second_dash_removed;
263            }
264            let value_completion = completion_token.split('=').collect::<Vec<&str>>();
265            if value_completion.len() > 1 {
266                for symbol in &mut symbols {
267                    if symbol.display_name() == value_completion[0] {
268                        for completion in symbol.complete(value_completion[1])? {
269                            completions.push(CompOut {
270                                name: format!("--{}={completion}", symbol.display_name()),
271                                desc: None,
272                            });
273                        }
274                        return Ok(completions);
275                    }
276                }
277            }
278
279            let mut version = version_flag();
280            let mut help = help_flag();
281            let mut built_in: Vec<&mut dyn Input> = vec![&mut help];
282            if has_version {
283                built_in.push(&mut version);
284            }
285
286            symbols
287                .iter()
288                .chain(built_in.iter())
289                .filter(|sym| sym.type_name() == InputType::Flag)
290                .filter(|sym| sym.display_name().starts_with(completion_token))
291                .for_each(|flag| {
292                    if flag.is_bool_flag() {
293                        completions.push(CompOut {
294                            name: format!("--{}", flag.display_name()),
295                            desc: flag.description(),
296                        });
297                    } else {
298                        completions.push(CompOut {
299                            name: format!("--{}=", flag.display_name()),
300                            desc: flag.description(),
301                        });
302                    }
303                });
304        } else {
305            let arg = symbols
306                .iter_mut()
307                .filter(|sym| sym.type_name() == InputType::Arg)
308                .nth(positional_args_so_far);
309
310            if let Some(arg) = arg {
311                for option in arg.complete(token)? {
312                    completions.push(CompOut {
313                        name: option.to_string(),
314                        desc: None,
315                    });
316                }
317            }
318        }
319
320        Ok(completions)
321    }
322
323    fn parse_args(&mut self, tokens: &[String]) -> CliResult<()> {
324        let subcommands = self.subcommand_docs();
325        let symbols = self.symbols();
326        let required_args = symbols.iter().filter(|f| !f.has_default()).count();
327
328        if tokens.is_empty() && (required_args > 0 || !subcommands.is_empty()) {
329            return Err(self.gen_help());
330        }
331
332        // try to match subcommands
333        if !tokens.is_empty() {
334            let token = &tokens[0];
335            if token == "--help" {
336                println!("{}", self.gen_help().msg);
337                return Ok(());
338            }
339
340            if token == "--version" {
341                let docs = &self.docs();
342                if let Some(version) = &docs.version {
343                    println!("{} -- {}", docs.cmd_path(), version);
344                    return Ok(());
345                }
346            }
347            if !subcommands.is_empty() {
348                for (idx, subcommand) in subcommands.iter().enumerate() {
349                    if &subcommand.name == token {
350                        return self.parse_subcommand(idx, &tokens[1..]);
351                    }
352                }
353
354                return Err(CliError::from(format!("{token} is not a valid subcommand")));
355            }
356        }
357
358        let mut symbols = self.symbols();
359
360        for token in tokens {
361            if token.starts_with('-') {
362                let mut token_matched = false;
363                for symbol in &mut symbols {
364                    if !symbol.parsed() && symbol.type_name() == InputType::Flag {
365                        let consumed = symbol.parse(token)?;
366                        if consumed {
367                            token_matched = true;
368                        }
369                    }
370                }
371                if !token_matched {
372                    return Err(CliError::from(format!(
373                        "Unexpected flag-like token found {token}"
374                    )));
375                }
376            } else {
377                'args: for symbol in &mut symbols {
378                    if !symbol.parsed() && symbol.type_name() == InputType::Arg {
379                        symbol.parse(token)?;
380                        break 'args;
381                    }
382                }
383            }
384        }
385
386        for symbol in symbols {
387            if symbol.type_name() == InputType::Arg && !symbol.has_default() && !symbol.parsed() {
388                return Err(CliError::from(format!(
389                    "Missing required argument: {}",
390                    symbol.display_name()
391                )));
392            }
393        }
394
395        self.call_handler()
396    }
397}