pub use std::net::Shutdown;
use std::os::unix::prelude::AsRawFd;
use crate::*;
pub struct Stream {
inner: StreamType,
}
impl AsRawFd for Stream {
fn as_raw_fd(&self) -> i32 {
match &self.inner {
StreamType::Tcp(s) => s.as_raw_fd(),
#[cfg(any(feature = "boringssl", feature = "openssl"))]
StreamType::TlsTcp(s) => s.as_raw_fd(),
}
}
}
impl Stream {
pub fn interest(&mut self) -> Interest {
match &mut self.inner {
StreamType::Tcp(s) => {
if !s.is_established() {
Interest::READABLE.add(Interest::WRITABLE)
} else {
Interest::READABLE
}
}
#[cfg(any(feature = "boringssl", feature = "openssl"))]
StreamType::TlsTcp(s) => s.interest(),
}
}
pub fn is_established(&mut self) -> bool {
match &mut self.inner {
StreamType::Tcp(s) => s.is_established(),
#[cfg(any(feature = "boringssl", feature = "openssl"))]
StreamType::TlsTcp(s) => !s.is_handshaking(),
}
}
pub fn is_handshaking(&self) -> bool {
match &self.inner {
StreamType::Tcp(_) => false,
#[cfg(any(feature = "boringssl", feature = "openssl"))]
StreamType::TlsTcp(s) => s.is_handshaking(),
}
}
pub fn do_handshake(&mut self) -> Result<()> {
match &mut self.inner {
StreamType::Tcp(_) => Ok(()),
#[cfg(any(feature = "boringssl", feature = "openssl"))]
StreamType::TlsTcp(s) => s.do_handshake(),
}
}
pub fn set_nodelay(&mut self, nodelay: bool) -> Result<()> {
match &mut self.inner {
StreamType::Tcp(s) => s.set_nodelay(nodelay),
#[cfg(any(feature = "boringssl", feature = "openssl"))]
StreamType::TlsTcp(s) => s.set_nodelay(nodelay),
}
}
pub fn shutdown(&mut self) -> Result<bool> {
let result = match &mut self.inner {
StreamType::Tcp(s) => s.shutdown(Shutdown::Both).map(|_| true),
#[cfg(any(feature = "boringssl", feature = "openssl"))]
StreamType::TlsTcp(s) => s.shutdown().map(|v| v == ShutdownResult::Received),
};
STREAM_SHUTDOWN.increment();
if result.is_err() {
STREAM_SHUTDOWN_EX.increment();
}
result
}
}
impl Drop for Stream {
fn drop(&mut self) {
STREAM_CLOSE.increment();
}
}
impl Debug for Stream {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::result::Result<(), std::fmt::Error> {
match &self.inner {
StreamType::Tcp(s) => write!(f, "{s:?}"),
#[cfg(any(feature = "boringssl", feature = "openssl"))]
StreamType::TlsTcp(s) => write!(f, "{s:?}"),
}
}
}
impl From<TcpStream> for Stream {
fn from(other: TcpStream) -> Self {
Self {
inner: StreamType::Tcp(other),
}
}
}
#[cfg(any(feature = "boringssl", feature = "openssl"))]
impl From<TlsTcpStream> for Stream {
fn from(other: TlsTcpStream) -> Self {
Self {
inner: StreamType::TlsTcp(other),
}
}
}
impl Read for Stream {
fn read(&mut self, buf: &mut [u8]) -> Result<usize> {
match &mut self.inner {
StreamType::Tcp(s) => s.read(buf),
#[cfg(any(feature = "boringssl", feature = "openssl"))]
StreamType::TlsTcp(s) => s.read(buf),
}
}
}
impl Write for Stream {
fn write(&mut self, buf: &[u8]) -> Result<usize> {
match &mut self.inner {
StreamType::Tcp(s) => s.write(buf),
#[cfg(any(feature = "boringssl", feature = "openssl"))]
StreamType::TlsTcp(s) => s.write(buf),
}
}
fn flush(&mut self) -> Result<()> {
match &mut self.inner {
StreamType::Tcp(s) => s.flush(),
#[cfg(any(feature = "boringssl", feature = "openssl"))]
StreamType::TlsTcp(s) => s.flush(),
}
}
}
impl event::Source for Stream {
fn register(&mut self, registry: &Registry, token: Token, interest: Interest) -> Result<()> {
match &mut self.inner {
StreamType::Tcp(s) => s.register(registry, token, interest),
#[cfg(any(feature = "boringssl", feature = "openssl"))]
StreamType::TlsTcp(s) => s.register(registry, token, interest),
}
}
fn reregister(
&mut self,
registry: &mio::Registry,
token: mio::Token,
interest: mio::Interest,
) -> Result<()> {
match &mut self.inner {
StreamType::Tcp(s) => s.reregister(registry, token, interest),
#[cfg(any(feature = "boringssl", feature = "openssl"))]
StreamType::TlsTcp(s) => s.reregister(registry, token, interest),
}
}
fn deregister(&mut self, registry: &mio::Registry) -> Result<()> {
match &mut self.inner {
StreamType::Tcp(s) => s.deregister(registry),
#[cfg(any(feature = "boringssl", feature = "openssl"))]
StreamType::TlsTcp(s) => s.deregister(registry),
}
}
}
enum StreamType {
Tcp(TcpStream),
#[cfg(any(feature = "boringssl", feature = "openssl"))]
TlsTcp(TlsTcpStream),
}