From 90cc96d63240b4e6ee2e5dc1dfd1dd5a302f23ce Mon Sep 17 00:00:00 2001 From: Philip Homburg Date: Wed, 4 Sep 2024 14:00:25 +0200 Subject: [PATCH] Add support to client transports for requests that may result in multiple responses. (#377) This adds ComposeRequestMulti and other *Multi types. The main change is to the stream transport, which is the only transport that implements SendRequestMulti. --- examples/client-transports.rs | 9 +- src/base/message.rs | 8 + src/net/client/mod.rs | 55 ++- src/net/client/multi_stream.rs | 9 +- src/net/client/request.rs | 363 ++++++++++++++- src/net/client/stream.rs | 669 +++++++++++++++++++++++++--- src/net/client/validator.rs | 2 +- src/net/server/tests/integration.rs | 6 +- src/resolv/stub/mod.rs | 4 +- src/stelline/client.rs | 3 +- src/validator/context.rs | 2 +- src/validator/mod.rs | 2 +- tests/net-client-cache.rs | 2 +- tests/net-client.rs | 9 +- 14 files changed, 1044 insertions(+), 99 deletions(-) diff --git a/examples/client-transports.rs b/examples/client-transports.rs index 92705bc0..cea3bce1 100644 --- a/examples/client-transports.rs +++ b/examples/client-transports.rs @@ -8,7 +8,9 @@ use domain::net::client::dgram_stream; use domain::net::client::multi_stream; use domain::net::client::protocol::{TcpConnect, TlsConnect, UdpConnect}; use domain::net::client::redundant; -use domain::net::client::request::{RequestMessage, SendRequest}; +use domain::net::client::request::{ + RequestMessage, RequestMessageMulti, SendRequest, +}; use domain::net::client::stream; use std::net::{IpAddr, SocketAddr}; use std::str::FromStr; @@ -33,7 +35,7 @@ async fn main() { let mut msg = msg.question(); msg.push((Name::vec_from_str("example.com").unwrap(), Rtype::AAAA)) .unwrap(); - let req = RequestMessage::new(msg); + let req = RequestMessage::new(msg).unwrap(); // Destination for UDP and TCP let server_addr = SocketAddr::new(IpAddr::from_str("::1").unwrap(), 53); @@ -234,7 +236,8 @@ async fn main() { } }; - let (tcp, transport) = stream::Connection::new(tcp_conn); + let (tcp, transport) = + stream::Connection::<_, RequestMessageMulti>>::new(tcp_conn); tokio::spawn(async move { transport.run().await; println!("single TCP run terminated"); diff --git a/src/base/message.rs b/src/base/message.rs index b1272e68..7f52ed42 100644 --- a/src/base/message.rs +++ b/src/base/message.rs @@ -439,6 +439,14 @@ impl Message { } } + /// Returns whether the message has a question that is either AXFR or + /// IXFR. + pub fn is_xfr(&self) -> bool { + self.first_question() + .map(|q| matches!(q.qtype(), Rtype::AXFR | Rtype::IXFR)) + .unwrap_or_default() + } + /// Returns the first question, if there is any. /// /// The method will return `None` both if there are no questions or if diff --git a/src/net/client/mod.rs b/src/net/client/mod.rs index c41b8a6b..36567610 100644 --- a/src/net/client/mod.rs +++ b/src/net/client/mod.rs @@ -32,7 +32,7 @@ //! 1) Creating a request message, //! 2) Creating a DNS transport, //! 3) Sending the request, and -//! 4) Receiving the reply. +//! 4) Receiving the reply or replies. //! //! The first and second step are independent and can happen in any order. //! The third step uses the resuts of the first and second step. @@ -87,7 +87,7 @@ //! tokio::spawn(transport.run()); //! # let req = domain::net::client::request::RequestMessage::new( //! # domain::base::MessageBuilder::new_vec() -//! # ); +//! # ).unwrap(); //! # let mut request = tcp_conn.send_request(req); //! # } //! ``` @@ -100,18 +100,18 @@ //! //! For example: //! ```no_run -//! # use domain::net::client::request::SendRequest; +//! # use domain::net::client::request::{RequestMessageMulti, SendRequest}; //! # use std::net::{IpAddr, SocketAddr}; //! # use std::str::FromStr; //! # async fn _test() { -//! # let (tls_conn, _) = domain::net::client::stream::Connection::new( +//! # let (tls_conn, _) = domain::net::client::stream::Connection::<_, RequestMessageMulti>>::new( //! # domain::net::client::protocol::TcpConnect::new( //! # SocketAddr::new(IpAddr::from_str("::1").unwrap(), 53) //! # ) //! # ); //! # let req = domain::net::client::request::RequestMessage::new( //! # domain::base::MessageBuilder::new_vec() -//! # ); +//! # ).unwrap(); //! let mut request = tls_conn.send_request(req); //! # } //! ``` @@ -128,22 +128,61 @@ //! //! For example: //! ```no_run -//! # use crate::domain::net::client::request::SendRequest; +//! # use crate::domain::net::client::request::{RequestMessageMulti, SendRequest}; //! # use std::net::{IpAddr, SocketAddr}; //! # use std::str::FromStr; //! # async fn _test() { -//! # let (tls_conn, _) = domain::net::client::stream::Connection::new( +//! # let (tls_conn, _) = domain::net::client::stream::Connection::<_, RequestMessageMulti>>::new( //! # domain::net::client::protocol::TcpConnect::new( //! # SocketAddr::new(IpAddr::from_str("::1").unwrap(), 53) //! # ) //! # ); //! # let req = domain::net::client::request::RequestMessage::new( //! # domain::base::MessageBuilder::new_vec() -//! # ); +//! # ).unwrap(); //! # let mut request = tls_conn.send_request(req); //! let reply = request.get_response().await; //! # } //! ``` +//! +//!
+//! +//! **Support for multiple responses:** +//! +//! [RequestMessage][request::RequestMessage] is designed for the most common +//! use case: single request, single response. +//! +//! However, zone transfers (e.g. using the `AXFR` or `IXFR` query types) can +//! result in multiple responses. Attempting to create a +//! [RequestMessage][request::RequestMessage] for such a query will result in +//! [Error::FormError][request::Error::FormError]. +//! +//! For zone transfers you should use +//! [RequestMessageMulti][request::RequestMessageMulti] instead which can be +//! used like so: +//! +//! ```no_run +//! # use crate::domain::net::client::request::{RequestMessage, SendRequestMulti}; +//! # use std::net::{IpAddr, SocketAddr}; +//! # use std::str::FromStr; +//! # async fn _test() { +//! # let (conn, _) = domain::net::client::stream::Connection::>, _>::new( +//! # domain::net::client::protocol::TcpConnect::new( +//! # SocketAddr::new(IpAddr::from_str("::1").unwrap(), 53) +//! # ) +//! # ); +//! # let req = domain::net::client::request::RequestMessageMulti::new( +//! # domain::base::MessageBuilder::new_vec() +//! # ).unwrap(); +//! # let mut request = conn.send_request(req); +//! while let Ok(reply) = request.get_response().await { +//! // ... +//! } +//! # } +//! ``` +//! +//!
+//! //! # Limitations //! diff --git a/src/net/client/multi_stream.rs b/src/net/client/multi_stream.rs index 6016178a..d0c65c75 100644 --- a/src/net/client/multi_stream.rs +++ b/src/net/client/multi_stream.rs @@ -6,7 +6,7 @@ use crate::base::Message; use crate::net::client::protocol::AsyncConnect; use crate::net::client::request::{ - ComposeRequest, Error, GetResponse, SendRequest, + ComposeRequest, Error, GetResponse, RequestMessageMulti, SendRequest, }; use crate::net::client::stream; use bytes::Bytes; @@ -19,6 +19,7 @@ use std::future::Future; use std::pin::Pin; use std::sync::Arc; use std::time::Duration; +use std::vec::Vec; use tokio::io; use tokio::io::{AsyncRead, AsyncWrite}; use tokio::sync::{mpsc, oneshot}; @@ -197,7 +198,7 @@ enum QueryState { ReceiveConn(oneshot::Receiver>), /// Start a query using the given stream transport. - StartQuery(Arc>), + StartQuery(Arc>>>), /// Get the result of the query. GetResult(stream::Request), @@ -222,7 +223,7 @@ struct ChanRespOk { id: u64, /// The new stream transport to use for sending a request. - conn: Arc>, + conn: Arc>>>, } impl Request { @@ -409,7 +410,7 @@ enum SingleConnState3 { None, /// Current stream transport. - Some(Arc>), + Some(Arc>>>), /// State that deals with an error getting a new octet stream from /// a connection stream. diff --git a/src/net/client/request.rs b/src/net/client/request.rs index dd4f5ddc..dd45a058 100644 --- a/src/net/client/request.rs +++ b/src/net/client/request.rs @@ -1,13 +1,12 @@ //! Constructing and sending requests. - -use crate::base::iana::Rcode; +use crate::base::iana::{Opcode, Rcode}; use crate::base::message::{CopyRecordsError, ShortMessage}; use crate::base::message_builder::{ - AdditionalBuilder, MessageBuilder, PushError, StaticCompressor, + AdditionalBuilder, MessageBuilder, PushError, }; use crate::base::opt::{ComposeOptData, LongOptData, OptRecord}; use crate::base::wire::{Composer, ParseError}; -use crate::base::{Header, Message, ParsedName, Rtype}; +use crate::base::{Header, Message, ParsedName, Rtype, StaticCompressor}; use crate::rdata::AllRecordData; use bytes::Bytes; use octseq::Octets; @@ -20,6 +19,9 @@ use std::vec::Vec; use std::{error, fmt}; use tracing::trace; +#[cfg(feature = "tsig")] +use crate::tsig; + //------------ ComposeRequest ------------------------------------------------ /// A trait that allows composing a request as a series. @@ -27,8 +29,8 @@ pub trait ComposeRequest: Debug + Send + Sync { /// Appends the final message to a provided composer. fn append_message( &self, - target: &mut Target, - ) -> Result<(), CopyRecordsError>; + target: Target, + ) -> Result, CopyRecordsError>; /// Create a message that captures the recorded changes. fn to_message(&self) -> Result>, Error>; @@ -62,6 +64,47 @@ pub trait ComposeRequest: Debug + Send + Sync { fn dnssec_ok(&self) -> bool; } +//------------ ComposeRequestMulti -------------------------------------------- + +/// A trait that allows composing a request as a series. +pub trait ComposeRequestMulti: Debug + Send + Sync { + /// Appends the final message to a provided composer. + fn append_message( + &self, + target: Target, + ) -> Result, CopyRecordsError>; + + /// Create a message that captures the recorded changes. + fn to_message(&self) -> Result>, Error>; + + /// Create a message that captures the recorded changes and convert to + /// a Vec. + + /// Return a reference to the current Header. + fn header(&self) -> &Header; + + /// Return a reference to a mutable Header to record changes to the header. + fn header_mut(&mut self) -> &mut Header; + + /// Set the UDP payload size. + fn set_udp_payload_size(&mut self, value: u16); + + /// Set the DNSSEC OK flag. + fn set_dnssec_ok(&mut self, value: bool); + + /// Add an EDNS option. + fn add_opt( + &mut self, + opt: &impl ComposeOptData, + ) -> Result<(), LongOptData>; + + /// Returns whether a message is an answer to the request. + fn is_answer(&self, answer: &Message<[u8]>) -> bool; + + /// Return the status of the DNSSEC OK flag. + fn dnssec_ok(&self) -> bool; +} + //------------ SendRequest --------------------------------------------------- /// Trait for starting a DNS request based on a request composer. @@ -76,6 +119,20 @@ pub trait SendRequest { ) -> Box; } +//------------ SendRequestMulti ----------------------------------------------- + +/// Trait for starting a DNS request based on a request composer. +/// +/// In the future, the return type of request should become an associated type. +/// However, the use of 'dyn Request' in redundant currently prevents that. +pub trait SendRequestMulti { + /// Request function that takes a ComposeRequestMulti type. + fn send_request( + &self, + request_msg: CR, + ) -> Box; +} + //------------ GetResponse --------------------------------------------------- /// Trait for getting the result of a DNS query. @@ -98,6 +155,28 @@ pub trait GetResponse: Debug { >; } +//------------ GetResponseMulti ---------------------------------------------- +/// Trait for getting a stream of result of a DNS query. +/// +/// In the future, the return type of get_response should become an associated +/// type. However, too many uses of 'dyn GetResponse' currently prevent that. +#[allow(clippy::type_complexity)] +pub trait GetResponseMulti: Debug { + /// Get the result of a DNS request. + /// + /// This function is intended to be cancel safe. + fn get_response( + &mut self, + ) -> Pin< + Box< + dyn Future>, Error>> + + Send + + Sync + + '_, + >, + >; +} + //------------ RequestMessage ------------------------------------------------ /// Object that implements the ComposeRequest trait for a Message object. @@ -114,15 +193,26 @@ pub struct RequestMessage> { } impl + Debug + Octets> RequestMessage { - /// Create a new BMB object. - pub fn new(msg: impl Into>) -> Self { + /// Create a new RequestMessage object. + pub fn new(msg: impl Into>) -> Result { let msg = msg.into(); + + // On UDP, IXFR results in a single response, so we need to accept it. + // We can reject AXFR because it always requires support for multiple + // responses. + if msg.header().opcode() == Opcode::QUERY + && msg.first_question().ok_or(Error::FormError)?.qtype() + == Rtype::AXFR + { + return Err(Error::FormError); + } + let header = msg.header(); - Self { + Ok(Self { msg, header, opt: None, - } + }) } /// Returns a mutable reference to the OPT record. @@ -209,12 +299,12 @@ impl + Debug + Octets + Send + Sync> ComposeRequest { fn append_message( &self, - target: &mut Target, - ) -> Result<(), CopyRecordsError> { + target: Target, + ) -> Result, CopyRecordsError> { let target = MessageBuilder::from_target(target) .map_err(|_| CopyRecordsError::Push(PushError::ShortBuf))?; - self.append_message_impl(target)?; - Ok(()) + let builder = self.append_message_impl(target)?; + Ok(builder) } fn to_vec(&self) -> Result, Error> { @@ -298,6 +388,222 @@ impl + Debug + Octets + Send + Sync> ComposeRequest } } +//------------ RequestMessageMulti -------------------------------------------- + +/// Object that implements the ComposeRequestMulti trait for a Message object. +#[derive(Clone, Debug)] +pub struct RequestMessageMulti +where + Octs: AsRef<[u8]>, +{ + /// Base message. + msg: Message, + + /// New header. + header: Header, + + /// The OPT record to add if required. + opt: Option>>, +} + +impl + Debug + Octets> RequestMessageMulti { + /// Create a new BMB object. + pub fn new(msg: impl Into>) -> Result { + let msg = msg.into(); + + // Only accept the streaming types (IXFR and AXFR). + if !msg.is_xfr() { + return Err(Error::FormError); + } + let header = msg.header(); + Ok(Self { + msg, + header, + opt: None, + }) + } + + /// Returns a mutable reference to the OPT record. + /// + /// Adds one if necessary. + fn opt_mut(&mut self) -> &mut OptRecord> { + self.opt.get_or_insert_with(Default::default) + } + + /// Appends the message to a composer. + fn append_message_impl( + &self, + mut target: MessageBuilder, + ) -> Result, CopyRecordsError> { + let source = &self.msg; + + *target.header_mut() = self.header; + + let source = source.question(); + let mut target = target.question(); + for rr in source { + target.push(rr?)?; + } + let mut source = source.answer()?; + let mut target = target.answer(); + for rr in &mut source { + let rr = rr? + .into_record::>>()? + .expect("record expected"); + target.push(rr)?; + } + + let mut source = + source.next_section()?.expect("section should be present"); + let mut target = target.authority(); + for rr in &mut source { + let rr = rr? + .into_record::>>()? + .expect("record expected"); + target.push(rr)?; + } + + let source = + source.next_section()?.expect("section should be present"); + let mut target = target.additional(); + for rr in source { + let rr = rr?; + if rr.rtype() != Rtype::OPT { + let rr = rr + .into_record::>>()? + .expect("record expected"); + target.push(rr)?; + } + } + + if let Some(opt) = self.opt.as_ref() { + target.push(opt.as_record())?; + } + + Ok(target) + } + + /// Create new message based on the changes to the base message. + fn to_message_impl(&self) -> Result>, Error> { + let target = + MessageBuilder::from_target(StaticCompressor::new(Vec::new())) + .expect("Vec is expected to have enough space"); + + let target = self.append_message_impl(target)?; + + // It would be nice to use .builder() here. But that one deletes all + // sections. We have to resort to .as_builder() which gives a + // reference and then .clone() + let result = target.as_builder().clone(); + let msg = Message::from_octets(result.finish().into_target()).expect( + "Message should be able to parse output from MessageBuilder", + ); + Ok(msg) + } +} + +impl + Debug + Octets + Send + Sync> ComposeRequestMulti + for RequestMessageMulti +{ + fn append_message( + &self, + target: Target, + ) -> Result, CopyRecordsError> { + let target = MessageBuilder::from_target(target) + .map_err(|_| CopyRecordsError::Push(PushError::ShortBuf))?; + let builder = self.append_message_impl(target)?; + Ok(builder) + } + + fn to_message(&self) -> Result>, Error> { + self.to_message_impl() + } + + fn header(&self) -> &Header { + &self.header + } + + fn header_mut(&mut self) -> &mut Header { + &mut self.header + } + + fn set_udp_payload_size(&mut self, value: u16) { + self.opt_mut().set_udp_payload_size(value); + } + + fn set_dnssec_ok(&mut self, value: bool) { + self.opt_mut().set_dnssec_ok(value); + } + + fn add_opt( + &mut self, + opt: &impl ComposeOptData, + ) -> Result<(), LongOptData> { + self.opt_mut().push(opt).map_err(|e| e.unlimited_buf()) + } + + fn is_answer(&self, answer: &Message<[u8]>) -> bool { + let answer_header = answer.header(); + let answer_hcounts = answer.header_counts(); + + // First check qr is set and IDs match. + if !answer_header.qr() || answer_header.id() != self.header.id() { + trace!( + "Wrong QR or ID: QR={}, answer ID={}, self ID={}", + answer_header.qr(), + answer_header.id(), + self.header.id() + ); + return false; + } + + // If the result is an error, then the question section can be empty. + // In that case we require all other sections to be empty as well. + if answer_header.rcode() != Rcode::NOERROR + && answer_hcounts.qdcount() == 0 + && answer_hcounts.ancount() == 0 + && answer_hcounts.nscount() == 0 + && answer_hcounts.arcount() == 0 + { + // We can accept this as a valid reply. + return true; + } + + // Now the question section in the reply has to be the same as in the + // query, except in the case of an AXFR subsequent response: + // + // https://datatracker.ietf.org/doc/html/rfc5936#section-2.2 + // 2.2. AXFR Response + // "The AXFR server MUST copy the Question section from the + // corresponding AXFR query message into the first response + // message's Question section. For subsequent messages, it MAY do + // the same or leave the Question section empty." + if self.msg.qtype() == Some(Rtype::AXFR) + && answer_hcounts.qdcount() == 0 + { + true + } else if answer_hcounts.qdcount() + != self.msg.header_counts().qdcount() + { + trace!("Wrong QD count"); + false + } else { + let res = answer.question() == self.msg.for_slice().question(); + if !res { + trace!("Wrong question"); + } + res + } + } + + fn dnssec_ok(&self) -> bool { + match &self.opt { + None => false, + Some(opt) => opt.dnssec_ok(), + } + } +} + //------------ Error --------------------------------------------------------- /// Error type for client transports. @@ -318,6 +624,9 @@ pub enum Error { /// Underlying transport not found in redundant connection RedundantTransportNotFound, + /// The message violated some constraints. + FormError, + /// Octet sequence too short to be a valid DNS message. ShortMessage, @@ -355,6 +664,14 @@ pub enum Error { /// An error happened in the datagram transport. Dgram(Arc), + #[cfg(feature = "unstable-server-transport")] + /// Zone write failed. + ZoneWrite, + + #[cfg(feature = "tsig")] + /// TSIG authentication failed. + Authentication(tsig::ValidationError), + #[cfg(feature = "unstable-validator")] /// An error happened during DNSSEC validation. Validation(crate::validator::context::Error), @@ -407,6 +724,9 @@ impl fmt::Display for Error { Error::ShortMessage => { write!(f, "octet sequence to short to be a valid message") } + Error::FormError => { + write!(f, "message violates a constraint") + } Error::StreamLongMessage => { write!(f, "message too long for stream transport") } @@ -436,6 +756,13 @@ impl fmt::Display for Error { write!(f, "no transport available") } Error::Dgram(err) => fmt::Display::fmt(err, f), + + #[cfg(feature = "unstable-server-transport")] + Error::ZoneWrite => write!(f, "zone write error"), + + #[cfg(feature = "tsig")] + Error::Authentication(err) => fmt::Display::fmt(err, f), + #[cfg(feature = "unstable-validator")] Error::Validation(_) => { write!(f, "error validating response") @@ -462,6 +789,7 @@ impl error::Error for Error { Error::MessageParseError => None, Error::RedundantTransportNotFound => None, Error::ShortMessage => None, + Error::FormError => None, Error::StreamLongMessage => None, Error::StreamIdleTimeout => None, Error::StreamReceiveError => None, @@ -473,6 +801,13 @@ impl error::Error for Error { Error::WrongReplyForQuery => None, Error::NoTransportAvailable => None, Error::Dgram(err) => Some(err), + + #[cfg(feature = "unstable-server-transport")] + Error::ZoneWrite => None, + + #[cfg(feature = "tsig")] + Error::Authentication(err) => Some(err), + #[cfg(feature = "unstable-validator")] Error::Validation(err) => Some(err), } diff --git a/src/net/client/stream.rs b/src/net/client/stream.rs index 2c82fecc..18ed58d3 100644 --- a/src/net/client/stream.rs +++ b/src/net/client/stream.rs @@ -3,23 +3,16 @@ // RFC 7766 describes DNS over TCP // RFC 7828 describes the edns-tcp-keepalive option -// TODO: -// - errors -// - connect errors? Retry after connection refused? -// - server errors -// - ID out of range -// - ID not in use -// - reply for wrong query -// - timeouts -// - request timeout -// - create new connection after end/failure of previous one - +use super::request::{ + ComposeRequest, ComposeRequestMulti, Error, GetResponse, + GetResponseMulti, SendRequest, SendRequestMulti, +}; +use crate::base::iana::{Rcode, Rtype}; use crate::base::message::Message; use crate::base::message_builder::StreamTarget; use crate::base::opt::{AllOptData, OptRecord, TcpKeepalive}; -use crate::net::client::request::{ - ComposeRequest, Error, GetResponse, SendRequest, -}; +use crate::base::{ParsedName, Serial}; +use crate::rdata::AllRecordData; use crate::utils::config::DefMinMax; use bytes::{Bytes, BytesMut}; use core::cmp; @@ -34,6 +27,7 @@ use std::vec::Vec; use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt}; use tokio::sync::{mpsc, oneshot}; use tokio::time::sleep; +use tracing::trace; //------------ Configuration Constants ---------------------------------------- @@ -77,9 +71,15 @@ const READ_REPLY_CHAN_CAP: usize = 8; /// Configuration for a stream transport connection. #[derive(Clone, Debug)] pub struct Config { - /// Response timeout. + /// Response timeout currently in effect. response_timeout: Duration, + /// Single response timeout. + single_response_timeout: Duration, + + /// Streaming response timeout. + streaming_response_timeout: Duration, + /// Default idle timeout. /// /// This value is used if the other side does not send a TcpKeepalive @@ -102,11 +102,53 @@ impl Config { } /// Sets the response timeout. + /// + /// For requests where ComposeRequest::is_streaming() returns true see + /// set_streaming_response_timeout() instead. + /// + /// Excessive values are quietly trimmed. + // + // XXX Maybe that’s wrong and we should rather return an error? pub fn set_response_timeout(&mut self, timeout: Duration) { self.response_timeout = RESPONSE_TIMEOUT.limit(timeout); + self.streaming_response_timeout = self.response_timeout; } - /// Sets the idle timeout. + /// Returns the streaming response timeout. + pub fn streaming_response_timeout(&self) -> Duration { + self.streaming_response_timeout + } + + /// Sets the streaming response timeout. + /// + /// Only used for requests where ComposeRequest::is_streaming() returns + /// true as it is typically desirable that such response streams be + /// allowed to complete even if the individual responses arrive very + /// slowly. + /// + /// Excessive values are quietly trimmed. + pub fn set_streaming_response_timeout(&mut self, timeout: Duration) { + self.streaming_response_timeout = RESPONSE_TIMEOUT.limit(timeout); + } + + /// Returns the initial idle timeout, if set. + pub fn idle_timeout(&self) -> Duration { + self.idle_timeout + } + + /// Sets the initial idle timeout. + /// + /// By default the stream is immediately closed if there are no pending + /// requests or responses. + /// + /// Set this to allow requests to be sent in sequence with delays between + /// such as a SOA query followed by AXFR for more efficient use of the + /// stream per RFC 9103. + /// + /// Note: May be overridden by an RFC 7828 edns-tcp-keepalive timeout + /// received from a server. + /// + /// Excessive values are quietly trimmed. pub fn set_idle_timeout(&mut self, timeout: Duration) { self.idle_timeout = IDLE_TIMEOUT.limit(timeout) } @@ -116,6 +158,8 @@ impl Default for Config { fn default() -> Self { Self { response_timeout: RESPONSE_TIMEOUT.default(), + single_response_timeout: RESPONSE_TIMEOUT.default(), + streaming_response_timeout: RESPONSE_TIMEOUT.default(), idle_timeout: IDLE_TIMEOUT.default(), } } @@ -125,19 +169,21 @@ impl Default for Config { /// A connection to a single stream transport. #[derive(Debug)] -pub struct Connection { +pub struct Connection { /// The sender half of the request channel. - sender: mpsc::Sender>, + sender: mpsc::Sender>, } -impl Connection { +impl Connection { /// Creates a new stream transport with default configuration. /// /// Returns a connection and a future that drives the transport using /// the provided stream. This future needs to be run while any queries /// are active. This is most easly achieved by spawning it into a runtime. /// It terminates when the last connection is dropped. - pub fn new(stream: Stream) -> (Self, Transport) { + pub fn new( + stream: Stream, + ) -> (Self, Transport) { Self::with_config(stream, Default::default()) } @@ -150,13 +196,17 @@ impl Connection { pub fn with_config( stream: Stream, config: Config, - ) -> (Self, Transport) { + ) -> (Self, Transport) { let (sender, transport) = Transport::new(stream, config); (Self { sender }, transport) } } -impl Connection { +impl Connection +where + Req: ComposeRequest + 'static, + ReqMulti: ComposeRequestMulti + 'static, +{ /// Start a DNS request. /// /// This function takes a precomposed message as a parameter and @@ -166,6 +216,8 @@ impl Connection { msg: Req, ) -> Result, Error> { let (sender, receiver) = oneshot::channel(); + let sender = ReplySender::Single(Some(sender)); + let msg = ReqSingleMulti::Single(msg); let req = ChanReq { sender, msg }; self.sender.send(req).await.map_err(|_| { // Send error. The receiver is gone, this means that the @@ -175,15 +227,47 @@ impl Connection { receiver.await.map_err(|_| Error::StreamReceiveError)? } - /// Returns a request handler for this connection. + /// Start a streaming request. + async fn handle_streaming_request_impl( + self, + msg: ReqMulti, + sender: mpsc::Sender>, Error>>, + ) -> Result<(), Error> { + let reply_sender = ReplySender::Stream(sender); + let msg = ReqSingleMulti::Multi(msg); + let req = ChanReq { + sender: reply_sender, + msg, + }; + self.sender.send(req).await.map_err(|_| { + // Send error. The receiver is gone, this means that the + // connection is closed. + Error::ConnectionClosed + })?; + Ok(()) + } + + /// Returns a request handler for a request. pub fn get_request(&self, request_msg: Req) -> Request { Request { fut: Box::pin(self.clone().handle_request_impl(request_msg)), } } + + /// Return a multiple-response request handler for a request. + fn get_streaming_request(&self, request_msg: ReqMulti) -> RequestMulti { + let (sender, receiver) = mpsc::channel(DEF_CHAN_CAP); + RequestMulti { + stream: receiver, + fut: Some(Box::pin( + self.clone() + .handle_streaming_request_impl(request_msg, sender), + )), + } + } } -impl Clone for Connection { +impl Clone for Connection { fn clone(&self) -> Self { Self { sender: self.sender.clone(), @@ -191,8 +275,10 @@ impl Clone for Connection { } } -impl SendRequest - for Connection +impl SendRequest for Connection +where + Req: ComposeRequest + 'static, + ReqMulti: ComposeRequestMulti + Debug + Send + Sync + 'static, { fn send_request( &self, @@ -202,6 +288,19 @@ impl SendRequest } } +impl SendRequestMulti for Connection +where + Req: ComposeRequest + Debug + Send + Sync + 'static, + ReqMulti: ComposeRequestMulti + 'static, +{ + fn send_request( + &self, + request_msg: ReqMulti, + ) -> Box { + Box::new(self.get_streaming_request(request_msg)) + } +} + //------------ Request ------------------------------------------------------- /// An active request. @@ -242,34 +341,140 @@ impl Debug for Request { } } +//------------ RequestMulti -------------------------------------------------- + +/// An active request. +pub struct RequestMulti { + /// Receiver for a stream of responses. + stream: mpsc::Receiver>, Error>>, + + /// The underlying future. + #[allow(clippy::type_complexity)] + fut: Option< + Pin> + Send + Sync>>, + >, +} + +impl RequestMulti { + /// Async function that waits for the future stored in Request to complete. + async fn get_response_impl( + &mut self, + ) -> Result>, Error> { + if self.fut.is_some() { + let fut = self.fut.take().expect("Some expected"); + fut.await?; + } + + // Fetch from the stream + self.stream + .recv() + .await + .ok_or(Error::ConnectionClosed) + .map_err(|_| Error::ConnectionClosed)? + } +} + +impl GetResponseMulti for RequestMulti { + fn get_response( + &mut self, + ) -> Pin< + Box< + dyn Future>, Error>> + + Send + + Sync + + '_, + >, + > { + let fut = self.get_response_impl(); + Box::pin(fut) + } +} + +impl Debug for RequestMulti { + fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result { + f.debug_struct("Request") + .field("fut", &format_args!("_")) + .finish() + } +} + //------------ Transport ----------------------------------------------------- /// The underlying machinery of a stream transport. #[derive(Debug)] -pub struct Transport { - /// The stream socket towards the remove end. +pub struct Transport { + /// The stream socket towards the remote end. stream: Stream, /// Transport configuration. config: Config, /// The receiver half of request channel. - receiver: mpsc::Receiver>, + receiver: mpsc::Receiver>, +} + +/// This is the type of sender in [ChanReq]. +#[derive(Debug)] +enum ReplySender { + /// Return channel for a single response. + Single(Option>), + + /// Return channel for a stream of responses. + Stream(mpsc::Sender>, Error>>), +} + +impl ReplySender { + /// Send a response. + async fn send(&mut self, resp: ChanResp) -> Result<(), ()> { + match self { + ReplySender::Single(sender) => match sender.take() { + Some(sender) => sender.send(resp).map_err(|_| ()), + None => Err(()), + }, + ReplySender::Stream(sender) => { + sender.send(resp.map(Some)).await.map_err(|_| ()) + } + } + } + + /// Send EOF on a response stream. + async fn send_eof(&mut self) -> Result<(), ()> { + match self { + ReplySender::Single(_) => { + panic!("cannot send EOF for Single"); + } + ReplySender::Stream(sender) => { + sender.send(Ok(None)).await.map_err(|_| ()) + } + } + } + + /// Report whether in stream mode or not. + fn is_stream(&self) -> bool { + matches!(self, Self::Stream(_)) + } +} + +#[derive(Debug)] +/// Enum that can either store a request for a single response or one for +/// multiple responses. +enum ReqSingleMulti { + /// Single response request. + Single(Req), + /// Multi-response request. + Multi(ReqMulti), } /// A message from a [`Request`] to start a new request. #[derive(Debug)] -struct ChanReq { +struct ChanReq { /// DNS request message - msg: Req, + msg: ReqSingleMulti, /// Sender to send result back to [`Request`] sender: ReplySender, } -/// This is the type of sender in [ChanReq]. -type ReplySender = oneshot::Sender; - /// A message back to [`Request`] returning a response. type ChanResp = Result, Error>; @@ -323,12 +528,60 @@ enum ConnState { WriteError(Error), } -impl Transport { +//--- Display +impl std::fmt::Display for ConnState { + fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result { + match self { + ConnState::Active(instant) => f.write_fmt(format_args!( + "Active (since {}s ago)", + instant + .map(|v| Instant::now().duration_since(v).as_secs()) + .unwrap_or_default() + )), + ConnState::Idle(instant) => f.write_fmt(format_args!( + "Idle (since {}s ago)", + Instant::now().duration_since(*instant).as_secs() + )), + ConnState::IdleTimeout => f.write_str("IdleTimeout"), + ConnState::ReadError(err) => { + f.write_fmt(format_args!("ReadError: {err}")) + } + ConnState::ReadTimeout => f.write_str("ReadTimeout"), + ConnState::WriteError(err) => { + f.write_fmt(format_args!("WriteError: {err}")) + } + } + } +} + +#[derive(Debug)] +/// State of an AXFR or IXFR responses stream for detecting the end of the +/// stream. +enum XFRState { + /// Start of AXFR. + AXFRInit, + /// After the first SOA record has been encountered. + AXFRFirstSoa(Serial), + /// Start of IXFR. + IXFRInit, + /// After the first SOA record has been encountered. + IXFRFirstSoa(Serial), + /// After the first SOA record in a diff section has been encountered. + IXFRFirstDiffSoa(Serial), + /// After the second SOA record in a diff section has been encountered. + IXFRSecondDiffSoa(Serial), + /// End of the stream has been found. + Done, + /// An error has occured. + Error, +} + +impl Transport { /// Creates a new transport. fn new( stream: Stream, config: Config, - ) -> (mpsc::Sender>, Self) { + ) -> (mpsc::Sender>, Self) { let (sender, receiver) = mpsc::channel(DEF_CHAN_CAP); ( sender, @@ -341,10 +594,11 @@ impl Transport { } } -impl Transport +impl Transport where Stream: AsyncRead + AsyncWrite, Req: ComposeRequest, + ReqMulti: ComposeRequestMulti, { /// Run the transport machinery. pub async fn run(mut self) { @@ -361,7 +615,8 @@ where idle_timeout: self.config.idle_timeout, send_keepalive: true, }; - let mut query_vec = Queries::new(); + let mut query_vec = + Queries::<(ChanReq, Option)>::new(); let mut reqmsg: Option> = None; let mut reqmsg_offset = 0; @@ -454,7 +709,7 @@ where &mut status); }; drop(opt_record); - Self::demux_reply(answer, &mut status, &mut query_vec); + Self::demux_reply(answer, &mut status, &mut query_vec).await; } res = write_stream.write(&msg[reqmsg_offset..]), if do_write => { @@ -479,9 +734,16 @@ where res = recv_fut, if !do_write => { match res { Some(req) => { + if req.sender.is_stream() { + self.config.response_timeout = + self.config.streaming_response_timeout; + } else { + self.config.response_timeout = + self.config.single_response_timeout; + } Self::insert_req( req, &mut status, &mut reqmsg, &mut query_vec - ) + ); } None => { // All references to the connection object have @@ -511,6 +773,8 @@ where } } + trace!("Closing TCP connecting in state: {}", status.state); + // Send FIN _ = write_stream.shutdown().await; } @@ -583,11 +847,14 @@ where } /// Reports an error to all outstanding queries. - fn error(error: Error, query_vec: &mut Queries>) { + fn error( + error: Error, + query_vec: &mut Queries<(ChanReq, Option)>, + ) { // Update all requests that are in progress. Don't wait for // any reply that may be on its way. - for item in query_vec.drain() { - _ = item.sender.send(Err(error.clone())); + for (mut req, _) in query_vec.drain() { + _ = req.sender.send(Err(error.clone())); } } @@ -612,16 +879,18 @@ where /// /// In addition, the status is updated to IdleTimeout or Idle if there /// are no remaining pending requests. - fn demux_reply( + async fn demux_reply( answer: Message, status: &mut Status, - query_vec: &mut Queries>, + query_vec: &mut Queries<(ChanReq, Option)>, ) { // We got an answer, reset the timer status.state = ConnState::Active(Some(Instant::now())); + let id = answer.header().id(); + // Get the correct query and send it the reply. - let req = match query_vec.try_remove(answer.header().id()) { + let (mut req, mut opt_xfr_data) = match query_vec.try_remove(id) { Some(req) => req, None => { // No query with this ID. We should @@ -629,12 +898,32 @@ where return; } }; - let answer = if req.msg.is_answer(answer.for_slice()) { + let mut send_eof = false; + let answer = if match &req.msg { + ReqSingleMulti::Single(msg) => msg.is_answer(answer.for_slice()), + ReqSingleMulti::Multi(msg) => { + let xfr_data = + opt_xfr_data.expect("xfr_data should be present"); + let (eof, xfr_data, is_answer) = + check_stream(msg, xfr_data, &answer); + send_eof = eof; + opt_xfr_data = Some(xfr_data); + is_answer + } + } { Ok(answer) } else { Err(Error::WrongReplyForQuery) }; - _ = req.sender.send(answer); + _ = req.sender.send(answer).await; + + if req.sender.is_stream() { + if send_eof { + _ = req.sender.send_eof().await; + } else { + query_vec.insert_at(id, (req, opt_xfr_data)); + } + } if query_vec.is_empty() { // Clear the activity timer. There is no need to do @@ -660,10 +949,10 @@ where /// idle. Addend a edns-tcp-keepalive option if needed. // Note: maybe reqmsg should be a return value. fn insert_req( - req: ChanReq, + mut req: ChanReq, status: &mut Status, reqmsg: &mut Option>, - query_vec: &mut Queries>, + query_vec: &mut Queries<(ChanReq, Option)>, ) { match &status.state { ConnState::Active(timer) => { @@ -695,12 +984,38 @@ where } } + let xfr_data = match &req.msg { + ReqSingleMulti::Single(_) => None, + ReqSingleMulti::Multi(msg) => { + let qtype = match msg.to_message().and_then(|m| { + m.sole_question() + .map_err(|_| Error::MessageParseError) + .map(|q| q.qtype()) + }) { + Ok(msg) => msg, + Err(e) => { + _ = req.sender.send(Err(e)); + return; + } + }; + if qtype == Rtype::AXFR { + Some(XFRState::AXFRInit) + } else if qtype == Rtype::IXFR { + Some(XFRState::IXFRInit) + } else { + // Stream requests should be either AXFR or IXFR. + _ = req.sender.send(Err(Error::FormError)); + return; + } + } + }; + // Note that insert may fail if there are too many - // outstanding queires. First call insert before checking + // outstanding queries. First call insert before checking // send_keepalive. - let (index, req) = match query_vec.insert(req) { + let (index, (req, _)) = match query_vec.insert((req, xfr_data)) { Ok(res) => res, - Err(req) => { + Err((mut req, _)) => { // Send an appropriate error and return. _ = req .sender @@ -717,11 +1032,21 @@ where // nature of its use of sequence numbers, is far more // resilient against forgery by third parties." - let hdr = req.msg.header_mut(); + let hdr = match &mut req.msg { + ReqSingleMulti::Single(msg) => msg.header_mut(), + ReqSingleMulti::Multi(msg) => msg.header_mut(), + }; hdr.set_id(index); if status.send_keepalive - && req.msg.add_opt(&TcpKeepalive::new(None)).is_ok() + && match &mut req.msg { + ReqSingleMulti::Single(msg) => { + msg.add_opt(&TcpKeepalive::new(None)).is_ok() + } + ReqSingleMulti::Multi(msg) => { + msg.add_opt(&TcpKeepalive::new(None)).is_ok() + } + } { status.send_keepalive = false; } @@ -732,7 +1057,7 @@ where } Err(err) => { // Take the sender out again and return the error. - if let Some(req) = query_vec.try_remove(index) { + if let Some((mut req, _)) = query_vec.try_remove(index) { _ = req.sender.send(Err(err)); } } @@ -748,14 +1073,224 @@ where } /// Convert the query message to a vector. - fn convert_query(msg: &Req) -> Result, Error> { - let mut target = StreamTarget::new_vec(); - msg.append_message(&mut target) - .map_err(|_| Error::StreamLongMessage)?; - Ok(target.into_target()) + fn convert_query( + msg: &ReqSingleMulti, + ) -> Result, Error> { + match msg { + ReqSingleMulti::Single(msg) => { + let mut target = StreamTarget::new_vec(); + msg.append_message(&mut target) + .map_err(|_| Error::StreamLongMessage)?; + Ok(target.into_target()) + } + ReqSingleMulti::Multi(msg) => { + let target = StreamTarget::new_vec(); + let target = msg + .append_message(target) + .map_err(|_| Error::StreamLongMessage)?; + Ok(target.finish().into_target()) + } + } } } +/// Upstate the response stream state based on a response message. +fn check_stream( + msg: &CRM, + mut xfr_state: XFRState, + answer: &Message, +) -> (bool, XFRState, bool) +where + CRM: ComposeRequestMulti, +{ + // First check if the reply matches the request. + // RFC 5936, Section 2.2.2: + // "In the first response message, this section MUST be copied from the + // query. In subsequent messages, this section MAY be copied from the + // query, or it MAY be empty. However, in an error response message + // (see Section 2.2), this section MUST be copied as well." + match xfr_state { + XFRState::AXFRInit | XFRState::IXFRInit => { + if !msg.is_answer(answer.for_slice()) { + xfr_state = XFRState::Error; + // If we detect an error, then keep the stream open. We are + // likely out of sync with respect to the sender. + return (false, xfr_state, false); + } + } + XFRState::AXFRFirstSoa(_) + | XFRState::IXFRFirstSoa(_) + | XFRState::IXFRFirstDiffSoa(_) + | XFRState::IXFRSecondDiffSoa(_) => + // No need to check anything. + {} + XFRState::Done => { + // We should not be here. Switch to error state. + xfr_state = XFRState::Error; + return (false, xfr_state, false); + } + XFRState::Error => + // Keep the stream open. + { + return (false, xfr_state, false) + } + } + + // Then check if the reply status an error. + if answer.header().rcode() != Rcode::NOERROR { + // Also check if this answers the question. + if !msg.is_answer(answer.for_slice()) { + xfr_state = XFRState::Error; + // If we detect an error, then keep the stream open. We are + // likely out of sync with respect to the sender. + return (false, xfr_state, false); + } + return (true, xfr_state, true); + } + + let ans_sec = match answer.answer() { + Ok(ans) => ans, + Err(_) => { + // Bad message, switch to error state. + xfr_state = XFRState::Error; + // If we detect an error, then keep the stream open. + return (true, xfr_state, false); + } + }; + for rr in + ans_sec.into_records::>>() + { + let rr = match rr { + Ok(rr) => rr, + Err(_) => { + // Bad message, switch to error state. + xfr_state = XFRState::Error; + return (true, xfr_state, false); + } + }; + match xfr_state { + XFRState::AXFRInit => { + // The first record has to be a SOA record. + if let AllRecordData::Soa(soa) = rr.data() { + xfr_state = XFRState::AXFRFirstSoa(soa.serial()); + continue; + } + // Bad data. Switch to error status. + xfr_state = XFRState::Error; + return (false, xfr_state, false); + } + XFRState::AXFRFirstSoa(serial) => { + if let AllRecordData::Soa(soa) = rr.data() { + if serial == soa.serial() { + // We found a match. + xfr_state = XFRState::Done; + continue; + } + + // Serial does not match. Move to error state. + xfr_state = XFRState::Error; + return (false, xfr_state, false); + } + + // Any other record, just continue. + } + XFRState::IXFRInit => { + // The first record has to be a SOA record. + if let AllRecordData::Soa(soa) = rr.data() { + xfr_state = XFRState::IXFRFirstSoa(soa.serial()); + continue; + } + // Bad data. Switch to error status. + xfr_state = XFRState::Error; + return (false, xfr_state, false); + } + XFRState::IXFRFirstSoa(serial) => { + // We have three possibilities: + // 1) The record is not a SOA. In that case the format is AXFR. + // 2) The record is a SOA and the serial is not the current + // serial. That is expected for an IXFR format. Move to + // IXFRFirstDiffSoa. + // 3) The record is a SOA and the serial is equal to the + // current serial. Treat this as a strange empty AXFR. + if let AllRecordData::Soa(soa) = rr.data() { + if serial == soa.serial() { + // We found a match. + xfr_state = XFRState::Done; + continue; + } + + xfr_state = XFRState::IXFRFirstDiffSoa(serial); + continue; + } + + // Any other record, move to AXFRFirstSoa. + xfr_state = XFRState::AXFRFirstSoa(serial); + } + XFRState::IXFRFirstDiffSoa(serial) => { + // Move to IXFRSecondDiffSoa if the record is a SOA record, + // otherwise stay in the current state. + if let AllRecordData::Soa(_) = rr.data() { + xfr_state = XFRState::IXFRSecondDiffSoa(serial); + continue; + } + + // Any other record, just continue. + } + XFRState::IXFRSecondDiffSoa(serial) => { + // Move to Done if the record is a SOA record and the + // serial is the one from the first SOA record, move to + // IXFRFirstDiffSoa for any other SOA record and + // otherwise stay in the current state. + if let AllRecordData::Soa(soa) = rr.data() { + if serial == soa.serial() { + // We found a match. + xfr_state = XFRState::Done; + continue; + } + + xfr_state = XFRState::IXFRFirstDiffSoa(serial); + continue; + } + + // Any other record, just continue. + } + XFRState::Done => { + // We got a record after we are done. Switch to error state. + xfr_state = XFRState::Error; + return (false, xfr_state, false); + } + XFRState::Error => panic!("should not be here"), + } + } + + // Check the final state. + match xfr_state { + XFRState::AXFRInit | XFRState::IXFRInit => { + // Still in one of the init state. So the data section was empty. + // Switch to error state. + xfr_state = XFRState::Error; + return (false, xfr_state, false); + } + XFRState::AXFRFirstSoa(_) + | XFRState::IXFRFirstDiffSoa(_) + | XFRState::IXFRSecondDiffSoa(_) => + // Just continue. + {} + XFRState::IXFRFirstSoa(_) => { + // We are still in IXFRFirstSoa. Assume the other side doesn't + // have anything more to say. We could check the SOA serial in + // the request. Just assume that we are done. + xfr_state = XFRState::Done; + return (true, xfr_state, true); + } + XFRState::Done => return (true, xfr_state, true), + XFRState::Error => unreachable!(), + } + + // (eof, xfr_data, is_answer) + (false, xfr_state, true) +} + //------------ Queries ------------------------------------------------------- /// Mapping outstanding queries to their ID. @@ -843,6 +1378,18 @@ impl Queries { Ok((idx, req)) } + /// Inserts the given query at a specified position. A pre-condition is + /// is that the slot has to be empty. + fn insert_at(&mut self, id: u16, req: T) { + let id = id as usize; + self.vec[id] = Some(req); + + self.count += 1; + if id == self.curr { + self.curr += 1; + } + } + /// Tries to remove and return the query at the given index. /// /// Returns `None` if there was no query there. diff --git a/src/net/client/validator.rs b/src/net/client/validator.rs index 1c22e19f..838da3f0 100644 --- a/src/net/client/validator.rs +++ b/src/net/client/validator.rs @@ -70,7 +70,7 @@ //! # let mut msg = msg.question(); //! # msg.push((Name::vec_from_str("example.com").unwrap(), Rtype::AAAA)) //! # .unwrap(); -//! let req = RequestMessage::new(msg); +//! let req = RequestMessage::new(msg).unwrap(); //! //! let ta = TrustAnchors::from_u8(b". 172800 IN DNSKEY 257 3 8 AwEAAaz/tAm8yTn4Mfeh5eyI96WSVexTBAvkMgJzkKTOiW1vkIbzxeF3+/4RgWOq7HrxRixHlFlExOLAJr5emLvN7SWXgnLh4+B5xQlNVz8Og8kvArMtNROxVQuCaSnIDdD5LKyWbRd2n9WGe2R8PzgCmr3EgVLrjyBxWezF0jLHwVN8efS3rCj/EWgvIWgb9tarpVUDK/b58Da+sqqls3eNbuv7pr+eoZG+SrDK6nWeL3c6H5Apxz7LjVc1uTIdsIXxuOLYA4/ilBmSVIzuDWfdRUfhHdY6+cn8HFRm+2hM8AnXGXws9555KrUB5qihylGa8subX2Nn6UwNR1AkUTV74bU= ;{id = 20326 (ksk), size = 2048b} ;;state=2 [ VALID ] ;;count=0 ;;lastchange=1683463064 ;;Sun May 7 12:37:44 2023").unwrap(); //! let vc = ValidationContext::new(ta, udptcp_conn.clone()); diff --git a/src/net/server/tests/integration.rs b/src/net/server/tests/integration.rs index bf6af2a9..ff18b633 100644 --- a/src/net/server/tests/integration.rs +++ b/src/net/server/tests/integration.rs @@ -16,6 +16,7 @@ use crate::base::iana::Rcode; use crate::base::name::{Name, ToName}; use crate::base::net::IpAddr; use crate::base::wire::Composer; +use crate::net::client::request::RequestMessageMulti; use crate::net::client::{dgram, stream}; use crate::net::server; use crate::net::server::buf::VecBufSource; @@ -193,7 +194,10 @@ fn mk_client_factory( move |source_addr| { let stream = stream_server_conn .connect(Some(SocketAddr::new(*source_addr, 0))); - let (conn, transport) = stream::Connection::new(stream); + let (conn, transport) = stream::Connection::< + _, + RequestMessageMulti>, + >::new(stream); tokio::spawn(transport.run()); Box::new(conn) }, diff --git a/src/resolv/stub/mod.rs b/src/resolv/stub/mod.rs index af66f170..f80ee069 100644 --- a/src/resolv/stub/mod.rs +++ b/src/resolv/stub/mod.rs @@ -395,7 +395,9 @@ impl<'a> Query<'a> { let msg = Message::from_octets(message.as_target().to_vec()) .expect("Message::from_octets should not fail"); - let request_msg = RequestMessage::new(msg); + let request_msg = RequestMessage::new(msg).map_err(|e| { + io::Error::new(io::ErrorKind::Other, e.to_string()) + })?; let transport = self.resolver.get_transport().await.map_err(|e| { io::Error::new(io::ErrorKind::Other, e.to_string()) diff --git a/src/stelline/client.rs b/src/stelline/client.rs index 6d8f85ba..3aea49b1 100644 --- a/src/stelline/client.rs +++ b/src/stelline/client.rs @@ -517,7 +517,8 @@ fn entry2reqmsg(entry: &Entry) -> RequestMessage> { header.set_cd(reply.cd); let msg = msg.into_message(); - let mut reqmsg = RequestMessage::new(msg); + let mut reqmsg = RequestMessage::new(msg) + .expect("should not fail unless the request is AXFR"); reqmsg.set_dnssec_ok(reply.fl_do); if reply.notify { reqmsg.header_mut().set_opcode(Opcode::NOTIFY); diff --git a/src/validator/context.rs b/src/validator/context.rs index 9052179e..d9e44ba1 100644 --- a/src/validator/context.rs +++ b/src/validator/context.rs @@ -2221,7 +2221,7 @@ where } }; let msg = Message::from_octets(octs).expect("should not fail"); - let mut req = RequestMessage::new(msg); + let mut req = RequestMessage::new(msg).expect("should not fail"); req.set_dnssec_ok(true); let mut request = upstream.send_request(req); diff --git a/src/validator/mod.rs b/src/validator/mod.rs index e2d22329..af70e18f 100644 --- a/src/validator/mod.rs +++ b/src/validator/mod.rs @@ -80,7 +80,7 @@ //! # let mut msg = msg.question(); //! # msg.push((Name::vec_from_str("example.com").unwrap(), Rtype::AAAA)) //! # .unwrap(); -//! let mut req = RequestMessage::new(msg); +//! let mut req = RequestMessage::new(msg).unwrap(); //! req.set_dnssec_ok(true); //! //! // Send a query message. diff --git a/tests/net-client-cache.rs b/tests/net-client-cache.rs index 01f97bb6..5f3f2b5e 100644 --- a/tests/net-client-cache.rs +++ b/tests/net-client-cache.rs @@ -84,7 +84,7 @@ async fn test_transport_error() { let mut msg = msg.question(); msg.push((Name::vec_from_str("example.com").unwrap(), Rtype::AAAA)) .unwrap(); - let req = RequestMessage::new(msg); + let req = RequestMessage::new(msg).unwrap(); let mut request = cached.send_request(req.clone()); let reply = request.get_response().await; diff --git a/tests/net-client.rs b/tests/net-client.rs index a1037172..0996b822 100644 --- a/tests/net-client.rs +++ b/tests/net-client.rs @@ -11,6 +11,7 @@ use domain::net::client::dgram; use domain::net::client::dgram_stream; use domain::net::client::multi_stream; use domain::net::client::redundant; +use domain::net::client::request::RequestMessageMulti; use domain::net::client::stream; use std::fs::File; use std::net::IpAddr; @@ -43,7 +44,8 @@ fn single() { let step_value = Arc::new(CurrStepValue::new()); let conn = Connection::new(stelline.clone(), step_value.clone()); - let (octstr, transport) = stream::Connection::new(conn); + let (octstr, transport) = + stream::Connection::<_, RequestMessageMulti>>::new(conn); tokio::spawn(async move { transport.run().await; }); @@ -138,7 +140,10 @@ fn tcp() { } }; - let (tcp, transport) = stream::Connection::new(tcp_conn); + let (tcp, transport) = stream::Connection::< + _, + RequestMessageMulti>, + >::new(tcp_conn); tokio::spawn(async move { transport.run().await; println!("single TCP run terminated");