use std::io::{self, ErrorKind, Read, Write};
use std::mem;
use std::net::TcpStream;
use std::sync::{Arc, Mutex, MutexGuard};
use std::time::{Duration, Instant};
use rustls::Connection;
const RX_CHUNK: usize = 32 * 1024;
const TX_CHUNK: usize = 16 * 1024;
pub struct TlsConn {
conn: Mutex<Connection>,
}
impl TlsConn {
pub fn new(conn: impl Into<Connection>) -> Arc<TlsConn> {
Arc::new(TlsConn {
conn: Mutex::new(conn.into()),
})
}
fn lock(&self) -> MutexGuard<'_, Connection> {
self.conn
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}
pub fn split(
self: &Arc<TlsConn>,
read_sock: TcpStream,
write_sock: TcpStream,
) -> (TlsReadHalf, TlsWriteHalf) {
(
TlsReadHalf {
conn: Arc::clone(self),
sock: read_sock,
rx: vec![0u8; RX_CHUNK],
rx_pos: 0,
rx_len: 0,
},
TlsWriteHalf {
conn: Arc::clone(self),
sock: write_sock,
tx: Vec::with_capacity(TX_CHUNK),
},
)
}
}
pub struct TlsReadHalf {
conn: Arc<TlsConn>,
sock: TcpStream,
rx: Vec<u8>,
rx_pos: usize,
rx_len: usize,
}
impl Read for TlsReadHalf {
fn read(&mut self, out: &mut [u8]) -> io::Result<usize> {
loop {
if let Some(n) = self.serve_plaintext(out)? {
return Ok(n);
}
if self.rx_pos < self.rx_len {
let mut conn = self.conn.lock();
let mut cursor: &[u8] = &self.rx[self.rx_pos..self.rx_len];
self.rx_pos += conn.read_tls(&mut cursor)?;
conn.process_new_packets()
.map_err(|err| io::Error::new(ErrorKind::InvalidData, err))?;
continue;
}
match self.sock.read(&mut self.rx) {
Ok(0) => {
return self
.serve_plaintext(out)
.map(|plaintext| plaintext.unwrap_or(0));
}
Ok(n) => {
self.rx_pos = 0;
self.rx_len = n;
}
Err(err) if err.kind() == ErrorKind::Interrupted => continue,
Err(err)
if err.kind() == ErrorKind::WouldBlock || err.kind() == ErrorKind::TimedOut =>
{
return Err(io::Error::from(ErrorKind::WouldBlock));
}
Err(err) => return Err(err),
}
}
}
}
impl TlsReadHalf {
fn serve_plaintext(&self, out: &mut [u8]) -> io::Result<Option<usize>> {
let mut conn = self.conn.lock();
match conn.reader().read(out) {
Ok(n) => Ok(Some(n)),
Err(err) if err.kind() == ErrorKind::WouldBlock => Ok(None),
Err(err) if err.kind() == ErrorKind::UnexpectedEof => Ok(Some(0)),
Err(err) => Err(err),
}
}
pub fn socket_handle(&self) -> io::Result<TcpStream> {
self.sock.try_clone()
}
}
pub struct TlsWriteHalf {
conn: Arc<TlsConn>,
sock: TcpStream,
tx: Vec<u8>,
}
impl Write for TlsWriteHalf {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
let mut offset = 0;
while offset < buf.len() {
let end = (offset + TX_CHUNK).min(buf.len());
{
let mut conn = self.conn.lock();
conn.writer().write_all(&buf[offset..end])?;
}
self.drain()?;
offset = end;
}
Ok(buf.len())
}
fn flush(&mut self) -> io::Result<()> {
self.drain()?;
self.sock.flush()
}
}
impl TlsWriteHalf {
fn drain(&mut self) -> io::Result<()> {
loop {
let mut cipher = mem::take(&mut self.tx);
cipher.clear();
let produced = {
let mut conn = self.conn.lock();
if !conn.wants_write() {
self.tx = cipher;
return Ok(());
}
conn.write_tls(&mut cipher)?
};
let write = if produced == 0 {
Ok(())
} else {
self.sock.write_all(&cipher)
};
self.tx = cipher;
write?;
if produced == 0 {
return Ok(());
}
}
}
}
pub fn drive_handshake(
conn: &mut Connection,
sock: &mut TcpStream,
deadline: Option<Duration>,
started: Instant,
) -> io::Result<()> {
while conn.is_handshaking() {
check_deadline(deadline, started)?;
if conn.wants_write() {
match conn.write_tls(sock) {
Ok(_) => {}
Err(err) if is_retryable(&err) => {}
Err(err) => return Err(err),
}
continue;
}
match conn.read_tls(sock) {
Ok(0) => {
return Err(io::Error::new(
ErrorKind::UnexpectedEof,
"peer closed during tls handshake",
));
}
Ok(_) => {
conn.process_new_packets()
.map_err(|err| io::Error::new(ErrorKind::InvalidData, err))?;
}
Err(err) if is_retryable(&err) => {}
Err(err) => return Err(err),
}
}
flush_final_flight(conn, sock)
}
fn flush_final_flight(conn: &mut Connection, sock: &mut TcpStream) -> io::Result<()> {
while conn.wants_write() {
match conn.write_tls(sock) {
Ok(0) => break,
Ok(_) => {}
Err(err) if is_retryable(&err) => break,
Err(err) => return Err(err),
}
}
Ok(())
}
fn check_deadline(deadline: Option<Duration>, started: Instant) -> io::Result<()> {
if let Some(deadline) = deadline {
if started.elapsed() >= deadline {
return Err(io::Error::new(
ErrorKind::TimedOut,
"tls handshake exceeded its deadline",
));
}
}
Ok(())
}
fn is_retryable(err: &io::Error) -> bool {
matches!(
err.kind(),
ErrorKind::WouldBlock | ErrorKind::TimedOut | ErrorKind::Interrupted
)
}