1use super::error::{CliError, CliErrorKind};
26use super::help::HelpView;
27use super::spec::{ArgSpec, Choice, CommandSpec, ValueHint};
28use clap::error::{ContextKind, ContextValue, ErrorKind};
29use clap::{ArgAction, ArgMatches, Command};
30use rich::{Console, Text};
31use std::ffi::OsString;
32use std::io::{IsTerminal, Write};
33
34impl CommandSpec {
35 pub fn from_clap(cmd: &clap::Command) -> CommandSpec {
46 let mut cmd = cmd.clone();
47 cmd.build();
48 let mut spec = convert(&cmd);
49 spec.bin_name = cmd
50 .get_bin_name()
51 .filter(|bin| *bin != cmd.get_name())
52 .map(str::to_string);
53 spec
54 }
55}
56
57fn convert(cmd: &Command) -> CommandSpec {
58 let text = |s: Option<&clap::builder::StyledStr>| s.map(|s| s.to_string());
59 let mut spec = CommandSpec::new(cmd.get_name());
60 spec.version = cmd.get_version().map(str::to_string);
61 spec.about = text(cmd.get_about()).unwrap_or_default();
62 spec.long_about = text(cmd.get_long_about());
63 spec.aliases = cmd.get_visible_aliases().map(str::to_string).collect();
64 if let Some(usage) = cmd.get_overridden_usage() {
65 spec.usage = usage
66 .to_string()
67 .lines()
68 .map(|l| l.trim().to_string())
69 .filter(|l| !l.is_empty())
70 .collect();
71 }
72 spec.args = cmd.get_arguments().map(convert_arg).collect();
73 spec.subcommands = cmd.get_subcommands().map(convert).collect();
74 spec.subcommand_heading = cmd.get_subcommand_help_heading().map(str::to_string);
75 spec.hidden = cmd.is_hide_set();
76 spec.subcommand_required = cmd.is_subcommand_required_set();
77 if let Some(after) = text(cmd.get_after_help()).or_else(|| text(cmd.get_after_long_help())) {
78 spec = spec.section("", after);
79 }
80 spec
81}
82
83fn convert_arg(arg: &clap::Arg) -> ArgSpec {
84 let action = arg.get_action();
85 let takes = action.takes_values();
86 let mut spec = ArgSpec::new(arg.get_id().as_str());
87 spec.long = arg.get_long().map(str::to_string);
88 spec.short = arg.get_short();
89 spec.short_aliases = arg.get_visible_short_aliases().unwrap_or_default();
90 spec.aliases = arg
91 .get_visible_aliases()
92 .unwrap_or_default()
93 .into_iter()
94 .map(str::to_string)
95 .collect();
96 spec.positional = arg.is_positional();
97 spec.global = arg.is_global_set();
98 spec.help = arg.get_help().map(|h| h.to_string()).unwrap_or_default();
99 spec.long_help = arg.get_long_help().map(|h| h.to_string());
100 spec.required = arg.is_required_set();
101 spec.hidden = arg.is_hide_set();
102 spec.heading = arg.get_help_heading().map(str::to_string);
103 spec.multiple = matches!(action, ArgAction::Count | ArgAction::Append)
104 || arg.get_num_args().is_some_and(|n| n.max_values() > 1);
105 if !arg.is_hide_env_set() {
106 spec.env = arg.get_env().map(|e| e.to_string_lossy().into_owned());
107 }
108 if takes {
109 let names: Vec<String> = arg
110 .get_value_names()
111 .map(|names| names.iter().map(|n| n.to_string()).collect())
112 .unwrap_or_default();
113 spec.value_name = Some(if names.is_empty() {
115 arg.get_id().as_str().to_string()
116 } else {
117 names.join(" ")
118 });
119 let choices: Vec<Choice> = arg
120 .get_possible_values()
121 .into_iter()
122 .filter(|v| !v.is_hide_set())
123 .map(|v| {
124 Choice::new(v.get_name())
125 .help(v.get_help().map(|h| h.to_string()).unwrap_or_default())
126 })
127 .collect();
128 spec.value = if !choices.is_empty() && !arg.is_hide_possible_values_set() {
129 ValueHint::Choices(choices)
130 } else {
131 match arg.get_value_hint() {
132 clap::ValueHint::FilePath => ValueHint::File,
133 clap::ValueHint::DirPath => ValueHint::Dir,
134 clap::ValueHint::AnyPath => ValueHint::Path,
135 clap::ValueHint::ExecutablePath
136 | clap::ValueHint::CommandName
137 | clap::ValueHint::CommandString
138 | clap::ValueHint::CommandWithArguments => ValueHint::Command,
139 clap::ValueHint::Url => ValueHint::Url,
140 _ => ValueHint::Any,
141 }
142 };
143 let defaults: Vec<String> = arg
144 .get_default_values()
145 .iter()
146 .map(|d| d.to_string_lossy().into_owned())
147 .collect();
148 if !defaults.is_empty() && !arg.is_hide_default_value_set() {
149 spec.default = Some(defaults.join(", "));
150 }
151 }
152 spec
153}
154
155impl CliError {
156 pub fn from_clap(err: &clap::Error, cmd: Option<&Command>) -> CliError {
161 let get = |kind: ContextKind| -> Option<String> {
162 match err.get(kind)? {
163 ContextValue::String(s) => Some(s.clone()),
164 ContextValue::Strings(v) => Some(v.join(", ")),
165 ContextValue::StyledStr(s) => Some(s.to_string()),
166 ContextValue::StyledStrs(v) => Some(
167 v.iter()
168 .map(|s| s.to_string())
169 .collect::<Vec<_>>()
170 .join(", "),
171 ),
172 ContextValue::Number(n) => Some(n.to_string()),
173 ContextValue::Bool(b) => Some(b.to_string()),
174 _ => None,
175 }
176 };
177 let list = |kind: ContextKind| -> Vec<String> {
178 match err.get(kind) {
179 Some(ContextValue::String(s)) => vec![s.clone()],
180 Some(ContextValue::Strings(v)) => v.clone(),
181 _ => Vec::new(),
182 }
183 };
184 let argument = get(ContextKind::InvalidArg).or_else(|| get(ContextKind::InvalidSubcommand));
185 let value = get(ContextKind::InvalidValue);
186 let (kind, message) = match err.kind() {
187 ErrorKind::UnknownArgument => (CliErrorKind::UnknownArgument, String::new()),
188 ErrorKind::InvalidSubcommand => (CliErrorKind::UnknownSubcommand, String::new()),
189 ErrorKind::InvalidValue if value.as_deref() == Some("") => {
190 (CliErrorKind::MissingValue, String::new())
191 }
192 ErrorKind::InvalidValue => (CliErrorKind::InvalidValue, String::new()),
193 ErrorKind::ValueValidation => {
194 let reason = std::error::Error::source(err).map(|e| e.to_string());
195 let message = format!(
196 "invalid value '{}' for '{}'{}",
197 value.as_deref().unwrap_or_default(),
198 argument.as_deref().unwrap_or_default(),
199 reason.map(|r| format!(": {r}")).unwrap_or_default()
200 );
201 (CliErrorKind::InvalidValue, message)
202 }
203 ErrorKind::NoEquals => (
204 CliErrorKind::MissingValue,
205 format!(
206 "equal sign is needed when assigning values to '{}'",
207 argument.as_deref().unwrap_or_default()
208 ),
209 ),
210 ErrorKind::TooManyValues => (CliErrorKind::UnexpectedValue, String::new()),
211 ErrorKind::TooFewValues | ErrorKind::WrongNumberOfValues => {
212 let expected = get(ContextKind::ExpectedNumValues)
213 .or_else(|| get(ContextKind::MinValues))
214 .unwrap_or_default();
215 let actual = get(ContextKind::ActualNumValues).unwrap_or_default();
216 let message = format!(
217 "{expected} values required by '{}'; {actual} provided",
218 argument.as_deref().unwrap_or_default()
219 );
220 (CliErrorKind::MissingValue, message)
221 }
222 ErrorKind::ArgumentConflict => (CliErrorKind::Conflict, String::new()),
223 ErrorKind::MissingRequiredArgument => (CliErrorKind::MissingRequired, String::new()),
224 ErrorKind::MissingSubcommand => {
225 let name = cmd.map(|c| c.get_name()).unwrap_or("command");
226 let message = format!("'{name}' requires a subcommand but one was not provided");
227 (CliErrorKind::MissingRequired, message)
228 }
229 _ => {
230 let rendered = err.to_string();
231 let first = rendered.lines().next().unwrap_or_default();
232 let message = first.strip_prefix("error: ").unwrap_or(first).to_string();
233 (CliErrorKind::Other, message)
234 }
235 };
236 let mut error = CliError::new(kind, message);
237 error.argument = argument;
238 error.value = match kind {
239 CliErrorKind::Conflict => get(ContextKind::PriorArg),
240 _ => value.filter(|v| !v.is_empty()),
241 };
242 for context in [
243 ContextKind::SuggestedArg,
244 ContextKind::SuggestedSubcommand,
245 ContextKind::SuggestedValue,
246 ContextKind::SuggestedCommand,
247 ] {
248 error.suggestions.extend(list(context));
249 }
250 error.possible_values = list(ContextKind::ValidValue);
251 error.usage = get(ContextKind::Usage)
252 .map(|u| u.trim_start_matches("Usage:").trim().to_string())
253 .or_else(|| cmd.map(|c| CommandSpec::from_clap(c).usage_lines().join("\n")));
254 if cmd.is_some_and(|c| !c.is_disable_help_flag_set()) {
255 error.help_flag = Some("--help".into());
256 }
257 error
258 }
259}
260
261pub fn print_help(cmd: &Command) {
263 Console::new().print(&HelpView::new(&CommandSpec::from_clap(cmd)));
264}
265
266#[derive(Debug)]
269pub struct ParseExit {
270 pub error: clap::Error,
272 pub output: String,
274 pub use_stderr: bool,
276 pub code: i32,
278}
279
280impl ParseExit {
281 pub fn print(&self) {
283 let _ = if self.use_stderr {
284 std::io::stderr().lock().write_all(self.output.as_bytes())
285 } else {
286 std::io::stdout().lock().write_all(self.output.as_bytes())
287 };
288 }
289
290 pub fn exit(&self) -> ! {
292 self.print();
293 let _ = std::io::stdout().flush();
294 std::process::exit(self.code)
295 }
296}
297
298impl std::fmt::Display for ParseExit {
299 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
300 self.error.fmt(f)
301 }
302}
303
304impl std::error::Error for ParseExit {}
305
306pub fn parse_or_exit(cmd: Command) -> ArgMatches {
309 let spec = CommandSpec::from_clap(&cmd);
310 parse_or_exit_with_spec(cmd, &spec)
311}
312
313pub fn parse_or_exit_with_spec(cmd: Command, spec: &CommandSpec) -> ArgMatches {
317 let args: Vec<OsString> = std::env::args_os().collect();
318 match cmd.clone().try_get_matches_from(&args) {
319 Ok(matches) => matches,
320 Err(err) => {
321 let console = if err.use_stderr() {
322 Console::builder()
323 .force_terminal(std::io::stderr().is_terminal())
324 .build()
325 } else {
326 Console::new()
327 };
328 render_exit(&console, &cmd, spec, &args, err).exit()
329 }
330 }
331}
332
333pub fn try_parse_from<I, T>(cmd: Command, args: I) -> Result<ArgMatches, ParseExit>
336where
337 I: IntoIterator<Item = T>,
338 T: Into<OsString> + Clone,
339{
340 let args: Vec<OsString> = args.into_iter().map(Into::into).collect();
341 match cmd.clone().try_get_matches_from(&args) {
342 Ok(matches) => Ok(matches),
343 Err(err) => {
344 let console = if err.use_stderr() {
345 Console::builder()
346 .force_terminal(std::io::stderr().is_terminal())
347 .build()
348 } else {
349 Console::new()
350 };
351 let spec = CommandSpec::from_clap(&cmd);
352 Err(render_exit(&console, &cmd, &spec, &args, err))
353 }
354 }
355}
356
357pub fn try_parse_with<I, T>(
360 console: &Console,
361 cmd: Command,
362 args: I,
363) -> Result<ArgMatches, ParseExit>
364where
365 I: IntoIterator<Item = T>,
366 T: Into<OsString> + Clone,
367{
368 let spec = CommandSpec::from_clap(&cmd);
369 try_parse_with_spec(console, cmd, &spec, args)
370}
371
372pub fn try_parse_with_spec<I, T>(
375 console: &Console,
376 cmd: Command,
377 spec: &CommandSpec,
378 args: I,
379) -> Result<ArgMatches, ParseExit>
380where
381 I: IntoIterator<Item = T>,
382 T: Into<OsString> + Clone,
383{
384 let args: Vec<OsString> = args.into_iter().map(Into::into).collect();
385 match cmd.clone().try_get_matches_from(&args) {
386 Ok(matches) => Ok(matches),
387 Err(err) => Err(render_exit(console, &cmd, spec, &args, err)),
388 }
389}
390
391fn subcommand_path(spec: &CommandSpec, args: &[OsString]) -> (Vec<String>, bool) {
393 let mut path: Vec<String> = Vec::new();
394 let mut current = spec;
395 let mut help_command = false;
396 for arg in args.iter().skip(1) {
397 let arg = arg.to_string_lossy();
398 if arg == "--" {
399 break;
400 }
401 if arg.starts_with('-') {
402 continue;
403 }
404 if !help_command && arg == "help" && current.find_subcommand("help").is_some() {
405 help_command = true;
407 continue;
408 }
409 match current.find_subcommand(&arg) {
410 Some(next) => {
411 path.push(next.name.clone());
412 current = next;
413 }
414 None => break,
415 }
416 }
417 (path, help_command)
418}
419
420fn render_exit(
421 console: &Console,
422 cmd: &Command,
423 spec: &CommandSpec,
424 args: &[OsString],
425 err: clap::Error,
426) -> ParseExit {
427 let (path, help_command) = subcommand_path(spec, args);
428 let names: Vec<&str> = path.iter().map(String::as_str).collect();
429 let mut sub = spec;
430 for name in &names {
431 sub = sub.find_subcommand(name).unwrap_or(sub);
432 }
433 let shown = std::iter::once(spec.display_name())
434 .chain(names.iter().copied())
435 .collect::<Vec<_>>()
436 .join(" ");
437 let output = match err.kind() {
438 ErrorKind::DisplayHelp | ErrorKind::DisplayHelpOnMissingArgumentOrSubcommand => {
439 let long = help_command || args.iter().any(|a| a == "--help");
440 let view = HelpView::for_path(spec, &names)
441 .unwrap_or_else(|| HelpView::new(spec))
442 .long(long);
443 console.render_export(&view)
444 }
445 ErrorKind::DisplayVersion => {
446 let version = sub
447 .version
448 .as_deref()
449 .or(spec.version.as_deref())
450 .unwrap_or_default();
451 console.render_export(&Text::new(format!("{shown} {version}")))
452 }
453 _ => {
454 let mut error = CliError::from_clap(&err, Some(cmd));
455 error.usage = Some(sub.usage_lines_as(&shown).join("\n"));
456 console.render_export(&error.to_diagnostic())
457 }
458 };
459 ParseExit {
460 use_stderr: err.use_stderr(),
461 code: err.exit_code(),
462 error: err,
463 output,
464 }
465}