use std::
{
cell::RefCell,
os::fd::AsFd,
sync::Arc,
time::Duration,
fmt,
fmt::Debug,
io::{ErrorKind, Write, prelude::*},
net::{SocketAddr, UdpSocket, TcpStream}
};
use socket2::{Socket, Domain, Type, Protocol, SockAddr};
use crate::{internal_error, internal_error_map, parsers::cfg_resolv_parser::ResolveConfEntry, error::*};
pub trait SocketTapCommon
{
fn get_remote_addr(&self) -> &SocketAddr;
}
pub trait SocketTap: SocketTapCommon + PartialEq<ResolveConfEntry>
{
fn is_tcp(&self) -> bool;
fn should_append_len(&self) -> bool;
fn is_encrypted(&self) -> bool;
fn send(&self, sndbuf: &[u8]) -> CDnsResult<usize> ;
fn recv(&self) -> CDnsResult<Option<Vec<u8>>>;
fn reconnect(&self) -> CDnsResult<()>;
}
impl Debug for dyn SocketTap
{
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result
{
write!(f, "{:?}", self)
}
}
#[cfg(feature = "use_sync_tls")]
pub(crate) mod with_tls
{
use std::cell::RefCell;
use std::io::{ErrorKind, Read, Write};
use std::net::TcpStream;
use std::os::fd::{AsFd, BorrowedFd};
use std::sync::Arc;
use std::time::Duration;
use rustls::pki_types::ServerName;
use rustls::{ClientConnection, RootCertStore, StreamOwned};
use crate::CDnsErrorDnsResponse;
use crate::network::{NetworkTap, SocketTap, new_tcp_stream};
use crate::parsers::cfg_resolv_parser::ResolveConfEntry;
use crate::{internal_error, internal_error_map, CDnsErrorType, CDnsResult, common::{DEF_USERAGENT}};
#[derive(Debug)]
pub(crate) struct TcpHttpsConnection
{
stream: TcpTlsConnection,
}
impl AsFd for TcpHttpsConnection
{
fn as_fd(&self) -> BorrowedFd<'_>
{
return self.stream.as_fd();
}
}
impl TcpHttpsConnection
{
fn connect(cfg: &ResolveConfEntry, conn_timeout: Option<Duration>, timeout: Option<Duration>) -> CDnsResult<Self>
{
let tls_conn = TcpTlsConnection::connect(cfg, conn_timeout, timeout)?;
return Ok(TcpHttpsConnection { stream: tls_conn });
}
}
#[derive(Debug)]
pub(crate) struct TcpTlsConnection
{
stream: StreamOwned<ClientConnection, TcpStream>,
}
impl AsFd for TcpTlsConnection
{
fn as_fd(&self) -> BorrowedFd<'_>
{
return self.stream.sock.as_fd();
}
}
impl TcpTlsConnection
{
fn connect(cfg: &ResolveConfEntry, conn_timeout: Option<Duration>, timeout: Option<Duration>) -> CDnsResult<Self>
{
let domain_name =
if let Some(domainname) = cfg.get_tls_domain()
{
ServerName::try_from(domainname.clone())
.map_err(|e|
internal_error_map!(CDnsErrorType::InternalError, "{}", e)
)?
}
else
{
internal_error!(CDnsErrorType::InternalError, "no domain is set for TLS conncection");
};
let config =
rustls
::ClientConfig
::builder_with_protocol_versions(&[&rustls::version::TLS12])
.with_root_certificates(RootCertStore{roots: webpki_roots::TLS_SERVER_ROOTS.into()})
.with_no_client_auth();
let conn =
rustls
::ClientConnection
::new(Arc::new(config), domain_name)
.map_err(|e| internal_error_map!(CDnsErrorType::InternalError, "{}", e))?;
let socket = new_tcp_stream(cfg, conn_timeout, timeout)?;
let mut tlssock = rustls::StreamOwned::new(conn, socket);
tlssock.flush().map_err(|e| internal_error_map!(CDnsErrorType::IoError, "{}", e))?;
return Ok( Self{ stream: tlssock } );
}
}
impl NetworkTap<TcpHttpsConnection>
{
fn sub_recv(&self, rcvbuf: &mut [u8]) -> CDnsResult<usize>
{
loop
{
match self.sock.borrow_mut().stream.stream.read(rcvbuf)
{
Ok(n) =>
{
return Ok(n);
},
Err(ref e) if e.kind() == ErrorKind::WouldBlock =>
{
return Ok(0);
},
Err(ref e) if e.kind() == ErrorKind::Interrupted =>
{
continue;
},
Err(e) =>
{
internal_error!(CDnsErrorType::IoError, "{}", e);
}
}
}
}
}
impl SocketTap for NetworkTap<TcpHttpsConnection>
{
fn is_encrypted(&self) -> bool
{
return true;
}
fn is_tcp(&self) -> bool
{
return true;
}
fn should_append_len(&self) -> bool
{
return false;
}
fn send(&self, sndbuf: &[u8]) -> CDnsResult<usize>
{
let url_path =
self.cfg.get_tls_path().map_or("", |f| f.as_str());
let host =
self.cfg.get_tls_domain().map_or(self.cfg.get_resolver_ip().to_string(), |f| f.clone());
let mut http_req =
[
"POST /", url_path," HTTP/1.1\r\n",
"Host: ", host.as_str(), "\r\n",
"Content-Type: application/dns-message\r\n",
"Accept: application/dns-message\r\n",
"User-Agent: ", DEF_USERAGENT, "\r\n",
"Content-Length: ", sndbuf.len().to_string().as_str(), "\r\n\r\n",
]
.concat()
.into_bytes();
println!("{}", http_req.len());
http_req.extend(sndbuf);
println!("{}", http_req.len());
return
self
.sock
.borrow_mut()
.stream
.stream
.write_all(&http_req)
.map_err(|e|
internal_error_map!(CDnsErrorType::IoError, "{}", e)
)
.map(|_| http_req.len());
}
fn recv(&self) -> CDnsResult<Option<Vec<u8>>>
{
panic!("DNS-over-HTTPS is not implemented");
}
fn reconnect(&self) -> CDnsResult<()>
{
let socket=
TcpHttpsConnection
::connect(
&self.cfg,
self.conn_timeout,
Some(self.timeout)
)?;
*self.sock.borrow_mut() = socket;
return Ok(());
}
}
impl NetworkTap<TcpTlsConnection>
{
fn sub_read(&self, rcvbuf: &mut [u8]) -> CDnsResult<usize>
{
loop
{
match self.sock.borrow_mut().stream.read(rcvbuf)
{
Ok(n) =>
{
return Ok(n);
},
Err(ref e) if e.kind() == ErrorKind::WouldBlock =>
{
return Ok(0);
},
Err(ref e) if e.kind() == ErrorKind::Interrupted =>
{
continue;
},
Err(e) =>
{
internal_error!(CDnsErrorType::IoError, "{}", e);
}
}
}
}
}
impl SocketTap for NetworkTap<TcpTlsConnection>
{
fn is_encrypted(&self) -> bool
{
return true;
}
fn is_tcp(&self) -> bool
{
return true;
}
fn should_append_len(&self) -> bool
{
return true;
}
fn send(&self, sndbuf: &[u8]) -> CDnsResult<usize>
{
return
self
.sock
.borrow_mut()
.stream
.write_all(sndbuf)
.map_err(|e| internal_error_map!(CDnsErrorType::IoError, "{}", e))
.map(|_| sndbuf.len());
}
fn recv(&self) -> CDnsResult<Option<Vec<u8>>>
{
let mut pkg_pen: [u8; 2] = [0, 0];
let n = self.sub_read(&mut pkg_pen)?;
if n == 0
{
internal_error!(CDnsErrorType::IoError, "tcp received zero len message!");
}
else if n != 2
{
internal_error!(CDnsErrorType::IoError, "tcp expected 2 bytes to be read!");
}
let ln = u16::from_be_bytes(pkg_pen);
let mut rcvbuf = vec![0_u8; ln as usize];
let mut n = self.sub_read(rcvbuf.as_mut_slice())?;
if n == 0
{
return Ok(None);
}
else if n == 1
{
n = self.sub_read(&mut rcvbuf[1..])?;
if n == 0
{
return Ok(None);
}
}
return Ok(Some(rcvbuf));
}
fn reconnect(&self) -> CDnsResult<()>
{
let socket=
TcpTlsConnection
::connect(
&self.cfg,
self.conn_timeout,
Some(self.timeout),
)?;
*self.sock.borrow_mut() = socket;
return Ok(());
}
}
impl NetworkTap<TcpTlsConnection>
{
pub(crate)
fn new_tls(
resolver: Arc<ResolveConfEntry>,
timeout: Duration,
conn_timeout: Option<Duration>
) -> CDnsResult<Box<dyn SocketTap>>
where
NetworkTap<TcpTlsConnection>: SocketTap + 'static
{
let socket=
TcpTlsConnection::connect(resolver.as_ref(), conn_timeout, Some(timeout))?;
let ret =
NetworkTap::<TcpTlsConnection>
{
sock:
RefCell::new(socket),
timeout:
timeout,
conn_timeout:
conn_timeout,
cfg:
resolver,
};
return Ok(Box::new(ret));
}
}
impl NetworkTap<TcpHttpsConnection>
{
pub(crate)
fn new_https(
resolver: Arc<ResolveConfEntry>,
timeout: Duration,
conn_timeout: Option<Duration>
) -> CDnsResult<Box<dyn SocketTap>>
where
NetworkTap<TcpHttpsConnection>: SocketTap + 'static
{
let socket=
TcpHttpsConnection::connect(resolver.as_ref(), conn_timeout, Some(timeout))?;
let ret =
NetworkTap::<TcpHttpsConnection>
{
sock:
RefCell::new(socket),
timeout:
timeout,
conn_timeout:
conn_timeout,
cfg:
resolver,
};
return Ok(Box::new(ret));
}
}
}
pub(crate) struct NetworkTap<T: AsFd>
{
sock: RefCell<T>,
timeout: Duration,
conn_timeout: Option<Duration>,
cfg: Arc<ResolveConfEntry>,
}
impl<T: AsFd> NetworkTap<T>
{
fn get_remote_addr(&self) -> &SocketAddr
{
return &self.cfg.get_resolver_sa();
}
}
impl<T: AsFd> SocketTapCommon for NetworkTap<T>
{
fn get_remote_addr(&self) -> &SocketAddr
{
return self.cfg.get_resolver_sa();
}
}
impl<T: AsFd> PartialEq<ResolveConfEntry> for NetworkTap<T>
{
fn eq(&self, other: &ResolveConfEntry) -> bool
{
return self.cfg.as_ref() == other;
}
}
impl<T: AsFd + Debug> fmt::Debug for NetworkTap<T>
{
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result
{
f
.debug_struct("NetworkTap")
.field("sock", &self.sock)
.field("timeout", &self.timeout)
.field("cfg", &self.cfg)
.finish()
}
}
impl NetworkTap<TcpStream>
{
pub(crate)
fn new_tcp(
resolver: Arc<ResolveConfEntry>,
timeout: Duration,
conn_timeout: Option<Duration>
) -> CDnsResult<Box<dyn SocketTap>>
where
NetworkTap<TcpStream>: SocketTap + 'static
{
let socket =
new_tcp_stream(&resolver, conn_timeout, Some(timeout))?;
let ret =
NetworkTap::<TcpStream>
{
sock:
RefCell::new(socket),
timeout:
timeout,
conn_timeout:
conn_timeout,
cfg:
resolver,
};
return Ok(Box::new(ret));
}
}
impl NetworkTap<UdpSocket>
{
fn new_udp_socket(resolver: &ResolveConfEntry, timeout: Duration) -> CDnsResult<UdpSocket>
{
let socket =
UdpSocket::bind(resolver.get_adapter_ip())
.map_err(|e| internal_error_map!(CDnsErrorType::InternalError, "{}", e))?;
socket
.set_nonblocking(false)
.map_err(|e| internal_error_map!(CDnsErrorType::IoError, "{}", e))?;
socket
.set_read_timeout(Some(timeout))
.map_err(|e| internal_error_map!(CDnsErrorType::IoError, "{}", e))?;
socket
.set_write_timeout(Some(timeout)) .map_err(|e| internal_error_map!(CDnsErrorType::IoError, "{}", e))?;
socket
.connect(&resolver.get_resolver_sa())
.map_err(|e| internal_error_map!(CDnsErrorType::IoError, "{}", e))?;
return Ok(socket);
}
pub(crate)
fn new_udp(
resolver: Arc<ResolveConfEntry>,
timeout: Duration
) -> CDnsResult<Box<dyn SocketTap>>
where
NetworkTap<UdpSocket>: SocketTap + 'static
{
let ret =
NetworkTap::<UdpSocket>
{
sock:
RefCell::new(
NetworkTap::<UdpSocket>::new_udp_socket(&resolver, timeout)?
),
timeout:
timeout,
conn_timeout:
None,
cfg:
resolver,
};
return Ok(Box::new(ret));
}
}
impl SocketTap for NetworkTap<UdpSocket>
{
fn is_tcp(&self) -> bool
{
return false;
}
fn is_encrypted(&self) -> bool
{
return false;
}
fn should_append_len(&self) -> bool
{
return false;
}
fn send(&self, sndbuf: &[u8]) -> CDnsResult<usize>
{
return
self
.sock
.borrow()
.send(sndbuf)
.map_err(|e| internal_error_map!(CDnsErrorType::IoError, "{}", e));
}
fn recv(&self) -> CDnsResult<Option<Vec<u8>>>
{
let mut rcvbuf = vec![0_u8; 1457];
let _n =
loop
{
match self.sock.borrow().recv_from(&mut rcvbuf)
{
Ok((rcv_len, rcv_src)) =>
{
if &rcv_src != self.get_remote_addr()
{
internal_error!(
CDnsErrorType::DnsResponse(CDnsErrorDnsResponse::DnsResponseFromUnknownDestination),
"received answer from unknown host: '{}' exp: '{}'",
self.get_remote_addr(),
rcv_src
);
}
break Ok(rcv_len);
},
Err(ref e) if e.kind() == ErrorKind::WouldBlock =>
{
return Ok(None);
},
Err(ref e) if e.kind() == ErrorKind::Interrupted =>
{
continue;
},
Err(e) =>
{
internal_error!(CDnsErrorType::IoError, "{}", e);
}
} }?;
return Ok(Some(rcvbuf));
}
fn reconnect(&self) -> CDnsResult<()>
{
let socket = NetworkTap::<UdpSocket>::new_udp_socket(&self.cfg, self.timeout)?;
*self.sock.borrow_mut() = socket;
return Ok(());
}
}
fn new_tcp_stream(cfg: &ResolveConfEntry, conn_timeout: Option<Duration>, timeout: Option<Duration>) -> CDnsResult<TcpStream>
{
let socket =
Socket::new(Domain::for_address(*cfg.get_resolver_sa()), Type::STREAM, Some(Protocol::TCP))
.map_err(|e| internal_error_map!(CDnsErrorType::IoError, "{}", e))?;
socket
.bind(&SockAddr::from(*cfg.get_adapter_ip()))
.map_err(|e| internal_error_map!(CDnsErrorType::IoError, "{}", e))?;
socket.set_tcp_nodelay(true).map_err(|e| internal_error_map!(CDnsErrorType::IoError, "{}", e))?;
socket.set_keepalive(false).map_err(|e| internal_error_map!(CDnsErrorType::IoError, "{}", e))?;
if let Some(c_timeout) = conn_timeout
{
socket
.connect_timeout(&SockAddr::from(*cfg.get_resolver_sa()), c_timeout)
.map_err(|e| internal_error_map!(CDnsErrorType::IoError, "{}", e))?;
}
else
{
socket
.connect(&SockAddr::from(*cfg.get_resolver_sa()))
.map_err(|e| internal_error_map!(CDnsErrorType::IoError, "{}", e))?;
}
let socket: TcpStream = socket.into();
socket
.set_nonblocking(false)
.map_err(|e| internal_error_map!(CDnsErrorType::IoError, "{}", e))?;
socket
.set_read_timeout(timeout) .map_err(|e| internal_error_map!(CDnsErrorType::IoError, "{}", e))?;
socket
.set_write_timeout(timeout) .map_err(|e| internal_error_map!(CDnsErrorType::IoError, "{}", e))?;
return Ok(socket);
}
impl NetworkTap<TcpStream>
{
fn sub_read(&self, rcvbuf: &mut [u8]) -> CDnsResult<usize>
{
loop
{
match self.sock.borrow_mut().read(rcvbuf)
{
Ok(n) =>
{
return Ok(n);
},
Err(ref e) if e.kind() == ErrorKind::WouldBlock =>
{
return Ok(0);
},
Err(ref e) if e.kind() == ErrorKind::Interrupted =>
{
continue;
},
Err(e) =>
{
internal_error!(CDnsErrorType::IoError, "{}", e);
}
}
}
}
}
impl SocketTap for NetworkTap<TcpStream>
{
fn is_tcp(&self) -> bool
{
return true;
}
fn is_encrypted(&self) -> bool
{
return false;
}
fn should_append_len(&self) -> bool
{
return true;
}
fn send(&self, sndbuf: &[u8]) -> CDnsResult<usize>
{
let res =
self
.sock
.borrow_mut()
.write(sndbuf);
match res
{
Ok(_) =>
return
res.map_err(|e| internal_error_map!(CDnsErrorType::IoError, "{}", e)),
Err(ref e)
if e.kind() == ErrorKind::NotConnected || e.kind() == ErrorKind::ConnectionAborted || e.kind() == ErrorKind::ConnectionReset =>
{
self.reconnect()?;
return
self
.sock
.borrow_mut()
.write(sndbuf)
.map_err(|e| internal_error_map!(CDnsErrorType::IoError, "{}", e));
},
Err(_) =>
return
res.map_err(|e| internal_error_map!(CDnsErrorType::IoError, "{}", e))
}
}
fn recv(&self) -> CDnsResult<Option<Vec<u8>>>
{
let mut pkg_pen: [u8; 2] = [0, 0];
let n = self.sub_read(&mut pkg_pen)?;
if n == 0
{
internal_error!(CDnsErrorType::IoError, "tcp received zero len message!");
}
else if n != 2
{
internal_error!(CDnsErrorType::IoError, "tcp expected 2 bytes to be read!");
}
let ln = u16::from_be_bytes(pkg_pen);
let mut rcvbuf = vec![0_u8; ln as usize];
let mut n = self.sub_read(rcvbuf.as_mut_slice())?;
if n == 0
{
return Ok(None);
}
else if n == 1
{
n = self.sub_read(&mut rcvbuf[1..])?;
if n == 0
{
return Ok(None);
}
}
return Ok(Some(rcvbuf));
}
fn reconnect(&self) -> CDnsResult<()>
{
let socket =
new_tcp_stream(&self.cfg, self.conn_timeout, Some(self.timeout))?;
*self.sock.borrow_mut() = socket;
return Ok(());
}
}