use crate::{
RuntimeError,
futures::{
net::{
exchange::Sends,
stream::{FinishTask, Io, Pipe, RecvTask, SendTask, Source},
},
tcp::{Connection, Listener},
tls::{
fd_io::FdIo,
handshake::{io_error, process, send_pending},
tls_task::TlsAcceptTask,
},
},
};
use rustls::ServerConfig;
use std::{
fmt,
io::{self, Read, Write},
net::SocketAddr,
sync::{Arc, Mutex, MutexGuard},
time::Duration,
};
struct TlsStream {
tls: Mutex<rustls::Connection>,
tcp: Connection,
}
const DRAIN_READS: usize = 16;
impl Drop for TlsStream {
fn drop(&mut self) {
let fd = self.tcp.pipe().fd();
let tls = self
.tls
.get_mut()
.unwrap_or_else(|poisoned| poisoned.into_inner());
drain(tls, fd);
tls.send_close_notify();
let _ = send_pending(tls, fd);
}
}
fn drain(tls: &mut rustls::Connection, fd: libc::c_int) {
for _ in 0..DRAIN_READS {
match tls.read_tls(&mut FdIo(fd)) {
Ok(0) | Err(_) => return,
Ok(_) => {}
}
if tls.process_new_packets().is_err() {
return;
}
}
}
#[derive(Clone)]
pub struct TlsConnection {
stream: Arc<TlsStream>,
}
impl TlsConnection {
pub(crate) fn new(tcp: Connection, tls: rustls::Connection) -> Self {
Self {
stream: Arc::new(TlsStream {
tls: Mutex::new(tls),
tcp,
}),
}
}
#[inline(always)]
pub(crate) fn pipe(&self) -> &Pipe {
self.stream.tcp.pipe()
}
fn session(&self) -> MutexGuard<'_, rustls::Connection> {
self.stream
.tls
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}
#[inline(always)]
fn source(&self) -> Source {
Source::Tls(self.clone())
}
pub(crate) fn read(&self, into: &mut Vec<u8>, room: usize) -> Result<Io, RuntimeError> {
let fd = self.pipe().fd();
let mut tls = self.session();
let mut ended = false;
loop {
let start = into.len();
into.resize(start + room, 0);
let read = tls.reader().read(&mut into[start..]);
match read {
Ok(0) => {
into.truncate(start);
return Ok(Io::Closed);
}
Ok(read) => {
into.truncate(start + read);
return Ok(Io::Moved(read));
}
Err(error) => {
into.truncate(start);
match error.kind() {
io::ErrorKind::WouldBlock => {}
io::ErrorKind::UnexpectedEof => return Ok(Io::Truncated),
_ => return Err(io_error(error)),
}
}
}
if ended {
return Ok(Io::Truncated);
}
if let Some(filter) = send_pending(&mut tls, fd)? {
return Ok(Io::Wait(filter));
}
match tls.read_tls(&mut FdIo(fd)) {
Ok(0) => ended = true,
Ok(_) => {}
Err(error) if error.kind() == io::ErrorKind::WouldBlock => {
return Ok(Io::Wait(libc::EVFILT_READ));
}
Err(error) if error.kind() == io::ErrorKind::Interrupted => continue,
Err(error) => return Err(io_error(error)),
}
process(&mut tls, fd)?;
}
}
pub(crate) fn write(&self, data: &[u8]) -> Result<Io, RuntimeError> {
let fd = self.pipe().fd();
let mut tls = self.session();
send_pending(&mut tls, fd)?;
let put = tls.writer().write(data).map_err(io_error)?;
send_pending(&mut tls, fd)?;
match put {
0 => Ok(Io::Wait(libc::EVFILT_WRITE)),
put => Ok(Io::Moved(put)),
}
}
pub(crate) fn say_goodbye(&self) {
self.session().send_close_notify();
}
pub(crate) fn flush(&self) -> Result<Io, RuntimeError> {
let fd = self.pipe().fd();
let mut tls = self.session();
match send_pending(&mut tls, fd)? {
Some(filter) => Ok(Io::Wait(filter)),
None => Ok(Io::Moved(0)),
}
}
pub fn send(&self, data: impl Into<Arc<[u8]>>) -> SendTask {
SendTask::new(self.source(), data.into())
}
pub fn recv(&self, max: usize) -> RecvTask {
RecvTask::some(self.source(), max)
}
pub fn recv_exact(&self, len: usize) -> RecvTask {
RecvTask::exact(self.source(), len)
}
pub fn recv_until(&self, delimiter: &[u8], max: usize) -> RecvTask {
RecvTask::until(self.source(), Arc::from(delimiter), max)
}
pub fn recv_to_end(&self) -> RecvTask {
RecvTask::to_end(self.source())
}
#[inline(always)]
pub fn local_addr(&self) -> SocketAddr {
self.stream.tcp.local_addr()
}
#[inline(always)]
pub fn peer_addr(&self) -> SocketAddr {
self.stream.tcp.peer_addr()
}
pub fn set_nodelay(&self, nodelay: bool) -> Result<(), RuntimeError> {
self.stream.tcp.set_nodelay(nodelay)
}
pub fn nodelay(&self) -> Result<bool, RuntimeError> {
self.stream.tcp.nodelay()
}
pub fn set_keepalive(&self, idle: Option<Duration>) -> Result<(), RuntimeError> {
self.stream.tcp.set_keepalive(idle)
}
pub fn alpn(&self) -> Option<Vec<u8>> {
self.session().alpn_protocol().map(<[u8]>::to_vec)
}
pub fn peer_certificates(&self) -> Vec<Vec<u8>> {
self.session()
.peer_certificates()
.map(|chain| chain.iter().map(|cert| cert.as_ref().to_vec()).collect())
.unwrap_or_default()
}
pub fn finish(&self) -> FinishTask {
FinishTask::new(self.source())
}
pub fn close(self) {
drop(self);
}
}
impl fmt::Debug for TlsConnection {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("TlsConnection")
.field("local", &self.local_addr())
.field("peer", &self.peer_addr())
.finish()
}
}
#[derive(Clone)]
pub struct TlsListener {
tcp: Listener,
config: Arc<ServerConfig>,
}
impl TlsListener {
pub(crate) fn new(tcp: Listener, config: Arc<ServerConfig>) -> Self {
Self { tcp, config }
}
#[inline(always)]
pub(crate) fn tcp(&self) -> &Listener {
&self.tcp
}
#[inline(always)]
pub(crate) fn config(&self) -> Arc<ServerConfig> {
self.config.clone()
}
pub fn accept(&self) -> TlsAcceptTask {
TlsAcceptTask::new(self.clone())
}
pub fn upgrade(&self, conn: Connection) -> TlsAcceptTask {
TlsAcceptTask::over(self.clone(), conn)
}
#[inline(always)]
pub fn local_addr(&self) -> SocketAddr {
self.tcp.local_addr()
}
pub fn close(self) {
drop(self);
}
}
impl fmt::Debug for TlsListener {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("TlsListener")
.field("local", &self.local_addr())
.finish()
}
}
impl Sends for TlsConnection {
fn send_all(&self, data: Arc<[u8]>) -> SendTask {
self.send(data)
}
}