diff --git a/Cargo.toml b/Cargo.toml index 58116f9..6ba3667 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -9,7 +9,7 @@ name = "ldns" path = "src/bin/ldns.rs" [dependencies] -clap = { version = "4", features = ["derive"] } +clap = { version = "4.3.4", features = ["derive"] } domain = "0.10.1" lexopt = "0.3.0" diff --git a/src/args.rs b/src/args.rs index 0659044..3883f51 100644 --- a/src/args.rs +++ b/src/args.rs @@ -1,3 +1,5 @@ +use crate::env::Env; + use super::commands::Command; use super::error::Error; @@ -9,8 +11,8 @@ pub struct Args { } impl Args { - pub fn execute(self) -> Result<(), Error> { - self.command.execute() + pub fn execute(self, env: impl Env) -> Result<(), Error> { + self.command.execute(env) } } diff --git a/src/bin/ldns.rs b/src/bin/ldns.rs index 6a8a5cc..beda509 100644 --- a/src/bin/ldns.rs +++ b/src/bin/ldns.rs @@ -9,14 +9,17 @@ use std::process::ExitCode; use dnst::try_ldns_compatibility; fn main() -> ExitCode { + let env = dnst::env::RealEnv; + let mut args = std::env::args_os(); args.next().unwrap(); - let args = try_ldns_compatibility(args).expect("ldns commmand is not recognized"); + let args = + try_ldns_compatibility(args).map(|args| args.expect("ldns commmand is not recognized")); - match args.execute() { + match args.and_then(|args| args.execute(&env)) { Ok(()) => ExitCode::SUCCESS, Err(err) => { - err.pretty_print(); + err.pretty_print(env); ExitCode::FAILURE } } diff --git a/src/commands/mod.rs b/src/commands/mod.rs index c602fe3..b7dbb3d 100644 --- a/src/commands/mod.rs +++ b/src/commands/mod.rs @@ -8,6 +8,7 @@ use std::str::FromStr; use nsec3hash::Nsec3Hash; +use crate::env::Env; use crate::Args; use super::error::Error; @@ -23,9 +24,9 @@ pub enum Command { } impl Command { - pub fn execute(self) -> Result<(), Error> { + pub fn execute(self, env: impl Env) -> Result<(), Error> { match self { - Self::Nsec3Hash(nsec3hash) => nsec3hash.execute(), + Self::Nsec3Hash(nsec3hash) => nsec3hash.execute(env), Self::Help(help) => help.execute(), } } diff --git a/src/commands/nsec3hash.rs b/src/commands/nsec3hash.rs index b673d50..1b175fe 100644 --- a/src/commands/nsec3hash.rs +++ b/src/commands/nsec3hash.rs @@ -1,3 +1,4 @@ +use crate::env::Env; use crate::error::Error; use clap::builder::ValueParser; use domain::base::iana::nsec3::Nsec3HashAlg; @@ -9,6 +10,7 @@ use lexopt::Arg; use octseq::OctetsBuilder; use ring::digest; use std::ffi::OsString; +use std::fmt::Write; use std::str::FromStr; use super::{parse_os, parse_os_with, LdnsCommand}; @@ -128,11 +130,13 @@ impl Nsec3Hash { } impl Nsec3Hash { - pub fn execute(self) -> Result<(), Error> { + pub fn execute(self, env: impl Env) -> Result<(), Error> { let hash = nsec3_hash(&self.name, self.algorithm, self.iterations, &self.salt) .to_string() .to_lowercase(); - println!("{}.", hash); + + let mut out = env.stdout(); + writeln!(out, "{}.", hash).unwrap(); Ok(()) } } @@ -179,3 +183,36 @@ where // For normal hash algorithms this should not fail. OwnerHash::from_octets(h.as_ref().to_vec()).expect("should not fail") } + +#[cfg(test)] +mod test { + use crate::env::fake::FakeCmd; + + #[test] + fn dnst_parse() { + let cmd = FakeCmd::new(["dnst", "nsec3-hash"]); + + assert!(cmd.parse().is_err()); + assert!(cmd.args(["-a"]).parse().is_err()); + } + + #[test] + fn dnst_run() { + let cmd = FakeCmd::new(["dnst", "nsec3-hash"]); + + let res = cmd.run(); + assert_eq!(res.exit_code, 2); + + let res = cmd.args(["example.test"]).run(); + assert_eq!(res.exit_code, 0); + assert_eq!(res.stdout, "o09614ibh1cq1rcc86289olr22ea0fso.\n") + } + + #[test] + fn ldns_parse() { + let cmd = FakeCmd::new(["ldns-nsec3-hash"]); + + assert!(cmd.parse().is_err()); + assert!(cmd.args(["-a"]).parse().is_err()); + } +} diff --git a/src/env/fake.rs b/src/env/fake.rs new file mode 100644 index 0000000..f02c66a --- /dev/null +++ b/src/env/fake.rs @@ -0,0 +1,134 @@ +use std::ffi::OsString; +use std::fmt; +use std::sync::Arc; +use std::sync::Mutex; + +use crate::{error::Error, parse_args, run, Args}; + +use super::Env; + +/// A command to run in a [`FakeEnv`] +/// +/// This is used for testing the utilities, running the real code in a fake +/// environment. +#[derive(Clone)] +pub struct FakeCmd { + /// The command to run, including `argv[0]` + cmd: Vec, +} + +/// The result of running a [`FakeCmd`] +/// +/// The fields are public to allow for easy assertions in tests. +pub struct FakeResult { + pub exit_code: u8, + pub stdout: String, + pub stderr: String, +} + +/// An environment that mocks interaction with the outside world +pub struct FakeEnv { + /// Description of the command being run + pub cmd: FakeCmd, + + /// The mocked stdout + pub stdout: FakeStream, + + /// The mocked stderr + pub stderr: FakeStream, + // pub stelline: Option, + // pub curr_step_value: Option>, +} + +impl Env for FakeEnv { + fn args_os(&self) -> impl Iterator { + self.cmd.cmd.iter().map(Into::into) + } + + fn stdout(&self) -> impl fmt::Write { + self.stdout.clone() + } + + fn stderr(&self) -> impl fmt::Write { + self.stderr.clone() + } +} + +impl FakeCmd { + /// Construct a new [`FakeCmd`] with a given command. + /// + /// The command can consist of multiple strings to specify a subcommand. + pub fn new>(cmd: impl IntoIterator) -> Self { + Self { + cmd: cmd.into_iter().map(Into::into).collect(), + } + } + + /// Add arguments to a clone of the [`FakeCmd`] + /// + /// ```rust,ignore + /// let cmd = FakeCmd::new(["dnst"]) + /// let sub1 = cmd.args(["sub1"]); // dnst sub1 + /// let sub2 = cmd.args(["sub2"]); // dnst sub2 + /// let sub3 = sub2.args(["sub3"]); // dnst sub2 sub3 + /// ``` + pub fn args>(&self, args: impl IntoIterator) -> Self { + let mut new = self.clone(); + new.cmd.extend(args.into_iter().map(Into::into)); + new + } + + /// Parse the arguments of this [`FakeCmd`] and return the result + pub fn parse(&self) -> Result { + let env = FakeEnv { + cmd: self.clone(), + stdout: Default::default(), + stderr: Default::default(), + }; + parse_args(env) + } + + /// Run the [`FakeCmd`] in a [`FakeEnv`], returning a [`FakeResult`] + pub fn run(&self) -> FakeResult { + let env = FakeEnv { + cmd: self.clone(), + stdout: Default::default(), + stderr: Default::default(), + }; + + let exit_code = run(&env); + + FakeResult { + exit_code, + stdout: env.get_stdout(), + stderr: env.get_stderr(), + } + } +} + +impl FakeEnv { + pub fn get_stdout(&self) -> String { + self.stdout.0.lock().unwrap().clone() + } + + pub fn get_stderr(&self) -> String { + self.stderr.0.lock().unwrap().clone() + } +} + +/// A type to used to mock stdout and stderr +#[derive(Clone, Default)] +pub struct FakeStream(Arc>); + +impl fmt::Write for FakeStream { + fn write_str(&mut self, s: &str) -> fmt::Result { + self.0.lock().unwrap().push_str(s); + Ok(()) + } +} + +impl fmt::Display for FakeStream { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str(self.0.lock().unwrap().as_ref()) + } +} diff --git a/src/env/mod.rs b/src/env/mod.rs new file mode 100644 index 0000000..717554e --- /dev/null +++ b/src/env/mod.rs @@ -0,0 +1,57 @@ +use std::ffi::OsString; +use std::fmt; + +mod real; + +#[cfg(test)] +pub mod fake; + +pub use real::RealEnv; + +pub trait Env { + // /// Make a network connection + // fn make_connection(&self); + + // /// Make a new [`StubResolver`] + // fn make_stub_resolver(&self); + + /// Get an iterator over the command line arguments passed to the program + /// + /// Equivalent to [`std::env::args_os`] + fn args_os(&self) -> impl Iterator; + + /// Get a reference to stdout + /// + /// Equivalent to [`std::io::stdout`] + fn stdout(&self) -> impl fmt::Write; + + /// Get a reference to stderr + /// + /// Equivalent to [`std::io::stderr`] + fn stderr(&self) -> impl fmt::Write; + + // /// Get a reference to stdin + // fn stdin(&self) -> impl io::Read; +} + +impl Env for &E { + // fn make_connection(&self) { + // todo!() + // } + + // fn make_stub_resolver(&self) { + // todo!() + // } + + fn args_os(&self) -> impl Iterator { + (**self).args_os() + } + + fn stdout(&self) -> impl fmt::Write { + (**self).stdout() + } + + fn stderr(&self) -> impl fmt::Write { + (**self).stderr() + } +} diff --git a/src/env/real.rs b/src/env/real.rs new file mode 100644 index 0000000..69fb5e0 --- /dev/null +++ b/src/env/real.rs @@ -0,0 +1,34 @@ +use std::ffi::OsString; +use std::fmt; +use std::io; + +use super::Env; + +/// Use real I/O +pub struct RealEnv; + +impl Env for RealEnv { + fn args_os(&self) -> impl Iterator { + std::env::args_os() + } + + fn stdout(&self) -> impl fmt::Write { + FmtWriter(io::stdout()) + } + + fn stderr(&self) -> impl fmt::Write { + FmtWriter(io::stderr()) + } +} + +struct FmtWriter(T); + +impl fmt::Write for FmtWriter { + fn write_str(&mut self, s: &str) -> std::fmt::Result { + self.0.write_all(s.as_bytes()).map_err(|_| fmt::Error) + } + + fn write_fmt(&mut self, args: fmt::Arguments<'_>) -> fmt::Result { + self.0.write_fmt(args).map_err(|_| fmt::Error) + } +} diff --git a/src/error.rs b/src/error.rs index e2c2470..e800a59 100644 --- a/src/error.rs +++ b/src/error.rs @@ -1,18 +1,18 @@ -use std::{error, fmt, io}; +use crate::env::Env; +use std::fmt::{self, Write}; +use std::{error, io}; //------------ Error --------------------------------------------------------- /// A program error. /// /// Such errors are highly likely to halt the program. -#[derive(Clone)] pub struct Error(Box); /// Information about an error. -#[derive(Clone)] struct Information { /// The primary error message. - primary: Box, + primary: PrimaryError, /// Layers of context to the error. /// @@ -20,13 +20,28 @@ struct Information { context: Vec>, } +#[derive(Debug)] +enum PrimaryError { + Clap(clap::Error), + Other(Box), +} + +impl fmt::Display for PrimaryError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + PrimaryError::Clap(e) => e.fmt(f), + PrimaryError::Other(e) => e.fmt(f), + } + } +} + //--- Interaction impl Error { /// Construct a new error from a string. pub fn new(error: &str) -> Self { Self(Box::new(Information { - primary: error.into(), + primary: PrimaryError::Other(error.into()), context: Vec::new(), })) } @@ -38,8 +53,21 @@ impl Error { } /// Pretty-print this error. - pub fn pretty_print(self) { + pub fn pretty_print(&self, env: impl Env) { use std::io::IsTerminal; + let mut err = env.stderr(); + + let error = match &self.0.primary { + // Clap errors are already styled. We don't want our own pretty + // styling around that and context does not make sense for command + // line arguments either. So we just print the styled string that + // clap produces and return. + PrimaryError::Clap(e) => { + let _ = writeln!(err, "{}", e.render().ansi()); + return; + } + PrimaryError::Other(error) => error, + }; // NOTE: This is a multicall binary, so argv[0] is necessary for // program operation. We would fail very early if it didn't exist. @@ -52,9 +80,24 @@ impl Error { "ERROR:" }; - eprint!("[{prog}] {error_marker} {}", self.0.primary); + let _ = write!(err, "[{prog}] {error_marker} {error}"); for context in &self.0.context { - eprint!("\n... while {context}"); + let _ = writeln!(err, "\n... while {context}"); + } + } + + pub fn exit_code(&self) -> u8 { + // Clap uses the exit code 2 and we want to keep that, but we aren't + // actually returning the clap error, so we replicate that behaviour + // here. + // + // Argument parsing errors from the ldns-xxx commands will not be clap + // errors and therefore be printed with an exit code of 1. This is + // expected because ldns also exits with 1. + if let PrimaryError::Clap(e) = &self.0.primary { + e.exit_code() as u8 + } else { + 1 } } } @@ -85,11 +128,20 @@ impl From for Error { } } +impl From for Error { + fn from(value: clap::Error) -> Self { + Self(Box::new(Information { + primary: PrimaryError::Clap(value), + context: Vec::new(), + })) + } +} + //--- Display, Debug impl fmt::Display for Error { fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result { - f.write_str(&self.0.primary) + self.0.primary.fmt(f) } } diff --git a/src/lib.rs b/src/lib.rs index a329971..8a14999 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,29 +1,53 @@ -use std::{ffi::OsString, path::Path}; +use std::ffi::OsString; +use std::path::Path; +use clap::Parser; use commands::{nsec3hash::Nsec3Hash, LdnsCommand}; +use env::Env; +use error::Error; pub use self::args::Args; pub mod args; pub mod commands; +pub mod env; pub mod error; -pub fn try_ldns_compatibility>(args: I) -> Option { +pub fn try_ldns_compatibility>( + args: I, +) -> Result, Error> { let mut args_iter = args.into_iter(); - let binary_path = args_iter.next()?; + let binary_path = args_iter.next().ok_or("Missing binary name")?; - let binary_name = Path::new(&binary_path).file_name()?.to_str()?; + let binary_name = Path::new(&binary_path) + .file_name() + .ok_or("Missing binary file name")? + .to_str() + .ok_or("Binary file name is not valid unicode")?; let res = match binary_name { "ldns-nsec3-hash" => Nsec3Hash::parse_ldns_args(args_iter), - _ => return None, + _ => return Ok(None), }; + res.map(Some) +} + +fn parse_args(env: impl Env) -> Result { + if let Some(args) = try_ldns_compatibility(env.args_os())? { + return Ok(args); + } + let args = Args::try_parse_from(env.args_os())?; + Ok(args) +} + +pub fn run(env: impl Env) -> u8 { + let res = parse_args(&env).and_then(|args| args.execute(&env)); match res { - Ok(args) => Some(args), + Ok(()) => 0, Err(err) => { - err.pretty_print(); - std::process::exit(1) + err.pretty_print(&env); + err.exit_code() } } } diff --git a/src/main.rs b/src/main.rs index 4a14114..b368f32 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,18 +1,6 @@ use std::process::ExitCode; -use clap::Parser; - fn main() -> ExitCode { - // If none of the ldns-* tools matched, then we continue with clap - // argument parsing. - let env_args = std::env::args_os(); - let args = dnst::try_ldns_compatibility(env_args).unwrap_or_else(dnst::Args::parse); - - match args.execute() { - Ok(()) => ExitCode::SUCCESS, - Err(err) => { - err.pretty_print(); - ExitCode::FAILURE - } - } + let env = dnst::env::RealEnv; + dnst::run(env).into() }