1use std::io;
5use std::pin::Pin;
6use std::task::{Context, Poll};
7
8use moirai_async::io::{AsyncRead, AsyncWrite};
9use moirai_async::net::TcpStream;
10use moirai_tls::{ServerName, TlsConnector, TlsStream};
11
12#[derive(Debug, Clone, PartialEq, Eq, Hash)]
14pub struct Origin {
15 pub secure: bool,
17 pub host: String,
19 pub port: u16,
21}
22
23impl Origin {
24 #[must_use]
26 pub fn authority(&self) -> String {
27 format!("{}:{}", self.host, self.port)
28 }
29
30 #[must_use]
33 pub fn host_header(&self) -> String {
34 let default = if self.secure { 443 } else { 80 };
35 if self.port == default {
36 self.host.clone()
37 } else {
38 format!("{}:{}", self.host, self.port)
39 }
40 }
41}
42
43pub enum Conn {
45 Plain(TcpStream),
47 Tls(Box<TlsStream<TcpStream>>),
49}
50
51impl Conn {
52 pub async fn connect(origin: &Origin, tls: &TlsConnector) -> io::Result<Self> {
57 let tcp = TcpStream::connect(&origin.authority()).await?;
58 if origin.secure {
59 let domain = ServerName::try_from(origin.host.clone())
60 .map_err(|e| io::Error::new(io::ErrorKind::InvalidInput, e))?;
61 let stream = tls.connect(domain, tcp).await?;
62 Ok(Conn::Tls(Box::new(stream)))
63 } else {
64 Ok(Conn::Plain(tcp))
65 }
66 }
67}
68
69impl AsyncRead for Conn {
70 #[inline]
71 fn poll_read(
72 self: Pin<&mut Self>,
73 cx: &mut Context<'_>,
74 buf: &mut [u8],
75 ) -> Poll<io::Result<usize>> {
76 match self.get_mut() {
77 Conn::Plain(s) => AsyncRead::poll_read(Pin::new(s), cx, buf),
78 Conn::Tls(s) => AsyncRead::poll_read(Pin::new(s.as_mut()), cx, buf),
79 }
80 }
81}
82
83impl AsyncWrite for Conn {
84 #[inline]
85 fn poll_write(
86 self: Pin<&mut Self>,
87 cx: &mut Context<'_>,
88 buf: &[u8],
89 ) -> Poll<io::Result<usize>> {
90 match self.get_mut() {
91 Conn::Plain(s) => AsyncWrite::poll_write(Pin::new(s), cx, buf),
92 Conn::Tls(s) => AsyncWrite::poll_write(Pin::new(s.as_mut()), cx, buf),
93 }
94 }
95
96 #[inline]
97 fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
98 match self.get_mut() {
99 Conn::Plain(s) => AsyncWrite::poll_flush(Pin::new(s), cx),
100 Conn::Tls(s) => AsyncWrite::poll_flush(Pin::new(s.as_mut()), cx),
101 }
102 }
103
104 #[inline]
105 fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
106 match self.get_mut() {
107 Conn::Plain(s) => AsyncWrite::poll_shutdown(Pin::new(s), cx),
108 Conn::Tls(s) => AsyncWrite::poll_shutdown(Pin::new(s.as_mut()), cx),
109 }
110 }
111}