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.
This commit is contained in:
Philip Homburg
2024-09-04 14:03:12 +02:00
parent 851d261d31
commit 90cc96d632
14 changed files with 1044 additions and 99 deletions
+6 -3
View File
@@ -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<Vec<u8>>>::new(tcp_conn);
tokio::spawn(async move {
transport.run().await;
println!("single TCP run terminated");
+8
View File
@@ -439,6 +439,14 @@ impl<Octs: Octets + ?Sized> Message<Octs> {
}
}
/// 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
+47 -8
View File
@@ -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<Vec<u8>>>::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<Vec<u8>>>::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;
//! # }
//! ```
//!
//! <div class="warning">
//!
//! **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::<RequestMessage<Vec<u8>>, _>::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 {
//! // ...
//! }
//! # }
//! ```
//!
//! </div>
//!
//! # Limitations
//!
+5 -4
View File
@@ -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<Req> {
ReceiveConn(oneshot::Receiver<ChanResp<Req>>),
/// Start a query using the given stream transport.
StartQuery(Arc<stream::Connection<Req>>),
StartQuery(Arc<stream::Connection<Req, RequestMessageMulti<Vec<u8>>>>),
/// Get the result of the query.
GetResult(stream::Request),
@@ -222,7 +223,7 @@ struct ChanRespOk<Req> {
id: u64,
/// The new stream transport to use for sending a request.
conn: Arc<stream::Connection<Req>>,
conn: Arc<stream::Connection<Req, RequestMessageMulti<Vec<u8>>>>,
}
impl<Req> Request<Req> {
@@ -409,7 +410,7 @@ enum SingleConnState3<Req> {
None,
/// Current stream transport.
Some(Arc<stream::Connection<Req>>),
Some(Arc<stream::Connection<Req, RequestMessageMulti<Vec<u8>>>>),
/// State that deals with an error getting a new octet stream from
/// a connection stream.
+349 -14
View File
@@ -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<Target: Composer>(
&self,
target: &mut Target,
) -> Result<(), CopyRecordsError>;
target: Target,
) -> Result<AdditionalBuilder<Target>, CopyRecordsError>;
/// Create a message that captures the recorded changes.
fn to_message(&self) -> Result<Message<Vec<u8>>, 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<Target: Composer>(
&self,
target: Target,
) -> Result<AdditionalBuilder<Target>, CopyRecordsError>;
/// Create a message that captures the recorded changes.
fn to_message(&self) -> Result<Message<Vec<u8>>, 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<CR> {
) -> Box<dyn GetResponse + Send + Sync>;
}
//------------ 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<CR> {
/// Request function that takes a ComposeRequestMulti type.
fn send_request(
&self,
request_msg: CR,
) -> Box<dyn GetResponseMulti + Send + Sync>;
}
//------------ 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<Output = Result<Option<Message<Bytes>>, Error>>
+ Send
+ Sync
+ '_,
>,
>;
}
//------------ RequestMessage ------------------------------------------------
/// Object that implements the ComposeRequest trait for a Message object.
@@ -114,15 +193,26 @@ pub struct RequestMessage<Octs: AsRef<[u8]>> {
}
impl<Octs: AsRef<[u8]> + Debug + Octets> RequestMessage<Octs> {
/// Create a new BMB object.
pub fn new(msg: impl Into<Message<Octs>>) -> Self {
/// Create a new RequestMessage object.
pub fn new(msg: impl Into<Message<Octs>>) -> Result<Self, Error> {
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<Octs: AsRef<[u8]> + Debug + Octets + Send + Sync> ComposeRequest
{
fn append_message<Target: Composer>(
&self,
target: &mut Target,
) -> Result<(), CopyRecordsError> {
target: Target,
) -> Result<AdditionalBuilder<Target>, 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<Vec<u8>, Error> {
@@ -298,6 +388,222 @@ impl<Octs: AsRef<[u8]> + Debug + Octets + Send + Sync> ComposeRequest
}
}
//------------ RequestMessageMulti --------------------------------------------
/// Object that implements the ComposeRequestMulti trait for a Message object.
#[derive(Clone, Debug)]
pub struct RequestMessageMulti<Octs>
where
Octs: AsRef<[u8]>,
{
/// Base message.
msg: Message<Octs>,
/// New header.
header: Header,
/// The OPT record to add if required.
opt: Option<OptRecord<Vec<u8>>>,
}
impl<Octs: AsRef<[u8]> + Debug + Octets> RequestMessageMulti<Octs> {
/// Create a new BMB object.
pub fn new(msg: impl Into<Message<Octs>>) -> Result<Self, Error> {
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<Vec<u8>> {
self.opt.get_or_insert_with(Default::default)
}
/// Appends the message to a composer.
fn append_message_impl<Target: Composer>(
&self,
mut target: MessageBuilder<Target>,
) -> Result<AdditionalBuilder<Target>, 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::<AllRecordData<_, ParsedName<_>>>()?
.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::<AllRecordData<_, ParsedName<_>>>()?
.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::<AllRecordData<_, ParsedName<_>>>()?
.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<Message<Vec<u8>>, 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<Octs: AsRef<[u8]> + Debug + Octets + Send + Sync> ComposeRequestMulti
for RequestMessageMulti<Octs>
{
fn append_message<Target: Composer>(
&self,
target: Target,
) -> Result<AdditionalBuilder<Target>, 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<Message<Vec<u8>>, 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<super::dgram::QueryError>),
#[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),
}
+608 -61
View File
@@ -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 thats 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<Req> {
pub struct Connection<Req, ReqMulti> {
/// The sender half of the request channel.
sender: mpsc::Sender<ChanReq<Req>>,
sender: mpsc::Sender<ChanReq<Req, ReqMulti>>,
}
impl<Req> Connection<Req> {
impl<Req, ReqMulti> Connection<Req, ReqMulti> {
/// 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: Stream) -> (Self, Transport<Stream, Req>) {
pub fn new<Stream>(
stream: Stream,
) -> (Self, Transport<Stream, Req, ReqMulti>) {
Self::with_config(stream, Default::default())
}
@@ -150,13 +196,17 @@ impl<Req> Connection<Req> {
pub fn with_config<Stream>(
stream: Stream,
config: Config,
) -> (Self, Transport<Stream, Req>) {
) -> (Self, Transport<Stream, Req, ReqMulti>) {
let (sender, transport) = Transport::new(stream, config);
(Self { sender }, transport)
}
}
impl<Req: ComposeRequest + 'static> Connection<Req> {
impl<Req, ReqMulti> Connection<Req, ReqMulti>
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<Req: ComposeRequest + 'static> Connection<Req> {
msg: Req,
) -> Result<Message<Bytes>, 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<Req: ComposeRequest + 'static> Connection<Req> {
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<Result<Option<Message<Bytes>>, 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<Req> Clone for Connection<Req> {
impl<Req, ReqMulti> Clone for Connection<Req, ReqMulti> {
fn clone(&self) -> Self {
Self {
sender: self.sender.clone(),
@@ -191,8 +275,10 @@ impl<Req> Clone for Connection<Req> {
}
}
impl<Req: ComposeRequest + Clone + 'static> SendRequest<Req>
for Connection<Req>
impl<Req, ReqMulti> SendRequest<Req> for Connection<Req, ReqMulti>
where
Req: ComposeRequest + 'static,
ReqMulti: ComposeRequestMulti + Debug + Send + Sync + 'static,
{
fn send_request(
&self,
@@ -202,6 +288,19 @@ impl<Req: ComposeRequest + Clone + 'static> SendRequest<Req>
}
}
impl<Req, ReqMulti> SendRequestMulti<ReqMulti> for Connection<Req, ReqMulti>
where
Req: ComposeRequest + Debug + Send + Sync + 'static,
ReqMulti: ComposeRequestMulti + 'static,
{
fn send_request(
&self,
request_msg: ReqMulti,
) -> Box<dyn GetResponseMulti + Send + Sync> {
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<Result<Option<Message<Bytes>>, Error>>,
/// The underlying future.
#[allow(clippy::type_complexity)]
fut: Option<
Pin<Box<dyn Future<Output = Result<(), Error>> + Send + Sync>>,
>,
}
impl RequestMulti {
/// Async function that waits for the future stored in Request to complete.
async fn get_response_impl(
&mut self,
) -> Result<Option<Message<Bytes>>, 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<Output = Result<Option<Message<Bytes>>, 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<Stream, Req> {
/// The stream socket towards the remove end.
pub struct Transport<Stream, Req, ReqMulti> {
/// The stream socket towards the remote end.
stream: Stream,
/// Transport configuration.
config: Config,
/// The receiver half of request channel.
receiver: mpsc::Receiver<ChanReq<Req>>,
receiver: mpsc::Receiver<ChanReq<Req, ReqMulti>>,
}
/// This is the type of sender in [ChanReq].
#[derive(Debug)]
enum ReplySender {
/// Return channel for a single response.
Single(Option<oneshot::Sender<ChanResp>>),
/// Return channel for a stream of responses.
Stream(mpsc::Sender<Result<Option<Message<Bytes>>, 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<Req, ReqMulti> {
/// Single response request.
Single(Req),
/// Multi-response request.
Multi(ReqMulti),
}
/// A message from a [`Request`] to start a new request.
#[derive(Debug)]
struct ChanReq<Req> {
struct ChanReq<Req, ReqMulti> {
/// DNS request message
msg: Req,
msg: ReqSingleMulti<Req, ReqMulti>,
/// Sender to send result back to [`Request`]
sender: ReplySender,
}
/// This is the type of sender in [ChanReq].
type ReplySender = oneshot::Sender<ChanResp>;
/// A message back to [`Request`] returning a response.
type ChanResp = Result<Message<Bytes>, Error>;
@@ -323,12 +528,60 @@ enum ConnState {
WriteError(Error),
}
impl<Stream, Req> Transport<Stream, Req> {
//--- 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<Stream, Req, ReqMulti> Transport<Stream, Req, ReqMulti> {
/// Creates a new transport.
fn new(
stream: Stream,
config: Config,
) -> (mpsc::Sender<ChanReq<Req>>, Self) {
) -> (mpsc::Sender<ChanReq<Req, ReqMulti>>, Self) {
let (sender, receiver) = mpsc::channel(DEF_CHAN_CAP);
(
sender,
@@ -341,10 +594,11 @@ impl<Stream, Req> Transport<Stream, Req> {
}
}
impl<Stream, Req> Transport<Stream, Req>
impl<Stream, Req, ReqMulti> Transport<Stream, Req, ReqMulti>
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<Req, ReqMulti>, Option<XFRState>)>::new();
let mut reqmsg: Option<Vec<u8>> = 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<ChanReq<Req>>) {
fn error(
error: Error,
query_vec: &mut Queries<(ChanReq<Req, ReqMulti>, Option<XFRState>)>,
) {
// 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<Bytes>,
status: &mut Status,
query_vec: &mut Queries<ChanReq<Req>>,
query_vec: &mut Queries<(ChanReq<Req, ReqMulti>, Option<XFRState>)>,
) {
// 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<Req>,
mut req: ChanReq<Req, ReqMulti>,
status: &mut Status,
reqmsg: &mut Option<Vec<u8>>,
query_vec: &mut Queries<ChanReq<Req>>,
query_vec: &mut Queries<(ChanReq<Req, ReqMulti>, Option<XFRState>)>,
) {
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<Vec<u8>, 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<Req, ReqMulti>,
) -> Result<Vec<u8>, 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<CRM>(
msg: &CRM,
mut xfr_state: XFRState,
answer: &Message<Bytes>,
) -> (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::<AllRecordData<Bytes, ParsedName<Bytes>>>()
{
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<T> Queries<T> {
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.
+1 -1
View File
@@ -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());
+5 -1
View File
@@ -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<Vec<u8>>,
>::new(stream);
tokio::spawn(transport.run());
Box::new(conn)
},
+3 -1
View File
@@ -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())
+2 -1
View File
@@ -517,7 +517,8 @@ fn entry2reqmsg(entry: &Entry) -> RequestMessage<Vec<u8>> {
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);
+1 -1
View File
@@ -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);
+1 -1
View File
@@ -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.
+1 -1
View File
@@ -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;
+7 -2
View File
@@ -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<Vec<u8>>>::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<Vec<u8>>,
>::new(tcp_conn);
tokio::spawn(async move {
transport.run().await;
println!("single TCP run terminated");