use std::net::SocketAddr;
use std::pin::Pin;
use std::task::{Context, Poll};
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf, ReadHalf, WriteHalf};
use tokio::net::TcpStream;
use tokio_rustls::client::TlsStream as ClientTlsStream;
use tokio_rustls::server::TlsStream as ServerTlsStream;
use hotaru_core::connection::{ConnMeta, ConnStream};
pub struct TlsMeta {
local: Option<SocketAddr>,
remote: Option<SocketAddr>,
}
impl ConnMeta for TlsMeta {
fn local_addr(&self) -> Option<SocketAddr> {
self.local
}
fn remote_addr(&self) -> Option<SocketAddr> {
self.remote
}
}
pub enum TlsStream {
Client(ClientTlsStream<TcpStream>),
Server(ServerTlsStream<TcpStream>),
}
impl AsyncRead for TlsStream {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
match self.get_mut() {
TlsStream::Client(s) => Pin::new(s).poll_read(cx, buf),
TlsStream::Server(s) => Pin::new(s).poll_read(cx, buf),
}
}
}
impl AsyncWrite for TlsStream {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<std::io::Result<usize>> {
match self.get_mut() {
TlsStream::Client(s) => Pin::new(s).poll_write(cx, buf),
TlsStream::Server(s) => Pin::new(s).poll_write(cx, buf),
}
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
match self.get_mut() {
TlsStream::Client(s) => Pin::new(s).poll_flush(cx),
TlsStream::Server(s) => Pin::new(s).poll_flush(cx),
}
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
match self.get_mut() {
TlsStream::Client(s) => Pin::new(s).poll_shutdown(cx),
TlsStream::Server(s) => Pin::new(s).poll_shutdown(cx),
}
}
}
impl ConnStream for TlsStream {
type ReadHalf = ReadHalf<TlsStream>;
type WriteHalf = WriteHalf<TlsStream>;
type Meta = TlsMeta;
fn split(self) -> (Self::ReadHalf, Self::WriteHalf, Self::Meta) {
let meta = TlsMeta {
local: self.local_addr().ok(),
remote: self.peer_addr().ok(),
};
let (r, w) = tokio::io::split(self);
(r, w, meta)
}
fn peer_addr(&self) -> std::io::Result<SocketAddr> {
match self {
TlsStream::Client(s) => s.get_ref().0.peer_addr(),
TlsStream::Server(s) => s.get_ref().0.peer_addr(),
}
}
fn local_addr(&self) -> std::io::Result<SocketAddr> {
match self {
TlsStream::Client(s) => s.get_ref().0.local_addr(),
TlsStream::Server(s) => s.get_ref().0.local_addr(),
}
}
}