use std::error::Error as StdError;
use std::fmt::Display;
use std::future::Future;
use std::net::{IpAddr, SocketAddr};
use std::time::Duration;
pub trait Spawn<F: Future<Output = ()>> {
fn spawn(&self, f: F);
}
#[derive(Debug, Clone, Default)]
pub struct TcpOpts {
pub nodelay: bool,
pub keepalive: Option<Duration>,
pub keepalive_interval: Option<Duration>,
pub keepalive_retries: Option<u32>,
pub bind_device: Option<String>,
pub user_timeout: Option<Duration>,
pub local_address: Option<IpAddr>,
pub send_buffer_size: Option<usize>,
pub recv_buffer_size: Option<usize>,
pub reuse_address: bool,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct TcpOptsSupport {
pub nodelay: bool,
pub keepalive: bool,
pub keepalive_interval: bool,
pub keepalive_retries: bool,
pub bind_device: bool,
pub user_timeout: bool,
pub local_address: bool,
pub send_buffer_size: bool,
pub recv_buffer_size: bool,
pub reuse_address: bool,
}
impl TcpOptsSupport {
pub const ALL: Self = Self {
nodelay: true,
keepalive: true,
keepalive_interval: true,
keepalive_retries: true,
bind_device: true,
user_timeout: true,
local_address: true,
send_buffer_size: true,
recv_buffer_size: true,
reuse_address: true,
};
pub const NONE: Self = Self {
nodelay: false,
keepalive: false,
keepalive_interval: false,
keepalive_retries: false,
bind_device: false,
user_timeout: false,
local_address: false,
send_buffer_size: false,
recv_buffer_size: false,
reuse_address: false,
};
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct UnsupportedTcpOpts {
missing: TcpOptsSupport,
}
impl UnsupportedTcpOpts {
pub fn names(&self) -> impl Iterator<Item = &'static str> {
let m = self.missing;
[
("nodelay", m.nodelay),
("keepalive", m.keepalive),
("keepalive_interval", m.keepalive_interval),
("keepalive_retries", m.keepalive_retries),
("bind_device", m.bind_device),
("user_timeout", m.user_timeout),
("local_address", m.local_address),
("send_buffer_size", m.send_buffer_size),
("recv_buffer_size", m.recv_buffer_size),
("reuse_address", m.reuse_address),
]
.into_iter()
.filter_map(|(name, missing)| missing.then_some(name))
}
}
impl Display for UnsupportedTcpOpts {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(
"this runtime cannot apply these TCP socket options, and does not ignore them:",
)?;
for (i, name) in self.names().enumerate() {
f.write_str(if i > 0 { ", " } else { " " })?;
f.write_str(name)?;
}
f.write_str(" (a runtime that does apply one declares it in TcpConnect::APPLIES)")
}
}
impl StdError for UnsupportedTcpOpts {}
impl TcpOpts {
pub fn reject_unsupported(&self, can: TcpOptsSupport) -> std::io::Result<()> {
let missing = TcpOptsSupport {
nodelay: self.nodelay && !can.nodelay,
keepalive: self.keepalive.is_some() && !can.keepalive,
keepalive_interval: self.keepalive_interval.is_some() && !can.keepalive_interval,
keepalive_retries: self.keepalive_retries.is_some() && !can.keepalive_retries,
bind_device: self.bind_device.is_some() && !can.bind_device,
user_timeout: self.user_timeout.is_some() && !can.user_timeout,
local_address: self.local_address.is_some() && !can.local_address,
send_buffer_size: self.send_buffer_size.is_some() && !can.send_buffer_size,
recv_buffer_size: self.recv_buffer_size.is_some() && !can.recv_buffer_size,
reuse_address: self.reuse_address && !can.reuse_address,
};
if missing == TcpOptsSupport::NONE {
return Ok(());
}
Err(std::io::Error::new(
std::io::ErrorKind::Unsupported,
UnsupportedTcpOpts { missing },
))
}
}
pub trait TcpConnect {
type Stream: hyper::rt::Read + hyper::rt::Write + Unpin;
const APPLIES: TcpOptsSupport = TcpOptsSupport::NONE;
fn connect(
&self,
addr: SocketAddr,
opts: &TcpOpts,
) -> impl Future<Output = std::io::Result<Self::Stream>>;
const SUPPORTS_UNIX: bool = false;
fn connect_unix(
&self,
path: &std::path::Path,
) -> impl Future<Output = std::io::Result<Self::Stream>> {
let _ = path;
async {
Err(std::io::Error::new(
std::io::ErrorKind::Unsupported,
UnixSocketsUnsupported,
))
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
#[error("this runtime does not connect to Unix-domain sockets")]
pub struct UnixSocketsUnsupported;
pub trait TcpAdoptStd: TcpConnect {
fn adopt(&self, std: std::net::TcpStream) -> std::io::Result<Self::Stream>;
}
pub trait Blocking {
fn run<T, F>(&self, f: F) -> impl Future<Output = Result<T, Cancelled>>
where
T: Send + 'static, F: FnOnce() -> T + Send + 'static; }
#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
#[error("blocking task pool went away before the work started")]
pub struct Cancelled;
#[cfg(test)]
mod tests {
use super::*;
use std::pin::Pin;
use std::task::Context;
use std::task::Poll;
#[test]
fn tcp_opts_default_is_conservative() {
let o = TcpOpts::default();
assert!(!o.nodelay, "the seam has no opinion about who is writing");
assert!(o.keepalive.is_none());
assert!(o.local_address.is_none());
assert!(o.send_buffer_size.is_none());
assert!(o.recv_buffer_size.is_none());
assert!(!o.reuse_address);
}
fn every_field_set() -> TcpOpts {
TcpOpts {
nodelay: true,
keepalive: Some(Duration::from_secs(30)),
keepalive_interval: Some(Duration::from_secs(5)),
keepalive_retries: Some(3),
bind_device: Some("lo".to_owned()),
user_timeout: Some(Duration::from_secs(20)),
local_address: Some(IpAddr::from([127, 0, 0, 1])),
send_buffer_size: Some(4096),
recv_buffer_size: Some(4096),
reuse_address: true,
}
}
const NAMES: &[&str] = &[
"nodelay",
"keepalive",
"keepalive_interval",
"keepalive_retries",
"bind_device",
"user_timeout",
"local_address",
"send_buffer_size",
"recv_buffer_size",
"reuse_address",
];
fn all_but(i: usize) -> TcpOptsSupport {
let mut can = TcpOptsSupport::ALL;
match i {
0 => can.nodelay = false,
1 => can.keepalive = false,
2 => can.keepalive_interval = false,
3 => can.keepalive_retries = false,
4 => can.bind_device = false,
5 => can.user_timeout = false,
6 => can.local_address = false,
7 => can.send_buffer_size = false,
8 => can.recv_buffer_size = false,
9 => can.reuse_address = false,
_ => unreachable!("one arm per NAMES entry"),
}
can
}
#[test]
fn reject_unsupported_is_a_no_op_against_all() {
assert!(
every_field_set()
.reject_unsupported(TcpOptsSupport::ALL)
.is_ok()
);
}
#[test]
fn a_runtime_that_applies_nothing_still_serves_a_caller_that_asked_for_nothing() {
assert!(
TcpOpts::default()
.reject_unsupported(TcpOptsSupport::NONE)
.is_ok()
);
}
#[test]
fn each_unappliable_option_is_named_on_its_own() {
for (i, name) in NAMES.iter().enumerate() {
let err = every_field_set()
.reject_unsupported(all_but(i))
.expect_err("the one option this runtime cannot apply was set");
let named: Vec<&str> = err
.get_ref()
.and_then(|e| e.downcast_ref::<UnsupportedTcpOpts>())
.expect("typed payload")
.names()
.collect();
assert_eq!(
named,
[*name],
"a withheld {name} must be the only option named"
);
assert!(err.to_string().contains(name), "{err}");
}
}
#[test]
fn the_message_names_the_constant_an_implementor_would_have_to_change() {
let err = every_field_set()
.reject_unsupported(all_but(0))
.expect_err("nodelay was withheld");
let msg = err.to_string();
assert!(msg.contains("TcpConnect::APPLIES"), "{msg}");
}
#[test]
fn all_offending_options_are_named_not_only_the_first() {
let err = every_field_set()
.reject_unsupported(TcpOptsSupport::NONE)
.expect_err("nothing can be applied and everything was asked for");
let msg = err.to_string();
for name in NAMES {
assert!(msg.contains(name), "{name} missing from: {msg}");
}
}
#[test]
fn the_error_is_unsupported_and_carries_a_typed_payload() {
const I: usize = 6;
let err = every_field_set()
.reject_unsupported(all_but(I))
.expect_err("one option was withheld");
assert_eq!(err.kind(), std::io::ErrorKind::Unsupported);
let payload = err
.get_ref()
.and_then(|e| e.downcast_ref::<UnsupportedTcpOpts>())
.expect("the typed payload survives the trip through io::Error");
assert_eq!(payload.names().collect::<Vec<_>>(), [NAMES[I]]);
assert_eq!(NAMES[I], "local_address", "the index still names it");
}
#[test]
fn an_option_left_unset_is_not_an_offence_even_when_unsupported() {
let opts = TcpOpts {
nodelay: true,
..TcpOpts::default()
};
let err = opts
.reject_unsupported(TcpOptsSupport::NONE)
.expect_err("nodelay was set and cannot be applied");
let payload = err
.get_ref()
.and_then(|e| e.downcast_ref::<UnsupportedTcpOpts>())
.expect("typed payload");
assert_eq!(payload.names().collect::<Vec<_>>(), ["nodelay"], "{err}");
}
#[test]
fn a_runtime_that_declares_nothing_applies_nothing() {
struct Forgetful;
struct NeverIo;
impl hyper::rt::Read for NeverIo {
fn poll_read(
self: Pin<&mut Self>,
_: &mut Context<'_>,
_: hyper::rt::ReadBufCursor<'_>,
) -> Poll<std::io::Result<()>> {
unreachable!("this runtime never connects")
}
}
impl hyper::rt::Write for NeverIo {
fn poll_write(
self: Pin<&mut Self>,
_: &mut Context<'_>,
_: &[u8],
) -> Poll<std::io::Result<usize>> {
unreachable!("this runtime never connects")
}
fn poll_flush(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<std::io::Result<()>> {
unreachable!("this runtime never connects")
}
fn poll_shutdown(
self: Pin<&mut Self>,
_: &mut Context<'_>,
) -> Poll<std::io::Result<()>> {
unreachable!("this runtime never connects")
}
}
impl TcpConnect for Forgetful {
type Stream = NeverIo;
async fn connect(&self, _: SocketAddr, _: &TcpOpts) -> std::io::Result<NeverIo> {
unreachable!("this runtime never connects")
}
}
assert_eq!(
<Forgetful as TcpConnect>::APPLIES,
TcpOptsSupport::NONE,
"a runtime that declares nothing must not claim to apply anything"
);
let err = every_field_set()
.reject_unsupported(<Forgetful as TcpConnect>::APPLIES)
.expect_err("a runtime that applies nothing must refuse everything asked of it");
let payload = err
.get_ref()
.and_then(|e| e.downcast_ref::<UnsupportedTcpOpts>())
.expect("typed payload");
assert_eq!(payload.names().collect::<Vec<_>>(), NAMES);
}
#[test]
fn spawn_is_generic_over_the_future_not_boxed() {
struct Immediate;
impl<F: std::future::Future<Output = ()>> Spawn<F> for Immediate {
fn spawn(&self, f: F) {
futures_executor::block_on(f)
}
}
let done = std::rc::Rc::new(std::cell::Cell::new(false));
let d = done.clone();
Immediate.spawn(async move { d.set(true) });
assert!(done.get());
}
}