use crate::modules::input::{Token, token};
use crate::{
RuntimeError,
constants::INLINE_PAYLOAD,
futures::{
net::{
address::{Target, sealed::Sealed as _},
exchange::{self, Stage},
socket::Options,
step::{Progress, settle, wait_on},
},
task::{
Nothing, Task,
sealed::{self, Step},
},
tcp::{
Connection,
tcp_task::{AcceptTask, ConnectTask, ListenTask},
},
tls::{
config::{self, ClientSettings, Keys, ServerSettings, tls_error},
connection::{TlsConnection, TlsListener},
handshake::handshake,
},
},
modules::park,
};
use rustls::pki_types::ServerName;
use std::{
mem,
net::{IpAddr, SocketAddr},
sync::Arc,
time::Duration,
};
const _: () = assert!(mem::size_of::<Result<TlsConnection, RuntimeError>>() <= INLINE_PAYLOAD);
const _: () = assert!(mem::size_of::<Result<TlsListener, RuntimeError>>() <= INLINE_PAYLOAD);
const _: () =
assert!(mem::size_of::<Result<(TlsConnection, SocketAddr), RuntimeError>>() <= INLINE_PAYLOAD);
fn server_name(target: &Target, given: Option<&str>) -> Result<ServerName<'static>, RuntimeError> {
let host = match (given, target) {
(Some(name), _) => name,
(None, Target::Addr(addr)) => return Ok(ServerName::IpAddress(addr.ip().into())),
(None, Target::Name(name)) => host_of(name).ok_or(RuntimeError::BadAddress)?,
};
if let Ok(ip) = host.parse::<IpAddr>() {
return Ok(ServerName::IpAddress(ip.into()));
}
ServerName::try_from(host.to_owned()).map_err(|_| RuntimeError::BadAddress)
}
fn host_of(text: &str) -> Option<&str> {
let (host, _) = text.rsplit_once(':')?;
let host = host
.strip_prefix('[')
.and_then(|inner| inner.strip_suffix(']'))
.unwrap_or(host);
(!host.is_empty()).then_some(host)
}
#[derive(Debug, Clone)]
#[must_use = "a task does nothing until it is run or spawned"]
pub struct TlsConnectTask {
target: Target,
server_name: Option<Arc<str>>,
settings: ClientSettings,
over: Option<Connection>,
connect: ConnectTask,
options: Options,
stage: Progress<Connecting>,
}
#[derive(Default)]
enum Connecting {
#[default]
Tcp,
Handshaking(Connection, rustls::Connection),
}
impl TlsConnectTask {
pub(crate) fn new(target: Target) -> Self {
Self {
connect: ConnectTask::new(target.clone()),
target,
server_name: None,
settings: ClientSettings::default(),
over: None,
options: Options::default(),
stage: Progress::default(),
}
}
pub fn nodelay(mut self, nodelay: bool) -> Self {
self.options.nodelay = nodelay;
self
}
pub fn keepalive(mut self, idle: Duration) -> Self {
self.options.keepalive = Some(idle);
self
}
pub fn server_name(mut self, name: &str) -> Self {
self.server_name = Some(Arc::from(name));
self
}
pub fn trust(mut self, pem: impl AsRef<[u8]>) -> Self {
self.settings.roots = Some(Arc::from(pem.as_ref()));
self
}
pub fn alpn<I, P>(mut self, protocols: I) -> Self
where
I: IntoIterator<Item = P>,
P: AsRef<[u8]>,
{
self.settings.alpn = protocols
.into_iter()
.map(|protocol| protocol.as_ref().to_vec())
.collect();
self
}
pub fn identity(mut self, cert: impl AsRef<[u8]>, key: impl AsRef<[u8]>) -> Self {
self.settings.identity = Some((Arc::from(cert.as_ref()), Arc::from(key.as_ref())));
self
}
pub(crate) fn over(conn: Connection) -> Self {
let mut task = Self::new(conn.peer_addr().target());
task.over = Some(conn);
task
}
fn begin(&mut self) {
self.connect = ConnectTask::new(self.target.clone()).with_options(self.options);
self.stage = Progress::default();
}
fn session(&self) -> Result<rustls::Connection, RuntimeError> {
let config = config::client_with(&self.settings)?;
let name = server_name(&self.target, self.server_name.as_deref())?;
let session = rustls::ClientConnection::new(config, name).map_err(tls_error)?;
Ok(session.into())
}
fn advance(
&mut self,
reactor_id: i32,
task_id: usize,
) -> Result<Step<Result<TlsConnection, RuntimeError>>, RuntimeError> {
loop {
match mem::take(&mut self.stage.0) {
Connecting::Tcp => match self.opened(reactor_id, task_id) {
Step::Done(Ok(tcp)) => {
let session = self.session()?;
self.stage.0 = Connecting::Handshaking(tcp, session);
}
Step::Done(Err(error)) => return Err(error),
Step::Park(park) => return Ok(Step::Park(park)),
},
Connecting::Handshaking(tcp, mut session) => {
let fd = tcp.pipe().fd();
match handshake(&mut session, fd)? {
Some(filter) => {
let step = wait_on(fd, filter)?;
self.stage.0 = Connecting::Handshaking(tcp, session);
return Ok(step);
}
None => return Ok(Step::Done(Ok(TlsConnection::new(tcp, session)))),
}
}
}
}
}
}
impl TlsConnectTask {
fn opened(&mut self, reactor_id: i32, task_id: usize) -> Step<Result<Connection, RuntimeError>> {
match &self.over {
Some(conn) => Step::Done(Ok(conn.clone())),
None => self.connect.step(token(), reactor_id, task_id),
}
}
}
#[derive(Debug, Clone)]
#[must_use = "a task does nothing until it is run or spawned"]
pub struct TlsListenTask {
listen: ListenTask,
settings: ServerSettings,
}
impl TlsListenTask {
pub(crate) fn new(target: Target, keys: Keys) -> Self {
Self {
listen: ListenTask::new(target),
settings: ServerSettings {
keys,
alpn: Arc::from([]),
client_roots: None,
},
}
}
pub fn alpn<I, P>(mut self, protocols: I) -> Self
where
I: IntoIterator<Item = P>,
P: AsRef<[u8]>,
{
self.settings.alpn = protocols
.into_iter()
.map(|protocol| protocol.as_ref().to_vec())
.collect();
self
}
pub fn require_client_cert(mut self, pem: impl AsRef<[u8]>) -> Self {
self.settings.client_roots = Some(Arc::from(pem.as_ref()));
self
}
pub fn backlog(mut self, backlog: u32) -> Self {
let options = Options {
backlog: Some(backlog),
..self.listen.options()
};
self.listen = self.listen.with_options(options);
self
}
pub fn reuse_port(mut self, reuse: bool) -> Self {
let options = Options {
reuse_port: reuse,
..self.listen.options()
};
self.listen = self.listen.with_options(options);
self
}
pub fn v6_only(mut self, only: bool) -> Self {
let options = Options {
v6_only: only,
..self.listen.options()
};
self.listen = self.listen.with_options(options);
self
}
fn listen(&self, reactor_id: i32, task_id: usize) -> Result<TlsListener, RuntimeError> {
let config = config::server_with(&self.settings)?;
let tcp = self.listen.execute(token(), reactor_id, task_id)?;
Ok(TlsListener::new(tcp, config))
}
}
#[derive(Debug, Clone)]
#[must_use = "a task does nothing until it is run or spawned"]
pub struct TlsAcceptTask {
listener: TlsListener,
accept: AcceptTask,
over: Option<Connection>,
stage: Progress<Accepting>,
}
#[derive(Default)]
enum Accepting {
#[default]
Tcp,
Handshaking(Connection, SocketAddr, rustls::Connection),
}
impl TlsAcceptTask {
pub(crate) fn new(listener: TlsListener) -> Self {
Self {
accept: AcceptTask::new(listener.tcp().clone()),
listener,
over: None,
stage: Progress::default(),
}
}
pub(crate) fn over(listener: TlsListener, conn: Connection) -> Self {
let mut task = Self::new(listener);
task.over = Some(conn);
task
}
fn opened(
&mut self,
reactor_id: i32,
task_id: usize,
) -> Step<Result<(Connection, SocketAddr), RuntimeError>> {
match &self.over {
Some(conn) => Step::Done(Ok((conn.clone(), conn.peer_addr()))),
None => self.accept.step(token(), reactor_id, task_id),
}
}
fn advance(
&mut self,
reactor_id: i32,
task_id: usize,
) -> Result<Step<Result<(TlsConnection, SocketAddr), RuntimeError>>, RuntimeError> {
loop {
match mem::take(&mut self.stage.0) {
Accepting::Tcp => match self.opened(reactor_id, task_id) {
Step::Done(Ok((tcp, peer))) => {
let session = rustls::ServerConnection::new(self.listener.config())
.map_err(tls_error)?;
self.stage.0 = Accepting::Handshaking(tcp, peer, session.into());
}
Step::Done(Err(error)) => return Err(error),
Step::Park(park) => return Ok(Step::Park(park)),
},
Accepting::Handshaking(tcp, peer, mut session) => {
let fd = tcp.pipe().fd();
match handshake(&mut session, fd)? {
Some(filter) => {
let step = wait_on(fd, filter)?;
self.stage.0 = Accepting::Handshaking(tcp, peer, session);
return Ok(step);
}
None => {
return Ok(Step::Done(Ok((TlsConnection::new(tcp, session), peer))));
}
}
}
}
}
}
}
#[derive(Debug, Clone)]
#[must_use = "a task does nothing until it is run or spawned"]
pub struct TlsRequestTask {
connect: TlsConnectTask,
data: Arc<[u8]>,
stage: Progress<Stage>,
}
impl TlsRequestTask {
pub(crate) fn new(target: Target, data: Arc<[u8]>) -> Self {
Self {
connect: TlsConnectTask::new(target),
data,
stage: Progress::default(),
}
}
pub fn server_name(mut self, name: &str) -> Self {
self.connect = self.connect.server_name(name);
self
}
pub fn trust(mut self, pem: impl AsRef<[u8]>) -> Self {
self.connect = self.connect.trust(pem);
self
}
pub fn alpn<I, P>(mut self, protocols: I) -> Self
where
I: IntoIterator<Item = P>,
P: AsRef<[u8]>,
{
self.connect = self.connect.alpn(protocols);
self
}
pub fn identity(mut self, cert: impl AsRef<[u8]>, key: impl AsRef<[u8]>) -> Self {
self.connect = self.connect.identity(cert, key);
self
}
fn advance(&mut self, reactor_id: i32, task_id: usize) -> Step<Result<Vec<u8>, RuntimeError>> {
exchange::advance(
&mut self.connect,
&mut self.stage.0,
&self.data,
reactor_id,
task_id,
)
}
}
impl sealed::Sealed for TlsConnectTask {}
impl sealed::Sealed for TlsListenTask {}
impl sealed::Sealed for TlsAcceptTask {}
impl sealed::Sealed for TlsRequestTask {}
impl Task for TlsConnectTask {
type Output = Result<TlsConnection, RuntimeError>;
type Input = Nothing;
fn execute(&self, _token: Token, reactor_id: i32, task_id: usize) -> Self::Output {
park::drive(self.clone(), reactor_id, task_id)
}
fn prepare(&mut self, _token: Token) {
self.begin();
}
fn blocking(&self, _token: Token) -> bool {
true
}
fn step(&mut self, _token: Token, reactor_id: i32, task_id: usize) -> Step<Self::Output> {
settle(self.advance(reactor_id, task_id))
}
}
impl Task for TlsListenTask {
type Output = Result<TlsListener, RuntimeError>;
type Input = Nothing;
fn execute(&self, _token: Token, reactor_id: i32, task_id: usize) -> Self::Output {
self.listen(reactor_id, task_id)
}
fn blocking(&self, _token: Token) -> bool {
true
}
}
impl Task for TlsAcceptTask {
type Output = Result<(TlsConnection, SocketAddr), RuntimeError>;
type Input = Nothing;
fn execute(&self, _token: Token, reactor_id: i32, task_id: usize) -> Self::Output {
park::drive(self.clone(), reactor_id, task_id)
}
fn prepare(&mut self, _token: Token) {
self.accept = AcceptTask::new(self.listener.tcp().clone());
self.stage = Progress::default();
}
fn step(&mut self, _token: Token, reactor_id: i32, task_id: usize) -> Step<Self::Output> {
settle(self.advance(reactor_id, task_id))
}
}
impl Task for TlsRequestTask {
type Output = Result<Vec<u8>, RuntimeError>;
type Input = Nothing;
fn execute(&self, _token: Token, reactor_id: i32, task_id: usize) -> Self::Output {
park::drive(self.clone(), reactor_id, task_id)
}
fn prepare(&mut self, _token: Token) {
self.connect.begin();
self.stage = Progress::default();
}
fn blocking(&self, _token: Token) -> bool {
self.connect.blocking(token())
}
fn step(&mut self, _token: Token, reactor_id: i32, task_id: usize) -> Step<Self::Output> {
self.advance(reactor_id, task_id)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::futures::net::address::sealed::Sealed;
#[test]
fn the_server_name_comes_from_where_it_connects() {
let named = server_name(&"example.com:443".target(), None).unwrap();
assert_eq!(named, ServerName::try_from("example.com").unwrap());
let v4 = server_name(&"127.0.0.1:443".target(), None).unwrap();
assert_eq!(
v4,
ServerName::IpAddress("127.0.0.1".parse::<IpAddr>().unwrap().into())
);
let v6 = server_name(&"[::1]:443".target(), None).unwrap();
assert_eq!(
v6,
ServerName::IpAddress("::1".parse::<IpAddr>().unwrap().into())
);
let given = server_name(&"127.0.0.1:443".target(), Some("localhost")).unwrap();
assert_eq!(given, ServerName::try_from("localhost").unwrap());
}
#[test]
fn no_usable_name_is_a_bad_address() {
assert_eq!(
server_name(&"no port here".target(), None),
Err(RuntimeError::BadAddress),
);
assert_eq!(
server_name(&"example.com:443".target(), Some("not a name!")),
Err(RuntimeError::BadAddress),
);
}
#[test]
fn the_host_is_picked_out() {
assert_eq!(host_of("example.com:443"), Some("example.com"));
assert_eq!(host_of("[fe80::1]:443"), Some("fe80::1"));
assert_eq!(host_of(":443"), None);
assert_eq!(host_of("no port"), None);
}
}