use crate::courierust_error::{Error, ErrorKind, Result};
use crate::courierust_io::{Read, Write};
use std::net::{SocketAddr, TcpListener, TcpStream};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::Duration;
pub(crate) mod poller;
pub mod stats;
pub(crate) mod udp;
impl Read for &TcpStream {
fn read(&mut self, buf: &mut [u8]) -> Result<usize> {
match std::io::Read::read(self, buf) {
Ok(n) => Ok(n),
Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => {
Err(Error::new(ErrorKind::WouldBlock))
}
Err(e) if e.kind() == std::io::ErrorKind::TimedOut => {
Err(Error::new(ErrorKind::Timeout))
}
Err(e) if e.raw_os_error() == Some(997) => Err(Error::new(ErrorKind::WouldBlock)),
Err(e) => Err(e.into()),
}
}
}
impl Write for &TcpStream {
fn write(&mut self, buf: &[u8]) -> Result<usize> {
match std::io::Write::write(self, buf) {
Ok(n) => Ok(n),
Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => {
Err(Error::new(ErrorKind::WouldBlock))
}
Err(e) => Err(e.into()),
}
}
fn flush(&mut self) -> Result<()> {
match std::io::Write::flush(self) {
Ok(()) => Ok(()),
Err(e) => Err(e.into()),
}
}
}
impl Read for Arc<TcpStream> {
fn read(&mut self, buf: &mut [u8]) -> Result<usize> {
let mut r: &TcpStream = self;
r.read(buf)
}
}
impl Write for Arc<TcpStream> {
fn write(&mut self, buf: &[u8]) -> Result<usize> {
let mut w: &TcpStream = self;
w.write(buf)
}
fn flush(&mut self) -> Result<()> {
let mut w: &TcpStream = self;
w.flush()
}
}
pub type Listener = TcpListener;
pub(crate) struct ConnStream {
peer: SocketAddr,
deadline: AtomicBool,
inner: ConnStreamKind,
}
enum ConnStreamKind {
Plain(TcpStream),
Tls {
socket: Arc<TcpStream>,
tls:
Box<std::sync::Mutex<crate::courierust_tls::TlsStream<Arc<TcpStream>, Arc<TcpStream>>>>,
},
}
impl ConnStream {
pub(crate) fn plain(stream: TcpStream) -> Self {
let peer = stream
.peer_addr()
.unwrap_or_else(|_| SocketAddr::from(([0, 0, 0, 0], 0)));
Self {
peer,
deadline: AtomicBool::new(false),
inner: ConnStreamKind::Plain(stream),
}
}
pub(crate) fn tls_client(
stream: TcpStream,
connector: &crate::courierust_tls::TlsConnector,
hostname: &str,
) -> crate::Result<Self> {
let peer = stream.peer_addr().map_err(|e| Error::io(e.to_string()))?;
let socket = Arc::new(stream);
let tls = connector
.connect(hostname, socket.clone(), socket.clone())
.map_err(|e| Error::io(e.to_string()))?;
Ok(Self {
peer,
deadline: AtomicBool::new(false),
inner: ConnStreamKind::Tls {
socket,
tls: Box::new(std::sync::Mutex::new(tls)),
},
})
}
pub(crate) fn tls_server(
tls: crate::courierust_tls::TlsStream<Arc<TcpStream>, Arc<TcpStream>>,
peer: SocketAddr,
) -> Self {
let socket = tls.underlying().clone();
Self {
peer,
deadline: AtomicBool::new(false),
inner: ConnStreamKind::Tls {
socket,
tls: Box::new(std::sync::Mutex::new(tls)),
},
}
}
pub(crate) fn peer_addr(&self) -> SocketAddr {
self.peer
}
pub(crate) fn alpn(&self) -> Option<Vec<u8>> {
match &self.inner {
ConnStreamKind::Tls { tls, .. } => {
tls.lock().ok().and_then(|g| g.alpn().map(|a| a.to_vec()))
}
ConnStreamKind::Plain(_) => None,
}
}
pub(crate) fn peek(&self, buf: &mut [u8]) -> std::io::Result<usize> {
match &self.inner {
ConnStreamKind::Plain(s) => s.peek(buf),
ConnStreamKind::Tls { .. } => Err(std::io::Error::other(
"peek is not supported on TLS streams",
)),
}
}
pub(crate) fn raw_fd(&self) -> crate::courierust_net::poller::Fd {
match &self.inner {
ConnStreamKind::Plain(s) => crate::courierust_net::poller::fd_of(s),
ConnStreamKind::Tls { socket, .. } => crate::courierust_net::poller::fd_of(socket),
}
}
pub(crate) fn configure(&self, read_timeout: Option<Duration>) -> Result<()> {
self.deadline.store(false, Ordering::Relaxed);
match &self.inner {
ConnStreamKind::Plain(s) => configure(s, read_timeout),
ConnStreamKind::Tls { socket, .. } => configure(socket, read_timeout),
}
}
pub(crate) fn set_deadline(&self, read_timeout: Option<Duration>) -> Result<()> {
self.configure(read_timeout)?;
self.deadline.store(true, Ordering::Relaxed);
Ok(())
}
pub(crate) fn linger_close(&self, budget: usize, deadline: Duration) {
let _ = self.configure(Some(deadline));
let mut sink = [0u8; 8 * 1024];
let mut left = budget;
while left > 0 {
let mut reader = self;
let want = core::cmp::min(left, sink.len());
match crate::courierust_io::Read::read(&mut reader, &mut sink[..want]) {
Ok(0) => break,
Ok(n) => left = left.saturating_sub(n),
Err(_) => break,
}
}
}
}
impl crate::courierust_io::Read for &ConnStream {
fn read(&mut self, buf: &mut [u8]) -> Result<usize> {
let read = match &self.inner {
ConnStreamKind::Plain(s) => {
let mut r: &TcpStream = s;
crate::courierust_io::Read::read(&mut r, buf)
}
ConnStreamKind::Tls { tls, .. } => {
let mut g = tls.lock().unwrap_or_else(|e| e.into_inner());
crate::courierust_io::Read::read(&mut *g, buf)
}
};
read.map_err(|e| {
if e.kind == ErrorKind::WouldBlock && self.deadline.load(Ordering::Relaxed) {
Error::timeout("read deadline expired")
} else {
e
}
})
}
}
impl crate::courierust_io::Write for &ConnStream {
fn write(&mut self, buf: &[u8]) -> Result<usize> {
match &self.inner {
ConnStreamKind::Plain(s) => {
let mut w: &TcpStream = s;
crate::courierust_io::Write::write(&mut w, buf)
}
ConnStreamKind::Tls { tls, .. } => {
let mut g = tls.lock().unwrap_or_else(|e| e.into_inner());
crate::courierust_io::Write::write(&mut *g, buf)
}
}
}
fn flush(&mut self) -> Result<()> {
match &self.inner {
ConnStreamKind::Plain(s) => {
let mut w: &TcpStream = s;
crate::courierust_io::Write::flush(&mut w)
}
ConnStreamKind::Tls { tls, .. } => {
let mut g = tls.lock().unwrap_or_else(|e| e.into_inner());
crate::courierust_io::Write::flush(&mut *g)
}
}
}
}
impl crate::courierust_io::Read for ConnStream {
fn read(&mut self, buf: &mut [u8]) -> Result<usize> {
(&*self).read(buf)
}
}
impl crate::courierust_io::Write for ConnStream {
fn write(&mut self, buf: &[u8]) -> Result<usize> {
(&*self).write(buf)
}
fn flush(&mut self) -> Result<()> {
(&*self).flush()
}
}
impl crate::courierust_io::Read for Arc<ConnStream> {
fn read(&mut self, buf: &mut [u8]) -> Result<usize> {
let mut r: &ConnStream = self;
r.read(buf)
}
}
impl crate::courierust_io::Write for Arc<ConnStream> {
fn write(&mut self, buf: &[u8]) -> Result<usize> {
let mut w: &ConnStream = self;
w.write(buf)
}
fn flush(&mut self) -> Result<()> {
let mut w: &ConnStream = self;
w.flush()
}
}
pub fn configure(stream: &TcpStream, read_timeout: Option<Duration>) -> Result<()> {
stream
.set_nodelay(true)
.map_err(|e| Error::io(e.to_string()))?;
stream
.set_read_timeout(read_timeout)
.map_err(|e| Error::io(e.to_string()))?;
stream
.set_write_timeout(read_timeout)
.map_err(|e| Error::io(e.to_string()))?;
Ok(())
}
pub fn connect(addr: &std::net::SocketAddr, timeout: Option<Duration>) -> Result<TcpStream> {
let stream = match timeout {
Some(t) => TcpStream::connect_timeout(addr, t),
None => TcpStream::connect(addr),
}
.map_err(|e| Error::io(e.to_string()))?;
stream
.set_nodelay(true)
.map_err(|e| Error::io(e.to_string()))?;
Ok(stream)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn deadline_read_timeout_is_reported_as_timeout() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let client = TcpStream::connect(addr).unwrap();
let (_peer, _) = listener.accept().unwrap();
let stream = ConnStream::plain(client);
stream
.set_deadline(Some(Duration::from_millis(50)))
.unwrap();
let mut sink = [0u8; 16];
let started = std::time::Instant::now();
let mut reader: &ConnStream = &stream;
let err = reader
.read(&mut sink)
.expect_err("nothing was ever sent on the connection");
assert_eq!(err.kind, ErrorKind::Timeout, "{err:?}");
assert!(
started.elapsed() >= Duration::from_millis(40),
"the read must wait out the deadline instead of failing at once"
);
}
}