From 32ca8f7ca073882ce22e4e3005dd540f297ef3ea Mon Sep 17 00:00:00 2001 From: Jannik Peters Date: Thu, 21 Nov 2024 16:40:41 +0100 Subject: [PATCH] Fix writing to stdout or file --- src/commands/signzone.rs | 57 ++++++++++++++++++++++++++++------------ src/env/mod.rs | 5 ++++ src/error.rs | 6 +++++ 3 files changed, 51 insertions(+), 17 deletions(-) diff --git a/src/commands/signzone.rs b/src/commands/signzone.rs index 5b1fc7c..edf345e 100644 --- a/src/commands/signzone.rs +++ b/src/commands/signzone.rs @@ -3,8 +3,10 @@ use core::str::FromStr; use std::cmp::min; use std::ffi::OsString; +use std::fmt; +use std::fmt::Write; use std::fs::File; -use std::io::Write; +use std::io; use std::path::{Path, PathBuf}; // TODO: use a re-export from domain? @@ -25,7 +27,7 @@ use domain::zonetree::types::StoredRecordData; use domain::zonetree::{StoredName, StoredRecord}; use lexopt::Arg; -use crate::env::Env; +use crate::env::{Env, Stream}; use crate::error::Error; use super::nsec3hash::Nsec3Hash; @@ -396,17 +398,13 @@ impl SignZone { }; let mut writer = if out_file.as_os_str() == "-" { - // Box::new(env.stdout()) as Box - // FIXME: env.stdout() uses impl fmt::Write, but because of - // domain::sign::records::SortedRecords::write_with_comments() - // we need io::Write here. - todo!() + FileOrStdout::Stdout(env.stdout()) } else { - Box::new(File::create(env.in_cwd(&out_file))?) as Box + FileOrStdout::File(File::create(env.in_cwd(&out_file))?) }; // Read the zone file. - let mut records = self.load_zone()?; + let mut records = self.load_zone(&env)?; // Import the specified keys. let mut keys = vec![]; @@ -471,12 +469,12 @@ impl SignZone { // " ;{... .}" but I find the spacing ugly and // would prefer for dnst to output " ; {... . }" // instead. - writer.write_all(b" ;{ flags: ")?; + writer.write_str(" ;{ flags: ")?; if nsec3.opt_out() { - writer.write_all(b"optout")?; + writer.write_str("optout")?; } else { - writer.write_all(b"-")?; + writer.write_str("-")?; } let next_owner_hash_hex = format!("{}", nsec3.next_owner()); @@ -505,9 +503,9 @@ impl SignZone { ZoneRecordData::Dnskey(dnskey) => { writer.write_fmt(format_args!(" ;{{id = {}", dnskey.key_tag()))?; if dnskey.is_secure_entry_point() { - writer.write_all(b" (ksk)")?; + writer.write_str(" (ksk)")?; } else if dnskey.is_zone_key() { - writer.write_all(b" (zsk)")?; + writer.write_str(" (zsk)")?; } let owner = r.owner().clone(); let dnskey = dnskey.clone(); @@ -522,9 +520,11 @@ impl SignZone { Ok(()) } - fn load_zone(&self) -> Result, Error> { - // TODO: load file with env.in_cwd(zonefile_path)? - let mut zone_file = File::open(&self.zonefile_path)?; + fn load_zone( + &self, + env: &impl Env, + ) -> Result, Error> { + let mut zone_file = File::open(env.in_cwd(&self.zonefile_path))?; let mut reader = inplace::Zonefile::load(&mut zone_file).unwrap(); if let Some(origin) = &self.origin { reader.set_origin(origin.clone()); @@ -707,3 +707,26 @@ enum SigningMode { // /// Only sign zone records, assume they are already hashed. // SignOnly, } + +//------------ FileOrStdout -------------------------------------------------- + +enum FileOrStdout { + File(T), + Stdout(Stream), +} + +impl fmt::Write for FileOrStdout { + fn write_str(&mut self, s: &str) -> std::fmt::Result { + match self { + FileOrStdout::File(f) => f.write_all(s.as_bytes()).map_err(|_| fmt::Error), + FileOrStdout::Stdout(o) => Ok(o.write_str(s)), + } + } + + fn write_fmt(&mut self, args: fmt::Arguments<'_>) -> fmt::Result { + match self { + FileOrStdout::File(f) => f.write_fmt(args).map_err(|_| fmt::Error), + FileOrStdout::Stdout(o) => Ok(o.write_fmt(args)), + } + } +} diff --git a/src/env/mod.rs b/src/env/mod.rs index dd62dcb..d70336e 100644 --- a/src/env/mod.rs +++ b/src/env/mod.rs @@ -55,6 +55,11 @@ impl Stream { // hard anyway. self.0.write_fmt(args).unwrap(); } + + pub fn write_str(&mut self, s: &str) { + // Same as with write_fmt... + self.0.write_str(s).unwrap(); + } } impl Env for &E { diff --git a/src/error.rs b/src/error.rs index ee7ad6b..24d5034 100644 --- a/src/error.rs +++ b/src/error.rs @@ -117,6 +117,12 @@ impl From for Error { } } +impl From for Error { + fn from(error: fmt::Error) -> Self { + Self::new(&error.to_string()) + } +} + impl From for Error { fn from(error: io::Error) -> Self { Self::new(&error.to_string())