diff --git a/Cargo.toml b/Cargo.toml index f9020ae5..e83152c1 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,5 +1,3 @@ [workspace] members = ["domain", "domain-core", "domain-resolv", "domain-sign", "domain-tsig", "domain-validate", "interop"] -[patch.crates-io] -ring = { git = "https://github.com/andrewtj/ring.git", rev = "eedb0514fb73e1034eefbec10140af8a51effff3"} \ No newline at end of file diff --git a/domain-core/src/message.rs b/domain-core/src/message.rs index 7c597526..ddebbdc2 100644 --- a/domain-core/src/message.rs +++ b/domain-core/src/message.rs @@ -83,9 +83,9 @@ impl Message { /// # Header Section /// impl> Message { - /// Returns a reference to the message header. - pub fn header(&self) -> &Header { - Header::for_message_slice(self.as_slice()) + /// Returns a the message header. + pub fn header(&self) -> Header { + *Header::for_message_slice(self.as_slice()) } /// Returns a mutable reference to the message header. @@ -94,18 +94,14 @@ impl> Message { Header::for_message_slice_mut(self.as_slice_mut()) } - /// Returns a reference the header counts of the message. - pub fn header_counts(&self) -> &HeaderCounts { - HeaderCounts::for_message_slice(self.as_slice()) + /// Returns the header counts of the message. + pub fn header_counts(&self) -> HeaderCounts { + *HeaderCounts::for_message_slice(self.as_slice()) } - /// Returns a mutable reference to the header counts. - /// - /// Since you can quite effectively break the message with this, it is - /// private. - pub fn header_counts_mut(&mut self) -> &mut HeaderCounts - where Octets: AsMut<[u8]> { - HeaderCounts::for_message_slice_mut(self.as_slice_mut()) + /// Returns the entire header section. + pub fn header_section(&self) -> HeaderSection { + *HeaderSection::for_message_slice(&self.as_slice()) } /// Returns whether the rcode is NoError. @@ -424,6 +420,11 @@ impl QuestionSection { } } + /// Returns the current position relative to the beginning of the message. + pub fn pos(&self) -> usize { + self.parser.pos() + } + /// Proceeds to the answer section. /// /// Skips over any remaining questions and then converts itself into the @@ -556,6 +557,11 @@ impl RecordSection { } } + /// Returns the current position relative to the beginning of the message. + pub fn pos(&self) -> usize { + self.parser.pos() + } + /// Trades `self` in for an iterator limited to a concrete record type. /// /// The record type is given through its record data type. Since the data diff --git a/domain-core/src/message_builder.rs b/domain-core/src/message_builder.rs index ccfb6212..4e371a41 100644 --- a/domain-core/src/message_builder.rs +++ b/domain-core/src/message_builder.rs @@ -8,10 +8,12 @@ use core::ops::{Deref, DerefMut}; #[cfg(feature = "bytes")] use bytes::BytesMut; use unwrap::unwrap; use crate::header::{Header, HeaderCounts, HeaderSection}; -use crate::iana::{OptionCode, OptRcode}; +use crate::iana::{OptionCode, OptRcode, Rcode, Rtype}; use crate::message::Message; use crate::name::{ToDname, Label}; -use crate::octets::{Compose, IntoOctets, Octets64, OctetsBuilder, ShortBuf}; +use crate::octets::{ + Compose, IntoOctets, Octets64, OctetsBuilder, OctetsRef, ShortBuf +}; use crate::opt::{OptHeader, OptData}; use crate::question::Question; use crate::rdata::RecordData; @@ -68,6 +70,45 @@ impl MessageBuilder> { } impl MessageBuilder { + /// Starts creating an answer for the given message. + /// + /// Specifically, this sets the ID, QR, OPCODE, RD, and RCODE fields + /// in the header and attempts to push the message’s questions to the + /// builder. If iterating of the questions fails, it adds what it can. + pub fn start_answer( + mut self, + msg: &Message, + rcode: Rcode, + ) -> Result, ShortBuf> + where Octets: AsRef<[u8]>, for<'a> &'a Octets: OctetsRef { + { + let header = self.header_mut(); + header.set_id(msg.header().id()); + header.set_qr(true); + header.set_opcode(msg.header().opcode()); + header.set_rd(msg.header().rd()); + header.set_rcode(rcode); + } + let mut builder = self.question(); + for item in msg.question() { + if let Ok(item) = item { + builder.push(item)?; + } + } + Ok(builder.answer()) + } + + /// Creates an AXFR request for the given domain. + pub fn request_axfr( + mut self, + apex: N + ) -> Result, ShortBuf> { + self.header_mut().set_random_id(); + let mut builder = self.question(); + builder.push((apex, Rtype::Axfr))?; + Ok(builder.answer()) + } + pub fn question(self) -> QuestionBuilder { QuestionBuilder::new(self) } @@ -96,6 +137,10 @@ impl MessageBuilder { &mut self.target } + pub fn as_slice(&self) -> &[u8] { + self.as_target().as_ref() + } + pub fn as_message(&self) -> Message<&[u8]> where Target: AsRef<[u8]> { unsafe { Message::from_octets_unchecked(self.target.as_ref()) } @@ -106,16 +151,16 @@ impl MessageBuilder { unsafe { Message::from_octets_unchecked(self.target.into_octets()) } } - pub fn header(&self) -> &Header { - Header::for_message_slice(self.target.as_ref()) + pub fn header(&self) -> Header { + *Header::for_message_slice(self.target.as_ref()) } pub fn header_mut(&mut self) -> &mut Header { Header::for_message_slice_mut(self.target.as_mut()) } - pub fn counts(&self) -> &HeaderCounts { - HeaderCounts::for_message_slice(self.target.as_ref()) + pub fn counts(&self) -> HeaderCounts { + *HeaderCounts::for_message_slice(self.target.as_ref()) } fn counts_mut(&mut self) -> &mut HeaderCounts { @@ -508,7 +553,7 @@ where Target: OctetsBuilder { fn push(&mut self, record: R) -> Result<(), ShortBuf> where N: ToDname, D: RecordData, R: Into> { record.into().compose(self.as_target_mut())?; - self.counts_mut().inc_ancount(); + self.counts_mut().inc_arcount(); Ok(()) } } @@ -592,7 +637,7 @@ impl OptBuilder { } pub fn rcode(&self) -> OptRcode { - self.opt_header().rcode(*self.header()) + self.opt_header().rcode(self.header()) } pub fn set_rcode(&mut self, rcode: OptRcode) { diff --git a/domain-core/src/octets.rs b/domain-core/src/octets.rs index b71ed4f6..b2444a7b 100644 --- a/domain-core/src/octets.rs +++ b/domain-core/src/octets.rs @@ -676,6 +676,9 @@ octets_array!(pub Octets2048 => 2048); octets_array!(pub Octets4096 => 4096); +#[cfg(feature = "smallvec")] +pub type OctetsVec = SmallVec<[u8; 24]>; + //------------ ShortBuf ------------------------------------------------------ /// An attempt was made to go beyond the end of a buffer. diff --git a/domain-core/src/rdata/rfc2845.rs b/domain-core/src/rdata/rfc2845.rs index 2ac1a2f8..ec0338fb 100644 --- a/domain-core/src/rdata/rfc2845.rs +++ b/domain-core/src/rdata/rfc2845.rs @@ -157,7 +157,6 @@ impl Tsig { /// /// [`fudge`]: #method.fudge /// [`time_signed`]: #method.time_signed - #[cfg(feature = "chrono")] pub fn is_valid_now(&self) -> bool { Time48::now().eq_fudged(self.time_signed, self.fudge.into()) } diff --git a/domain-core/src/serial.rs b/domain-core/src/serial.rs index 1fb9090f..12d379d2 100644 --- a/domain-core/src/serial.rs +++ b/domain-core/src/serial.rs @@ -6,7 +6,7 @@ //! [`Serial`]: struct.Serial.html use core::{cmp, fmt, str}; -#[cfg(feature = "bytes")] use chrono::{Utc, TimeZone}; +use chrono::{DateTime, Utc, TimeZone}; use crate::cmp::CanonicalOrd; #[cfg(feature = "bytes")] use crate::master::scan::{ CharSource, Scan, ScanError, Scanner, SyntaxError @@ -42,7 +42,6 @@ pub struct Serial(pub u32); impl Serial { /// Returns a serial number for the current Unix time. - #[cfg(feature = "chrono")] pub fn now() -> Self { Utc::now().into() } @@ -88,7 +87,7 @@ impl Serial { /// In RRSIG records, the expiration and inception time is given as /// serial values. Their master file format can either be the signature /// value or a specific date in `YYYYMMDDHHmmSS` format. - #[cfg(all(feature="bytes"))] + #[cfg(feature="bytes")] pub fn scan_rrsig( scanner: &mut Scanner ) -> Result { @@ -182,7 +181,6 @@ impl From for u32 { } } -#[cfg(feature = "chrono")] impl From> for Serial { fn from(value: DateTime) -> Self { let mut value = value.timestamp(); diff --git a/domain-sign/Cargo.toml b/domain-sign/Cargo.toml index 1932cee9..cfd19830 100644 --- a/domain-sign/Cargo.toml +++ b/domain-sign/Cargo.toml @@ -19,13 +19,13 @@ path = "src/lib.rs" bytes = "0.4" derive_more = "^0.15" openssl = { version = "^0.10", optional = true } -ring = { version = "0.15.0-alpha", optional = true } +ring = { version = "0.16", optional = true } unwrap = "^1.2" [dependencies.domain-core] path = "../domain-core" -version = "0.4.1" +version = "0.5.0-pre" [features] ringsigner = ["ring"] -default = ["ringsigner"] \ No newline at end of file +default = ["ringsigner"] diff --git a/domain-tsig/Cargo.toml b/domain-tsig/Cargo.toml index 1b600e59..2cfc8072 100644 --- a/domain-tsig/Cargo.toml +++ b/domain-tsig/Cargo.toml @@ -17,10 +17,13 @@ path = "src/lib.rs" [dependencies] bytes = "^0.4" -derive_more = "^0.14" -ring = "0.15.0-alpha" +derive_more = "^0.99" +ring = "0.16" +smallvec = "1.0" +unwrap = "1.2" [dependencies.domain-core] path = "../domain-core" -version = "0.4.1" +version = "0.5.0-pre" +features = ["std", "smallvec"] diff --git a/domain-tsig/src/lib.rs b/domain-tsig/src/lib.rs index ec906bbc..b173fde1 100644 --- a/domain-tsig/src/lib.rs +++ b/domain-tsig/src/lib.rs @@ -57,15 +57,18 @@ use std::collections::HashMap; use bytes::{BigEndian, ByteOrder, Bytes, BytesMut}; use derive_more::Display; use ring::{constant_time, hmac, rand, hkdf::KeyType}; +use unwrap::unwrap; +use domain_core::header::HeaderSection; use domain_core::iana::{Class, Rcode, TsigRcode}; use domain_core::message::Message; use domain_core::message_builder::{ - AdditionalBuilder, MessageBuilder, SectionBuilder, RecordSectionBuilder + AdditionalBuilder, MessageBuilder, RecordSectionBuilder }; use domain_core::name::{ - Dname, Label, ParsedDname, ParsedDnameError, ToDname, ToLabelIter + Dname, Label, ParsedDname, ToDname, ToLabelIter }; -use domain_core::parse::ShortBuf; +use domain_core::octets::{OctetsBuilder, OctetsRef, OctetsVec, ShortBuf}; +use domain_core::parse::ParseError; use domain_core::record::Record; use domain_core::rdata::rfc2845::{Time48, Tsig}; @@ -102,7 +105,7 @@ pub struct Key { key: hmac::Key, /// The name of the key as a domain name. - name: Dname, + name: Dname, /// Minimum length of received signatures. /// @@ -141,7 +144,7 @@ impl Key { pub fn new( algorithm: Algorithm, key: &[u8], - name: Dname, + name: Dname, min_mac_len: Option, signing_len: Option ) -> Result { @@ -166,7 +169,7 @@ impl Key { pub fn generate( algorithm: Algorithm, rng: &dyn rand::SecureRandom, - name: Dname, + name: Dname, min_mac_len: Option, signing_len: Option ) -> Result<(Self, Bytes), GenerateKeyError> { @@ -240,7 +243,7 @@ impl Key { } /// Returns a reference to the name of this key. - pub fn name(&self) -> &Dname { + pub fn name(&self) -> &Dname { &self.name } @@ -260,9 +263,11 @@ impl Key { } /// Checks whether the key in the record is this key. - fn check_tsig( - &self, tsig: &Record> - ) -> Result<(), ValidationError> { + fn check_tsig( + &self, + tsig: &MessageTsig + ) -> Result<(), ValidationError> + where for<'o> &'o Octets: OctetsRef { if *tsig.owner() != self.name || *tsig.data().algorithm() != self.algorithm().to_dname() { @@ -305,17 +310,14 @@ impl Key { /// /// The method fails if the TSIG record doesn’t fit into the message /// anymore, in which case the builder is returned unharmed. - fn complete_message( + fn complete_message( &self, - mut message: AdditionalBuilder, + message: &mut AdditionalBuilder, variables: &Variables, mac: &[u8], - ) -> Result { + ) -> Result<(), ShortBuf> { let id = message.header().id(); - match message.push(variables.to_tsig(self, mac, id)) { - Ok(()) => Ok(message.freeze()), - Err(_) => Err(message) - } + variables.push_tsig(self, mac, id, message) } } @@ -378,14 +380,18 @@ impl + Clone> KeyStore for K { } } -impl KeyStore for HashMap<(Dname, Algorithm), K, S> -where K: AsRef + Clone, S: hash::BuildHasher { +impl KeyStore for HashMap<(Dname, Algorithm), K, S> +where + K: AsRef + Clone, + S: hash::BuildHasher +{ type Key = K; fn get_key( &self, name: &N, algorithm: Algorithm ) -> Option { - let name = name.to_name(); // XXX This seems a bit wasteful. + // XXX This seems a bit wasteful. + let name = unwrap!(name.to_dname::()); self.get(&(name, algorithm)).cloned() } } @@ -430,9 +436,10 @@ impl> ClientTransaction { /// recommended default value for _fudge:_ 300 seconds. /// /// [`request_with_fudge`]: #method.request_with_fudge - pub fn request( - key: K, message: AdditionalBuilder - ) -> Result<(Message, Self), AdditionalBuilder> { + pub fn request( + key: K, + message: &mut AdditionalBuilder + ) -> Result { Self::request_with_fudge(key, message, 300) } @@ -455,20 +462,21 @@ impl> ClientTransaction { /// the untouched message. /// /// [`request`]: #method.request - pub fn request_with_fudge( - key: K, message: AdditionalBuilder, fudge: u16 - ) -> Result<(Message, Self), AdditionalBuilder> { + pub fn request_with_fudge( + key: K, + message: &mut AdditionalBuilder, + fudge: u16 + ) -> Result { let variables = Variables::new( Time48::now(), fudge, TsigRcode::NoError, None ); let (mut context, mac) = SigningContext::request( - key, message.so_far(), &variables + key, message.as_slice(), None, &variables ); let mac = context.key().signature_slice(&mac); context.apply_signature(mac); - context.key().complete_message(message, &variables, mac).map(|message| { - (message, ClientTransaction { context }) - }) + context.key().complete_message(message, &variables, mac)?; + Ok(ClientTransaction { context }) } /// Validates an answer. @@ -481,22 +489,32 @@ impl> ClientTransaction { /// whether this record is a correct record for this transaction and if /// it correctly signs the answer for this transaction. If any of this /// fails, returns an error. - pub fn answer( - &self, message: &mut Message - ) -> Result<(), ValidationError> { - let (variables, tsig) = match self.context.extract_answer_tsig( - message)? { + pub fn answer( + &self, message: &mut Message + ) -> Result<(), ValidationError> + where + Octets: AsRef<[u8]> + AsMut<[u8]>, + for<'a> &'a Octets: OctetsRef + { + let tsig = match self.context.get_answer_tsig(message)? { Some(some) => some, None => return Err(ValidationError::ServerUnsigned) }; - + let mut header = message.header_section(); + header.header_mut().set_id(tsig.data().original_id()); + header.counts_mut().dec_arcount(); let signature = self.context.answer( - message.as_slice(), &variables + header.as_slice(), + Some(&message.as_slice()[ + mem::size_of::()..tsig.start + ]), + &tsig.variables() ); self.context.key().compare_signatures( &signature, tsig.data().mac().as_ref() )?; self.context.check_answer_time(message, &tsig)?; + remove_tsig(tsig.into_original_id(), message); Ok(()) } @@ -541,10 +559,15 @@ impl> ServerTransaction { /// If anything is wrong with the message with regards to TSIG, the /// function returns the error message that should be returned to the /// client as the error case of the result. - pub fn request>( - store: &S, - message: &mut Message - ) -> Result, Message> { + pub fn request( + store: &Store, + message: &mut Message + ) -> Result, ServerError> + where + Store: KeyStore, + Octets: AsRef<[u8]> + AsMut<[u8]>, + for<'o> &'o Octets: OctetsRef + { SigningContext::server_request( store, message ).map(|context| context.map(|context| ServerTransaction { context })) @@ -562,9 +585,9 @@ impl> ServerTransaction { /// If appending the TSIG record fails, which can only happen if there /// isn’t enough space left, it returns the builder unchanged as the /// error case. - pub fn answer( - self, message: AdditionalBuilder - ) -> Result { + pub fn answer( + self, message: &mut AdditionalBuilder + ) -> Result<(), ShortBuf> { self.answer_with_fudge(message, 300) } @@ -576,16 +599,16 @@ impl> ServerTransaction { /// The default, suggested by the RFC, is 300. /// /// [`answer`]: #method.answer - pub fn answer_with_fudge( + pub fn answer_with_fudge( self, - message: AdditionalBuilder, + message: &mut AdditionalBuilder, fudge: u16 - ) -> Result { + ) -> Result<(), ShortBuf> { let variables = Variables::new( Time48::now(), fudge, TsigRcode::NoError, None ); let (mac, key) = self.context.final_answer( - message.so_far(), &variables + message.as_slice(), None, &variables ); let mac = key.as_ref().signature_slice(&mac); key.as_ref().complete_message(message, &variables, mac) @@ -643,9 +666,9 @@ impl> ClientSequence { /// returns the builder untouched as the error case. Otherwise, it will /// freeze the message and return both it and a new value of a client /// sequence. - pub fn request( - key: K, message: AdditionalBuilder - ) -> Result<(Message, Self), AdditionalBuilder> { + pub fn request( + key: K, message: &mut AdditionalBuilder + ) -> Result { Self::request_with_fudge(key, message, 300) } @@ -658,20 +681,19 @@ impl> ClientSequence { /// seconds. /// /// [`request`]: #method.request - pub fn request_with_fudge( - key: K, message: AdditionalBuilder, fudge: u16 - ) -> Result<(Message, Self), AdditionalBuilder> { + pub fn request_with_fudge( + key: K, message: &mut AdditionalBuilder, fudge: u16 + ) -> Result { let variables = Variables::new( Time48::now(), fudge, TsigRcode::NoError, None ); let (mut context, mac) = SigningContext::request( - key, message.so_far(), &variables + key, message.as_slice(), None, &variables ); let mac = context.key().signature_slice(&mac); context.apply_signature(mac); - context.key().complete_message(message, &variables, mac).map(|message| { - (message, ClientSequence { context, first: true, unsigned: 0 }) - }) + context.key().complete_message(message, &variables, mac)?; + Ok(ClientSequence { context, first: true, unsigned: 0 }) } /// Validates an answer. @@ -682,10 +704,14 @@ impl> ClientSequence { /// /// If it doesn’t or if there had been more than 99 unsigned messages in /// the sequence since the last signed one, returns an error. - pub fn answer( + pub fn answer( &mut self, - message: &mut Message - ) -> Result<(), ValidationError> { + message: &mut Message + ) -> Result<(), ValidationError> + where + Octets: AsRef<[u8]> + AsMut<[u8]>, + for<'a> &'a Octets: OctetsRef + { if self.first { self.answer_first(message) } @@ -711,18 +737,27 @@ impl> ClientSequence { } /// Checks the first answer in the sequence. - fn answer_first( + fn answer_first( &mut self, - message: &mut Message - ) -> Result<(), ValidationError> { - let (variables, tsig) = match self.context.extract_answer_tsig( - message)? { + message: &mut Message + ) -> Result<(), ValidationError> + where + Octets: AsRef<[u8]> + AsMut<[u8]>, + for<'a> &'a Octets: OctetsRef + { + let tsig = match self.context.get_answer_tsig(message)? { Some(some) => some, None => return Err(ValidationError::ServerUnsigned) }; - + let mut header = message.header_section(); + header.header_mut().set_id(tsig.data().original_id()); + header.counts_mut().dec_arcount(); let signature = self.context.first_answer( - message.as_slice(), &variables + header.as_slice(), + Some(&message.as_slice()[ + mem::size_of::()..tsig.start + ]), + &tsig.variables() ); self.context.key().compare_signatures( &signature, tsig.data().mac().as_ref() @@ -730,39 +765,52 @@ impl> ClientSequence { self.context.apply_signature(tsig.data().mac().as_ref()); self.context.check_answer_time(message, &tsig)?; self.first = false; + remove_tsig(tsig.into_original_id(), message); Ok(()) } /// Checks any subsequent answer in the sequence. - fn answer_subsequent( + fn answer_subsequent( &mut self, - message: &mut Message, - ) -> Result<(), ValidationError> { - match self.context.extract_answer_tsig(message)? { - Some((variables, tsig)) => { - // Check the MAC. - let signature = self.context.signed_subsequent( - message.as_slice(), &variables - ); - self.context.key().compare_signatures( - &signature, tsig.data().mac().as_ref() - )?; - self.context.apply_signature(tsig.data().mac().as_ref()); - self.context.check_answer_time(message, &tsig)?; - self.unsigned = 0; - Ok(()) - } + message: &mut Message, + ) -> Result<(), ValidationError> + where + Octets: AsRef<[u8]> + AsMut<[u8]>, + for<'a> &'a Octets: OctetsRef + { + let tsig = match self.context.get_answer_tsig(message)? { + Some(tsig) => tsig, None => { if self.unsigned < 99 { self.context.unsigned_subsequent(message.as_slice()); self.unsigned += 1; - Ok(()) + return Ok(()) } else { - Err(ValidationError::TooManyUnsigned) + return Err(ValidationError::TooManyUnsigned) } } - } + }; + + // Check the MAC. + let mut header = message.header_section(); + header.header_mut().set_id(tsig.data().original_id()); + header.counts_mut().dec_arcount(); + let signature = self.context.signed_subsequent( + header.as_slice(), + Some(&message.as_slice()[ + mem::size_of::()..tsig.start + ]), + &tsig.variables() + ); + self.context.key().compare_signatures( + &signature, tsig.data().mac().as_ref() + )?; + self.context.apply_signature(tsig.data().mac().as_ref()); + self.context.check_answer_time(message, &tsig)?; + self.unsigned = 0; + remove_tsig(tsig.into_original_id(), message); + Ok(()) } /// Returns a reference to the transaction’s key. @@ -813,10 +861,15 @@ impl> ServerSequence { /// If anything is wrong with the message with regards to TSIG, the /// function returns the error message that should be returned to the /// client as the error case of the result. - pub fn request>( - store: &S, - message: &mut Message - ) -> Result, Message> { + pub fn request( + store: &Store, + message: &mut Message + ) -> Result, ServerError> + where + Store: KeyStore, + Octets: AsRef<[u8]> + AsMut<[u8]>, + for<'o> &'o Octets: OctetsRef + { SigningContext::server_request( store, message ).map(|context| context.map(|context| { @@ -831,9 +884,9 @@ impl> ServerSequence { /// it attempts to add a TSIG record to the additional section, if that /// fails because there wasn’t enough space in the builder, returns the /// unchanged builder as an error. - pub fn answer( - &mut self, message: AdditionalBuilder - ) -> Result { + pub fn answer( + &mut self, message: &mut AdditionalBuilder + ) -> Result<(), ShortBuf> { self.answer_with_fudge(message, 300) } @@ -842,24 +895,25 @@ impl> ServerSequence { /// This is nearly identical to [`answer`] except that it allows to /// specify the ‘fudge’ which declares the number of seconds the /// receiver’s clock may be off from this systems current time. - pub fn answer_with_fudge( + pub fn answer_with_fudge( &mut self, - message: AdditionalBuilder, + message: &mut AdditionalBuilder, fudge: u16 - ) -> Result { + ) -> Result<(), ShortBuf> { let variables = Variables::new( Time48::now(), fudge, TsigRcode::NoError, None ); let mac = if self.first { self.first = false; - self.context.first_answer(message.so_far(), &variables) + self.context.first_answer(message.as_slice(), None, &variables) } else { - self.context.signed_subsequent(message.so_far(), &variables) + self.context.signed_subsequent( + message.as_slice(), None, &variables + ) }; let mac = self.key().signature_slice(&mac); self.key().complete_message(message, &variables, mac) - .map_err(|_| ShortBuf) } /// Returns a reference to the transaction’s key. @@ -907,14 +961,19 @@ impl> SigningContext { /// correctly signed with a known key. Returns `Ok(None)` if there was /// no TSIG record at all. Returns an error with a message to be returned /// to the client otherwise. - fn server_request>( - store: &S, - message: &mut Message - ) -> Result, Message> { + fn server_request( + store: &Store, + message: &mut Message + ) -> Result, ServerError> + where + Store: KeyStore, + Octets: AsRef<[u8]> + AsMut<[u8]>, + for<'a> &'a Octets: OctetsRef + { // 4.5 Server TSIG checks // // First, do we have a valid TSIG? - let tsig = match extract_tsig(message) { + let tsig = match MessageTsig::from_message(message) { Some(tsig) => tsig, None => return Ok(None) }; @@ -922,35 +981,33 @@ impl> SigningContext { // 4.5.1. KEY check and error handling let algorithm = match Algorithm::from_dname(tsig.data().algorithm()) { Some(algorithm) => algorithm, - None => { - return Err( - Self::unsigned_error(message, &tsig, TsigRcode::BadKey) - ) - } + None => return Err(ServerError::unsigned(TsigRcode::BadKey)), }; let key = match store.get_key(tsig.owner(), algorithm) { Some(key) => key, - None => { - return Err( - Self::unsigned_error(message, &tsig, TsigRcode::BadKey) - ) - } + None => return Err(ServerError::unsigned(TsigRcode::BadKey)), }; - let variables = Variables::from_tsig(&tsig); + let variables = tsig.variables(); // 4.5.3 MAC check // // Contrary to RFC 2845, this must be done before the time check. - update_id(message, &tsig); + let mut header = message.header_section(); + header.header_mut().set_id(tsig.data().original_id()); + header.counts_mut().dec_arcount(); let (mut context, signature) = Self::request( - key, message.as_slice(), &variables + key, header.as_slice(), + Some(&message.as_slice()[ + mem::size_of::()..tsig.start + ]), + &variables ); let res = context.key.as_ref().compare_signatures( &signature, tsig.data().mac().as_ref() ); if let Err(err) = res { - return Err(Self::unsigned_error(message, &tsig, match err { + return Err(ServerError::unsigned(match err { ValidationError::BadTrunc => TsigRcode::BadTrunc, ValidationError::BadKey => TsigRcode::BadKey, _ => TsigRcode::FormErr, @@ -965,22 +1022,17 @@ impl> SigningContext { // Note that we are not doing the caching of the most recent // time_signed because, well, that’ll require mutexes and stuff. if !tsig.data().is_valid_now() { - let mut response = MessageBuilder::new_udp(); - response.start_answer(message, Rcode::NotAuth); - return Err( - // unwrap: answer should always fit. - context.signed_error( - response.additional(), - &Variables::new( - variables.time_signed, - variables.fudge, - TsigRcode::BadTime, - Some(Time48::now()) - ) + return Err(ServerError::signed( + context, + Variables::new( + variables.time_signed, + variables.fudge, + TsigRcode::BadTime, + Some(Time48::now()) ) - ) + )) } - + remove_tsig(tsig.into_original_id(), message); Ok(Some(context)) } @@ -989,25 +1041,23 @@ impl> SigningContext { /// This is the first part of the code shared by the various answer /// functions of `ClientTransaction` and `ClientSequence`. It does /// everything that needs to be done before actually verifying the - /// signature: Extract the TSIG record, handle unsigned errors, check - /// that the key and algorithm correspond to our key and algorithm, and - /// update the message ID if necessary. + /// signature: Find the TSIG record, handle unsigned errors, check + /// that the key and algorithm correspond to our key and algorithm. /// /// Since there may be unsigned messages in client sequences, returns /// `Ok(None)` if there is no TSIG at all. Otherwise, if all steps /// succeed, returns the TSIG variables and the TSIG record. If there /// is an error, returns that. - #[allow(clippy::type_complexity)] - fn extract_answer_tsig( + /// + /// Because the returned TSIG record references the message, so it will + /// later have to have the TSIG record stripped off and the ID updated. + fn get_answer_tsig<'a, Octets>( &self, - message: &mut Message - ) -> Result< - Option<(Variables, Record>)>, - ValidationError - > - { + message: &'a Message + ) -> Result>, ValidationError> + where Octets: AsRef<[u8]>, for<'o> &'o Octets: OctetsRef { // Extract TSIG or bail out. - let tsig = match extract_tsig(message) { + let tsig = match MessageTsig::from_message(message) { Some(tsig) => tsig, None => return Ok(None) }; @@ -1025,9 +1075,7 @@ impl> SigningContext { // Check that the server used the correct key and algorithm. self.key().check_tsig(&tsig)?; - // Fix up message and return. - update_id(message, &tsig); - Ok(Some((Variables::from_tsig(&tsig), tsig))) + Ok(Some(tsig)) } /// Checks the timing values of an answer TSIG. @@ -1036,11 +1084,12 @@ impl> SigningContext { /// answer methods of `ClientTransaction` and `ClientSequence`. It /// checks for timing errors reported by the server as well as the /// time signed in the signature. - fn check_answer_time( + fn check_answer_time<'a, Octets>( &self, - message: &Message, - tsig: &Record> - ) -> Result<(), ValidationError> { + message: &'a Message, + tsig: &MessageTsig<'a, Octets>, + ) -> Result<(), ValidationError> + where Octets: AsRef<[u8]>, for<'o> &'o Octets: OctetsRef { if message.header().rcode() == Rcode::NotAuth && tsig.data().error() == TsigRcode::BadTime { @@ -1100,11 +1149,15 @@ impl> SigningContext { /// message. fn request( key: K, - message: &[u8], + first: &[u8], + second: Option<&[u8]>, variables: &Variables ) -> (Self, hmac::Tag) { let mut context = key.as_ref().signing_context(); - context.update(message); + context.update(first); + if let Some(second) = second { + context.update(second) + } variables.sign(key.as_ref(), &mut context); let signature = context.sign(); (Self::new(key), signature) @@ -1119,11 +1172,15 @@ impl> SigningContext { /// itself will _not_ change. fn answer( &self, - message: &[u8], + first: &[u8], + second: Option<&[u8]>, variables: &Variables ) -> hmac::Tag { let mut context = self.context.clone(); - context.update(message); + context.update(first); + if let Some(second) = second { + context.update(second) + } variables.sign(self.key.as_ref(), &mut context); context.sign() } @@ -1133,10 +1190,14 @@ impl> SigningContext { /// This is like `answer` above but it doesn’t need to clone the context. fn final_answer( mut self, - message: &[u8], + first: &[u8], + second: Option<&[u8]>, variables: &Variables ) -> (hmac::Tag, K) { - self.context.update(message); + self.context.update(first); + if let Some(second) = second { + self.context.update(second) + } variables.sign(self.key.as_ref(), &mut self.context); (self.context.sign(), self.key) } @@ -1146,7 +1207,8 @@ impl> SigningContext { /// This is like `answer` but it resets the context. fn first_answer( &mut self, - message: &[u8], + first: &[u8], + second: Option<&[u8]>, variables: &Variables ) -> hmac::Tag { // Replace current context with new context. @@ -1154,7 +1216,10 @@ impl> SigningContext { mem::swap(&mut self.context, &mut context); // Update the old context with message and variables, return signature - context.update(message); + context.update(first); + if let Some(second) = second { + context.update(second) + } variables.sign(self.key.as_ref(), &mut context); context.sign() } @@ -1172,7 +1237,8 @@ impl> SigningContext { /// Resets the context. fn signed_subsequent( &mut self, - message: &[u8], + first: &[u8], + second: Option<&[u8]>, variables: &Variables ) -> hmac::Tag { // Replace current context with new context. @@ -1180,44 +1246,78 @@ impl> SigningContext { mem::swap(&mut self.context, &mut context); // Update the old context with message and timers, return signature - context.update(message); + context.update(first); + if let Some(second) = second { + context.update(second) + } variables.sign_timers(&mut context); context.sign() } +} - /// Creates an unsigned error response for the given message. - fn unsigned_error( - msg: &Message, - tsig: &Record>, - error: TsigRcode - ) -> Message { - let mut res = MessageBuilder::new_udp(); - res.start_answer(msg, Rcode::NotAuth); - let mut res = res.additional(); - res.push(( - tsig.owner(), tsig.class(), tsig.ttl(), - Tsig::new( - tsig.data().algorithm(), - tsig.data().time_signed(), - tsig.data().fudge(), - Bytes::new(), - msg.header().id(), - error, - Bytes::new() - ) - )).unwrap(); - res.freeze() + +//------------ MessageTsig --------------------------------------------------- + +/// The TSIG record of a message. +struct MessageTsig<'a, Octets> +where for<'o> &'o Octets: OctetsRef { + /// The actual record. + record: Record< + ParsedDname<&'a Octets>, + Tsig<<&'a Octets as OctetsRef>::Range, ParsedDname<&'a Octets>> + >, + + /// The index of the start of the record. + start: usize, +} + +impl<'a, Octets> MessageTsig<'a, Octets> +where for<'o> &'o Octets: OctetsRef { + /// Get the TSIG record from a message. + /// + /// Checks that there is exactly one TSIG record in the additional + /// section, that it is the last record in this section. If that is true, + /// returns the parsed TSIG records. + fn from_message(msg: &'a Message) -> Option + where Octets: AsRef<[u8]>, for<'o> &'o Octets: OctetsRef { + let mut section = msg.additional().ok()?; + let mut start = section.pos(); + let mut record = section.next()?; + loop { + record = match section.next() { + Some(record) => record, + None => break + }; + start = section.pos(); + } + record.ok()?.into_record::>().ok()?.map(|record| { + MessageTsig { record, start } + }) } - /// Trades the context for a signed error response. - fn signed_error( - self, - message: AdditionalBuilder, - variables: &Variables, - ) -> Message { - let (mac, key) = self.final_answer(message.so_far(), variables); - let mac = key.as_ref().signature_slice(&mac); - key.as_ref().complete_message(message, variables, mac).unwrap() + fn variables(&self) -> Variables { + Variables::new( + self.record.data().time_signed(), + self.record.data().fudge(), + self.record.data().error(), + self.record.data().other_time(), + ) + } + + fn into_original_id(self) -> u16 { + self.record.data().original_id() + } +} + +impl<'a, Octets> std::ops::Deref for MessageTsig<'a, Octets> +where for<'o> &'o Octets: OctetsRef { + type Target = Record< + ParsedDname<&'a Octets>, + Tsig<<&'a Octets as OctetsRef>::Range, ParsedDname<&'a Octets>> + >; + + fn deref(&self) -> &Self::Target { + &self.record } } @@ -1260,28 +1360,20 @@ impl Variables { } } - /// Creates a new value from a given TSIG record. - fn from_tsig(record: &Record>) -> Self { - Variables::new( - record.data().time_signed(), - record.data().fudge(), - record.data().error(), - record.data().other_time(), - ) - } - /// Produces a TSIG record from this value and some more data. - fn to_tsig>( + fn push_tsig( &self, key: &Key, - hmac: S, - original_id: u16 - ) -> Record> { - let other = match self.other { - Some(time) => time.into_bytes(), - None => Bytes::new() + hmac: &[u8], + original_id: u16, + builder: &mut AdditionalBuilder, + ) -> Result<(), ShortBuf> { + let other = self.other.map(Time48::into_octets); + let other = match other { + Some(ref time) => time.as_ref(), + None => b"" }; - Record::new( + builder.push(( key.name.clone(), Class::Any, 0, @@ -1289,12 +1381,12 @@ impl Variables { key.algorithm().to_dname(), self.time_signed, self.fudge, - hmac.into(), + hmac, original_id, self.error, other, ) - ) + )) } /// Applies the variables to a signing context. @@ -1422,10 +1514,10 @@ impl Algorithm { } /// Returns a domain name for this value. - pub fn to_dname(self) -> Dname { + pub fn to_dname(self) -> Dname<&'static [u8]> { unsafe { - Dname::from_bytes_unchecked( - Bytes::from_static(self.into_wire_slice()) + Dname::from_octets_unchecked( + self.into_wire_slice() ) } } @@ -1476,45 +1568,88 @@ impl fmt::Display for Algorithm { //------------ Helper Functions ---------------------------------------------- -/// Extracts the TSIG record from a message. -/// -/// Checks that there is exactly one TSIG record in the additional -/// section, that it is the last record in this section. If that is true, -/// returns both the message without that TSIG record and the TSIG record -/// itself. -/// -/// Note that the function does _not_ update the message ID. -fn extract_tsig( - msg: &mut Message -) -> Option>> { - let additional = match msg.additional() { - Ok(additional) => additional, - Err(_) => return None, - }; - let mut seen = false; - for record in additional.limit_to::>() { - if seen || record.is_err() { - return None - } - seen = true - } - let tsig = match msg.extract_last() { - Some(tsig) => tsig, - None => return None - }; - Some(tsig) +fn remove_tsig(original_id: u16, message: &mut Message) +where + Octets: AsRef<[u8]> + AsMut<[u8]>, + for<'o> &'o Octets: OctetsRef +{ + message.header_mut().set_id(original_id); + message.remove_last_additional(); } -/// Updates message’s ID to tsig’s original ID. -fn update_id( - message: &mut Message, - tsig: &Record> -) { - if message.header().id() != tsig.data().original_id() { - message.update_header(|header| { - header.set_id(tsig.data().original_id()) - }) +//------------ ServerError --------------------------------------------------- + +/// A TSIG record of a received request couldn’t be validated. +/// +/// A value of this type carries all information necessary to produce the +/// error response to be send back to the client. +#[derive(Clone, Debug)] +pub struct ServerError(ServerErrorInner); + +#[derive(Clone, Debug)] +enum ServerErrorInner { + /// Return an unsigned error message. + /// + /// To crate the actual message, we need the original message with the + /// TSIG intact as the last additional record. + Unsigned { + error: TsigRcode, + }, + + /// Return a signed error message. + Signed { + context: SigningContext, + variables: Variables, + } +} + +impl> ServerError { + fn unsigned(error: TsigRcode) -> Self { + ServerError(ServerErrorInner::Unsigned { error }) + } + + fn signed(context: SigningContext, variables: Variables) -> Self { + ServerError(ServerErrorInner::Signed { context, variables }) + } + + pub fn build_message, Target: OctetsBuilder>( + self, + msg: &Message, + builder: MessageBuilder, + ) -> Result, ShortBuf> + where for<'a> &'a Octets: OctetsRef { + let builder = builder.start_answer(msg, Rcode::NotAuth)?; + let mut builder = builder.additional(); + match self.0 { + ServerErrorInner::Unsigned { error } => { + let tsig = { + MessageTsig::from_message( + msg + ).expect("missing or malformed TSIG record") + }; + builder.push(( + tsig.owner(), tsig.class(), tsig.ttl(), + Tsig::new( + tsig.data().algorithm(), + tsig.data().time_signed(), + tsig.data().fudge(), + b"", + msg.header().id(), + error, + b"", + ) + ))?; + } + ServerErrorInner::Signed { context, variables } => { + let (mac, key) = context.final_answer( + builder.as_slice(), None, &variables + ); + let mac = key.as_ref().signature_slice(&mac); + key.as_ref().complete_message(&mut builder, &variables, mac)?; + } + } + Ok(builder) } } @@ -1616,10 +1751,9 @@ pub enum ValidationError { TooManyUnsigned, } -impl From for ValidationError { - fn from(_: ParsedDnameError) -> Self { +impl From for ValidationError { + fn from(_: ParseError) -> Self { ValidationError::FormErr } } - diff --git a/domain-validate/Cargo.toml b/domain-validate/Cargo.toml index b7d7a519..b2dc9584 100644 --- a/domain-validate/Cargo.toml +++ b/domain-validate/Cargo.toml @@ -18,8 +18,8 @@ path = "src/lib.rs" [dependencies] bytes = "0.4" derive_more = "^0.15" -ring = "=0.15.0-alpha3" +ring = "0.16" [dependencies.domain-core] path = "../domain-core" -version = "0.4.1" \ No newline at end of file +version = "0.5.0-pre" diff --git a/domain/Cargo.toml b/domain/Cargo.toml index bac74111..e780ebb0 100644 --- a/domain/Cargo.toml +++ b/domain/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "domain" -version = "0.4.1" +version = "0.5.0-pre" edition = "2018" authors = ["Martin Hoffmann "] description = "A DNS library for Rust – Meta Crate." @@ -19,11 +19,11 @@ validate = ["domain-validate"] [dependencies.domain-core] path = "../domain-core" -version = "0.4.1" +version = "0.5.0-pre" [dependencies.domain-resolv] path = "../domain-resolv" -version = "0.4.1" +version = "0.5.0-pre" optional = true [dependencies.domain-sign] diff --git a/interop/Cargo.toml b/interop/Cargo.toml index 4219ec5f..f1c8561f 100644 --- a/interop/Cargo.toml +++ b/interop/Cargo.toml @@ -11,5 +11,6 @@ domain-core = { path = "../domain-core" } domain-resolv = { path = "../domain-resolv" } domain-tsig = { path = "../domain-tsig" } bytes = "0.4" -ring = "0.15.0-alpha" +ring = "0.16" +unwrap = "1.2" diff --git a/interop/tests/tsig.rs b/interop/tests/tsig.rs index bb4c90dc..ad78d817 100644 --- a/interop/tests/tsig.rs +++ b/interop/tests/tsig.rs @@ -1,6 +1,4 @@ //! Tests the TSIG implementation. -extern crate interop; -extern crate ring; use std::{env, fs, io, thread}; use std::io::{Read, Write}; @@ -9,16 +7,24 @@ use std::process::Command; use std::str::FromStr; use std::time::Duration; use ring::rand::SystemRandom; +use unwrap::unwrap; use interop::nsd; -use interop::domain::core::{Dname, Message, MessageBuilder, Record}; +use interop::domain::core::message::Message; use interop::domain::core::message_builder::{ - AdditionalBuilder, RecordSectionBuilder, SectionBuilder + AdditionalBuilder, AnswerBuilder, MessageBuilder, RecordSectionBuilder, + StreamTarget, }; +use interop::domain::core::name::Dname; use interop::domain::core::iana::{Rcode, Rtype}; use interop::domain::core::rdata::{A, Soa}; use interop::domain::core::utils::base64; use interop::domain::tsig; +type TestMessage = Message>; +type TestBuilder = MessageBuilder>>; +type TestAnswer = AnswerBuilder>>; +type TestAdditional = AdditionalBuilder>>; + //------------ Tests -------------------------------------------------------- @@ -66,21 +72,27 @@ fn tsig_client_nsd() { let res = thread::spawn(move || { // Create an AXFR request and send it to NSD. - let request = MessageBuilder::request_axfr( - Dname::from_str("example.com.").unwrap() - ).additional(); - let (msg, tran) = tsig::ClientTransaction::request(&key, request) - .unwrap(); + let request = TestBuilder::new_stream_vec(); + let mut request = unwrap!(request.request_axfr( + unwrap!(Dname::>::from_str("example.com.")) + )).additional(); + let tran = unwrap!( + tsig::ClientTransaction::request(&key, &mut request) + ); let sock = UdpSocket::bind("127.0.0.1:54320").unwrap(); - sock.send_to(msg.as_ref(), "127.0.0.1:54321").unwrap(); + unwrap!(sock.send_to( + request.as_target().as_dgram_slice(), + "127.0.0.1:54321" + )); let mut answer = loop { let mut buf = vec![0; 512]; let (len, addr) = sock.recv_from(buf.as_mut()).unwrap(); if addr != SocketAddr::from_str("127.0.0.1:54321").unwrap() { continue; } - let answer = Message::from_bytes(buf[..len].into()).unwrap(); - if answer.header().id() == msg.header().id() { + buf.truncate(len); + let answer = Message::from_octets(buf).unwrap(); + if answer.header().id() == request.header().id() { break answer; } }; @@ -91,7 +103,7 @@ fn tsig_client_nsd() { // Shut down NSD just to be sure. let _ = nsd.kill(); - res.unwrap(); // Panic if the thread paniced. + unwrap!(res); // Panic if the thread paniced. } /// Tests the TSIG server implementation against drill as a client. @@ -113,25 +125,34 @@ fn tsig_server_drill() { loop { let mut buf = vec![0; 512]; let (len, addr) = sock.recv_from(buf.as_mut()).unwrap(); - let mut request = match Message::from_bytes(buf[..len].into()) { + buf.truncate(len); + let mut request = match Message::from_octets(buf) { Ok(request) => request, Err(_) => continue, }; - let mut answer = MessageBuilder::new_udp(); - answer.start_answer(&request, Rcode::NoError); - let tran = match tsig::ServerTransaction::request(&&key, - &mut request) { + let answer = TestBuilder::new_stream_vec(); + let answer = unwrap!( + answer.start_answer(&request, Rcode::NoError) + ); + let tran = match tsig::ServerTransaction::request( + &&key, &mut request + ) { Ok(Some(tran)) => tran, Ok(None) => { - sock.send_to(answer.freeze().as_slice(), addr).unwrap(); + sock.send_to(answer.as_slice(), addr).unwrap(); continue; } Err(error) => { - sock.send_to(error.as_slice(), addr).unwrap(); + let answer = unwrap!(error.build_message( + &request, + TestBuilder::new_stream_vec() + )); + sock.send_to(answer.as_slice(), addr).unwrap(); continue; } }; - let answer = tran.answer(answer.additional()).unwrap(); + let mut answer = answer.additional(); + unwrap!(tran.answer(&mut answer)); sock.send_to(answer.as_slice(), addr).unwrap(); } }); @@ -187,14 +208,15 @@ fn tsig_client_sequence_nsd() { } let res = thread::spawn(move || { - let mut sock = TcpStream::connect("127.0.0.1:54323").unwrap(); - let request = MessageBuilder::request_axfr( - Dname::from_str("example.com.").unwrap() - ).additional(); - let (msg, mut tran) = tsig::ClientSequence::request(&key, request) - .unwrap(); - sock.write_all(&(msg.len() as u16).to_be_bytes()).unwrap(); - sock.write_all(msg.as_slice()).unwrap(); + let mut sock = unwrap!(TcpStream::connect("127.0.0.1:54323")); + let request = TestBuilder::new_stream_vec(); + let mut request = unwrap!(request.request_axfr( + unwrap!(Dname::>::from_str("example.com.")) + )).additional(); + let mut tran = unwrap!( + tsig::ClientSequence::request(&key, &mut request) + ); + unwrap!(sock.write_all(request.as_target().as_stream_slice())); loop { let mut len = [0u8; 2]; sock.read_exact(&mut len).unwrap(); @@ -202,8 +224,8 @@ fn tsig_client_sequence_nsd() { assert!(len != 0); let mut buf = vec![0; len]; sock.read_exact(&mut buf).unwrap(); - let mut answer = Message::from_bytes(buf.into()).unwrap(); - tran.answer(&mut answer).unwrap(); + let mut answer = unwrap!(Message::from_octets(buf)); + unwrap!(tran.answer(&mut answer)); // Last message has SOA as last record in answer section. // We don’t care about details. if answer.answer().unwrap().last().unwrap().unwrap().rtype() @@ -211,7 +233,7 @@ fn tsig_client_sequence_nsd() { break } } - tran.done().unwrap() + unwrap!(tran.done()) }).join(); // Shut down NSD just to be sure. @@ -243,25 +265,30 @@ fn tsig_server_sequence_drill() { let len = u16::from_be_bytes(buf) as usize; let mut buf = vec![0; len]; sock.read_exact(&mut buf).unwrap(); - let mut request = Message::from_bytes(buf.into()).unwrap(); + let mut request = Message::from_octets(buf).unwrap(); let mut tran = tsig::ServerSequence::request(&&key, &mut request) .unwrap().unwrap(); + let mut answer = make_first_axfr(&request); + unwrap!(tran.answer(&mut answer)); send_tcp( &mut sock, - tran.answer(make_first_axfr(&request)).unwrap().as_ref() + answer.as_target().as_stream_slice() ).unwrap(); for two in 0..10u8 { for one in 0..10u8 { + let mut answer = make_middle_axfr(&request, one, two); + unwrap!(tran.answer(&mut answer)); send_tcp( &mut sock, - tran.answer(make_middle_axfr(&request, one, two)) - .unwrap().as_ref() + answer.as_target().as_stream_slice() ).unwrap(); } } + let mut answer = make_last_axfr(&request); + unwrap!(tran.answer(&mut answer)); send_tcp( &mut sock, - tran.answer(make_last_axfr(&request)).unwrap().as_ref() + answer.as_target().as_stream_slice() ).unwrap(); } }); @@ -286,50 +313,55 @@ fn send_tcp(sock: &mut TcpStream, msg: &[u8]) -> Result<(), io::Error> { sock.write_all(msg) } -fn make_first_axfr(request: &Message) -> AdditionalBuilder { - let mut msg = MessageBuilder::new_tcp(1024); - msg.start_answer(request, Rcode::NoError); - let mut msg = msg.answer(); - msg.push(make_soa()).unwrap(); - msg.push(make_a(0, 0, 0)).unwrap(); +fn make_first_axfr(request: &TestMessage) -> TestAdditional { + let msg = TestBuilder::new_stream_vec(); + let mut msg = unwrap!(msg.start_answer(request, Rcode::NoError)); + push_soa(&mut msg); + push_a(&mut msg, 0, 0, 0); msg.additional() } -fn make_middle_axfr(request: &Message, one: u8, two: u8) -> AdditionalBuilder { - let mut msg = MessageBuilder::new_tcp(1024); - msg.start_answer(request, Rcode::NoError); - let mut msg = msg.answer(); - msg.push(make_a(1, one, two)).unwrap(); +fn make_middle_axfr( + request: &TestMessage, + one: u8, + two: u8 +) -> TestAdditional { + let msg = TestBuilder::new_stream_vec(); + let mut msg = unwrap!(msg.start_answer(request, Rcode::NoError)); + push_a(&mut msg, 1, one, two); msg.additional() } -fn make_last_axfr(request: &Message) -> AdditionalBuilder { - let mut msg = MessageBuilder::new_tcp(1024); - msg.start_answer(request, Rcode::NoError); - let mut msg = msg.answer(); - msg.push(make_a(2, 0, 0)).unwrap(); - msg.push(make_soa()).unwrap(); +fn make_last_axfr(request: &TestMessage) -> TestAdditional { + let msg = TestBuilder::new_stream_vec(); + let mut msg = unwrap!(msg.start_answer(request, Rcode::NoError)); + push_a(&mut msg, 2, 0, 0); + push_soa(&mut msg); msg.additional() } -fn make_soa() -> Record> { - ( - Dname::from_str("example.com.").unwrap(), - 3600, - Soa::new( - Dname::from_str("mname.example.com.").unwrap(), - Dname::from_str("rname.example.com.").unwrap(), - 12.into(), - 3600, 3600, 3600, 3600 +fn push_soa(builder: &mut TestAnswer) { + unwrap!(builder.push( + ( + Dname::>::from_str("example.com.").unwrap(), + 3600, + Soa::new( + Dname::>::from_str("mname.example.com.").unwrap(), + Dname::>::from_str("rname.example.com.").unwrap(), + 12.into(), + 3600, 3600, 3600, 3600 + ) ) - ).into() + )) } -fn make_a(zero: u8, one: u8, two: u8) -> Record { - ( - Dname::from_str("example.com.").unwrap(), - 3600, - A::from_octets(10, zero, one, two) - ).into() +fn push_a(builder: &mut TestAnswer, zero: u8, one: u8, two: u8) { + unwrap!(builder.push( + ( + Dname::>::from_str("example.com.").unwrap(), + 3600, + A::from_octets(10, zero, one, two) + ) + )) }