use crate::{
either::Either,
event::{self, builder::StreamTcpConnectErrorReason, EndpointPublisher},
msg,
path::secret,
stream::{
application::Stream,
client::{rpc as rpc_internal, tokio as client},
endpoint,
environment::{
tokio::{self as env, Environment},
udp as udp_pool, Environment as _,
},
recv, socket,
},
};
use s2n_quic::server::Name;
use s2n_quic_core::time::{Clock, Timestamp};
use std::{io, net::SocketAddr, sync::Arc, time::Duration};
use tokio::net::TcpStream;
pub mod rpc {
pub use crate::stream::client::rpc::{InMemoryResponse, Request, Response};
}
#[allow(async_fn_in_trait)]
pub trait Handshake: Clone {
async fn handshake_with_entry(
&self,
remote_handshake_addr: SocketAddr,
server_name: Name,
) -> std::io::Result<(secret::map::Peer, secret::HandshakeKind)>;
fn local_addr(&self) -> std::io::Result<SocketAddr>;
fn map(&self) -> &secret::Map;
}
impl Handshake for crate::psk::client::Provider {
async fn handshake_with_entry(
&self,
remote_handshake_addr: SocketAddr,
server_name: Name,
) -> std::io::Result<(secret::map::Peer, secret::HandshakeKind)> {
self.handshake_with_entry(remote_handshake_addr, server_name)
.await
}
fn local_addr(&self) -> std::io::Result<SocketAddr> {
self.local_addr()
}
fn map(&self) -> &secret::Map {
self.map()
}
}
#[derive(Clone)]
pub struct Client<H: Handshake + Clone, S: event::Subscriber + Clone> {
env: Environment<S>,
handshake: H,
default_protocol: socket::Protocol,
linger: Option<Duration>,
}
impl<H: Handshake + Clone, S: event::Subscriber + Clone> Client<H, S> {
#[inline]
pub fn new(handshake: H, subscriber: S) -> io::Result<Self> {
Self::builder().build(handshake, subscriber)
}
#[inline]
pub fn builder() -> Builder {
Builder::default()
}
pub fn drop_state(&self) {
self.handshake.map().drop_state()
}
pub fn handshake_state(&self) -> &H {
&self.handshake
}
#[inline]
pub async fn handshake_with(
&self,
remote_handshake_addr: SocketAddr,
server_name: Name,
) -> io::Result<secret::HandshakeKind> {
let (_peer, kind) = self
.handshake
.handshake_with_entry(remote_handshake_addr, server_name)
.await?;
Ok(kind)
}
#[inline]
async fn handshake_for_connect(
&self,
remote_handshake_addr: SocketAddr,
server_name: Name,
) -> io::Result<secret::map::Peer> {
let (peer, _kind) = self
.handshake
.handshake_with_entry(remote_handshake_addr, server_name)
.await?;
Ok(peer)
}
#[inline]
pub async fn connect(
&self,
handshake_addr: SocketAddr,
acceptor_addr: SocketAddr,
server_name: Name,
) -> io::Result<Stream<S>> {
match self.default_protocol {
socket::Protocol::Udp => {
self.connect_udp(handshake_addr, acceptor_addr, server_name)
.await
}
socket::Protocol::Tcp => {
self.connect_tcp(handshake_addr, acceptor_addr, server_name)
.await
}
protocol => Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!("invalid default protocol {protocol:?}"),
)),
}
}
pub async fn rpc<Req, Res>(
&self,
handshake_addr: SocketAddr,
acceptor_addr: SocketAddr,
request: Req,
response: Res,
server_name: Name,
) -> io::Result<Res::Output>
where
Req: rpc::Request,
Res: rpc::Response,
{
match self.default_protocol {
socket::Protocol::Udp => {
self.rpc_udp(
handshake_addr,
acceptor_addr,
request,
response,
server_name,
)
.await
}
socket::Protocol::Tcp => {
self.rpc_tcp(
handshake_addr,
acceptor_addr,
request,
response,
server_name,
)
.await
}
protocol => Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!("invalid default protocol {protocol:?}"),
)),
}
}
#[inline]
pub async fn connect_udp(
&self,
handshake_addr: SocketAddr,
acceptor_addr: SocketAddr,
server_name: Name,
) -> io::Result<Stream<S>> {
let handshake = self.handshake_for_connect(handshake_addr, server_name);
let mut stream = client::connect_udp(handshake, acceptor_addr, &self.env).await?;
Self::write_prelude(&mut stream).await?;
Ok(stream)
}
#[inline]
pub async fn rpc_udp<Req, Res>(
&self,
handshake_addr: SocketAddr,
acceptor_addr: SocketAddr,
request: Req,
response: Res,
server_name: Name,
) -> io::Result<Res::Output>
where
Req: rpc::Request,
Res: rpc::Response,
{
let handshake = self.handshake_for_connect(handshake_addr, server_name);
let stream = client::connect_udp(handshake, acceptor_addr, &self.env).await?;
rpc_internal::from_stream(stream, request, response).await
}
#[inline]
pub async fn connect_tcp(
&self,
handshake_addr: SocketAddr,
acceptor_addr: SocketAddr,
server_name: Name,
) -> io::Result<Stream<S>> {
let handshake = self.handshake_for_connect(handshake_addr, server_name);
let mut stream =
client::connect_tcp(handshake, acceptor_addr, &self.env, self.linger).await?;
Self::write_prelude(&mut stream).await?;
Ok(stream)
}
#[inline]
pub async fn connect_tls(
&self,
addr: SocketAddr,
server_name: Name,
config: &impl crate::stream::TlsConnectionBuilder,
) -> io::Result<Stream<S>> {
let stream = client::connect_tls(
addr,
server_name,
config,
&self.env,
self.linger,
self.handshake.map(),
)
.await?;
Ok(stream)
}
#[inline]
pub async fn rpc_tcp<Req, Res>(
&self,
handshake_addr: SocketAddr,
acceptor_addr: SocketAddr,
request: Req,
response: Res,
server_name: Name,
) -> io::Result<Res::Output>
where
Req: rpc::Request,
Res: rpc::Response,
{
let handshake = self.handshake_for_connect(handshake_addr, server_name);
let stream = client::connect_tcp(handshake, acceptor_addr, &self.env, self.linger).await?;
rpc_internal::from_stream(stream, request, response).await
}
#[inline]
pub async fn connect_tls_with(
&self,
stream: TcpStream,
server_name: Name,
config: &impl crate::stream::TlsConnectionBuilder,
) -> io::Result<Stream<S>> {
let stream = client::connect_tls_with(
stream,
server_name,
config,
&self.env,
self.linger,
self.handshake.map(),
)
.await?;
Ok(stream)
}
#[inline]
pub async fn connect_tcp_with(
&self,
handshake_addr: SocketAddr,
stream: TcpStream,
server_name: Name,
) -> io::Result<Stream<S>> {
let handshake = self
.handshake_for_connect(handshake_addr, server_name)
.await?;
let mut stream = client::connect_tcp_with(handshake, stream, &self.env).await?;
Self::write_prelude(&mut stream).await?;
Ok(stream)
}
#[inline]
async fn write_prelude(stream: &mut Stream<S>) -> io::Result<()> {
stream
.write_from(&mut s2n_quic_core::buffer::reader::storage::Empty)
.await
.map(|_| ())
}
}
#[derive(Default)]
pub struct Builder {
default_protocol: Option<socket::Protocol>,
background_threads: Option<usize>,
linger: Option<Duration>,
send_buffer: Option<usize>,
recv_buffer: Option<usize>,
}
impl Builder {
pub fn with_tcp(self, enabled: bool) -> Self {
self.with_default_protocol(if enabled {
socket::Protocol::Tcp
} else {
socket::Protocol::Udp
})
}
pub fn with_udp(self, enabled: bool) -> Self {
self.with_default_protocol(if enabled {
socket::Protocol::Udp
} else {
socket::Protocol::Tcp
})
}
pub fn with_default_protocol(mut self, protocol: socket::Protocol) -> Self {
self.default_protocol = Some(protocol);
self
}
pub fn with_background_threads(mut self, threads: usize) -> Self {
self.background_threads = Some(threads);
self
}
pub fn with_linger(mut self, linger: Duration) -> Self {
self.linger = Some(linger);
self
}
pub fn with_send_buffer(mut self, bytes: usize) -> Self {
self.send_buffer = Some(bytes);
self
}
pub fn with_recv_buffer(mut self, bytes: usize) -> Self {
self.recv_buffer = Some(bytes);
self
}
#[inline]
pub fn build<H: Handshake + Clone, S: event::Subscriber + Clone>(
self,
handshake: H,
subscriber: S,
) -> io::Result<Client<H, S>> {
let mut local_addr = handshake.local_addr()?;
local_addr.set_port(0);
let mut options = socket::Options::new(local_addr);
options.send_buffer = self.send_buffer;
options.recv_buffer = self.recv_buffer;
let mut env = env::Builder::new(subscriber).with_socket_options(options);
let pool = udp_pool::Config::new(handshake.map().clone());
env = env.with_pool(pool);
if let Some(threads) = self.background_threads {
env = env.with_threads(threads);
}
let env = env.build()?;
let default_protocol = self.default_protocol.unwrap_or(socket::Protocol::Udp);
let linger = self.linger;
Ok(Client {
env,
handshake,
default_protocol,
linger,
})
}
}
#[inline]
pub async fn connect_udp<H, Sub>(
handshake: H,
acceptor_addr: SocketAddr,
env: &Environment<Sub>,
) -> io::Result<Stream<Sub>>
where
H: core::future::Future<Output = io::Result<secret::map::Peer>>,
Sub: event::Subscriber + Clone,
{
let entry = handshake.await?;
let stream = if env.has_recv_pool() {
let peer = env::udp::Pooled(acceptor_addr.into());
endpoint::open_stream(env, entry, peer, None)?
} else {
let peer = env::udp::Owned(acceptor_addr.into(), recv_buffer());
endpoint::open_stream(env, entry, peer, None)?
};
let stream = stream.connect()?;
debug_assert_eq!(stream.protocol(), socket::Protocol::Udp);
Ok(stream)
}
struct DropGuard<'a, S: event::Subscriber + Clone> {
env: &'a Environment<S>,
start: Timestamp,
reason: Option<StreamTcpConnectErrorReason>,
}
impl<S: event::Subscriber + Clone> Drop for DropGuard<'_, S> {
fn drop(&mut self) {
if let Some(reason) = self.reason.take() {
let now = self.env.clock().get_time();
self.env
.endpoint_publisher_with_time(now)
.on_stream_connect_error(event::builder::StreamConnectError {
reason,
latency: now.saturating_duration_since(self.start),
});
}
}
}
#[inline]
pub async fn connect_tcp<H, Sub>(
handshake: H,
acceptor_addr: SocketAddr,
env: &Environment<Sub>,
linger: Option<Duration>,
) -> io::Result<Stream<Sub>>
where
H: core::future::Future<Output = io::Result<secret::map::Peer>>,
Sub: event::Subscriber + Clone,
{
let start = env.clock().get_time();
let mut guard = DropGuard {
env,
reason: Some(StreamTcpConnectErrorReason::AbortedPendingBoth),
start,
};
let connect = TcpStream::connect(acceptor_addr);
tokio::pin!(handshake);
tokio::pin!(connect);
let mut error = None;
let mut socket = None;
let mut peer = None;
while (socket.is_none() || peer.is_none()) && error.is_none() {
tokio::select! {
connected = &mut connect, if socket.is_none() => {
let now = env.clock().get_time();
env.endpoint_publisher_with_time(now).on_stream_tcp_connect(event::builder::StreamTcpConnect {
error: connected.is_err(),
latency: now.saturating_duration_since(start),
});
match connected {
Ok(v) => {
socket = Some(Ok(v));
guard.reason = match guard.reason.clone() {
Some(StreamTcpConnectErrorReason::AbortedPendingBoth) => Some(
StreamTcpConnectErrorReason::AbortedPendingHandshake
),
other => other,
};
},
Err(e) => {
guard.reason = Some(StreamTcpConnectErrorReason::TcpConnect);
error = Some(e);
socket = Some(Err(()));
}
}
}
handshaked = &mut handshake, if peer.is_none() => {
match handshaked {
Ok(v) => {
peer = Some(Ok(v));
guard.reason = match guard.reason.clone() {
Some(StreamTcpConnectErrorReason::AbortedPendingBoth) => Some(
StreamTcpConnectErrorReason::AbortedPendingConnect
),
other => other,
};
},
Err(e) => {
guard.reason = Some(StreamTcpConnectErrorReason::Handshake);
error = Some(e);
peer = Some(Err(()));
}
}
}
}
}
if error.is_none() {
guard.reason = None;
}
env.endpoint_publisher()
.on_stream_connect(event::builder::StreamConnect {
error: error.is_some(),
handshake_success: match &peer {
Some(Ok(_)) => event::builder::MaybeBoolCounter::Success,
Some(Err(_)) => event::builder::MaybeBoolCounter::Failure,
None => event::builder::MaybeBoolCounter::Aborted,
},
tcp_success: match &socket {
Some(Ok(_)) => event::builder::MaybeBoolCounter::Success,
Some(Err(_)) => event::builder::MaybeBoolCounter::Failure,
None => event::builder::MaybeBoolCounter::Aborted,
},
});
let (Some(Ok(socket)), Some(Ok(entry))) = (socket, peer) else {
#[expect(
clippy::unwrap_used,
reason = "if socket or peer isn't present the error is always set, as documented above"
)]
return Err(error.unwrap());
};
let _ = socket.set_nodelay(true);
if linger.is_some() {
#[allow(deprecated)]
let _ = socket.set_linger(linger);
}
let peer_addr = if acceptor_addr.ip().is_unspecified() {
socket.peer_addr()?
} else {
acceptor_addr
}
.into();
let local_port = socket.local_addr()?.port();
let peer = env::tcp::Registered {
socket,
peer_addr,
local_port,
recv_buffer: recv_buffer(),
};
let stream = endpoint::open_stream(env, entry, peer, None)?;
let stream = stream.connect()?;
debug_assert_eq!(stream.protocol(), socket::Protocol::Tcp);
Ok(stream)
}
#[inline]
pub async fn connect_tcp_with<Sub>(
entry: secret::map::Peer,
socket: TcpStream,
env: &Environment<Sub>,
) -> io::Result<Stream<Sub>>
where
Sub: event::Subscriber + Clone,
{
let local_port = socket.local_addr()?.port();
let peer_addr = socket.peer_addr()?.into();
let peer = env::tcp::Registered {
socket,
peer_addr,
local_port,
recv_buffer: recv_buffer(),
};
let stream = endpoint::open_stream(env, entry, peer, None)?;
let stream = stream.connect()?;
debug_assert_eq!(stream.protocol(), socket::Protocol::Tcp);
Ok(stream)
}
#[inline]
fn recv_buffer() -> recv::shared::RecvBuffer {
let recv_buffer = recv::buffer::Local::new(msg::recv::Message::new(9000), None);
Either::A(recv_buffer)
}
#[inline]
pub async fn connect_tls<Sub>(
addr: SocketAddr,
server_name: Name,
config: &impl crate::stream::TlsConnectionBuilder,
env: &Environment<Sub>,
linger: Option<Duration>,
map: &crate::path::secret::Map,
) -> io::Result<Stream<Sub>>
where
Sub: event::Subscriber + Clone,
{
let start = env.clock().get_time();
let mut guard = DropGuard {
env,
reason: Some(StreamTcpConnectErrorReason::AbortedPendingBoth),
start,
};
let connected = TcpStream::connect(addr).await;
let kernel_start_time = env.clock().get_time();
env.endpoint_publisher_with_time(kernel_start_time)
.on_stream_tcp_connect(event::builder::StreamTcpConnect {
error: connected.is_err(),
latency: kernel_start_time.saturating_duration_since(start),
});
let socket = match connected {
Ok(v) => {
guard.reason = Some(StreamTcpConnectErrorReason::AbortedPendingHandshake);
v
}
Err(e) => {
guard.reason = Some(StreamTcpConnectErrorReason::TcpConnect);
return Err(e);
}
};
let stream = NegotiateTls {
socket,
addr,
server_name,
config,
env,
linger,
map,
start,
kernel_start_time,
guard,
}
.negotiate()
.await?;
Ok(stream)
}
#[inline]
pub async fn connect_tls_with<Sub>(
socket: TcpStream,
server_name: Name,
config: &impl crate::stream::TlsConnectionBuilder,
env: &Environment<Sub>,
linger: Option<Duration>,
map: &crate::path::secret::Map,
) -> io::Result<Stream<Sub>>
where
Sub: event::Subscriber + Clone,
{
let start = env.clock().get_time();
let addr = socket.peer_addr()?;
let guard = DropGuard {
env,
reason: Some(StreamTcpConnectErrorReason::AbortedPendingHandshake),
start,
};
let stream = NegotiateTls {
socket,
addr,
server_name,
config,
env,
linger,
map,
start,
kernel_start_time: start,
guard,
}
.negotiate()
.await?;
Ok(stream)
}
struct NegotiateTls<'a, C, Sub>
where
C: crate::stream::TlsConnectionBuilder,
Sub: event::Subscriber + Clone,
{
socket: TcpStream,
addr: SocketAddr,
server_name: Name,
config: &'a C,
env: &'a Environment<Sub>,
linger: Option<Duration>,
map: &'a crate::path::secret::Map,
start: Timestamp,
kernel_start_time: Timestamp,
guard: DropGuard<'a, Sub>,
}
impl<C, Sub> NegotiateTls<'_, C, Sub>
where
C: crate::stream::TlsConnectionBuilder,
Sub: event::Subscriber + Clone,
{
#[inline]
async fn negotiate(self) -> io::Result<Stream<Sub>> {
let NegotiateTls {
socket,
addr,
server_name,
config,
env,
linger,
map,
start,
kernel_start_time,
mut guard,
} = self;
let _ = socket.set_nodelay(true);
if linger.is_some() {
#[allow(deprecated)]
let _ = socket.set_linger(linger);
}
let mut connection = config.build_connection(s2n_tls::enums::Mode::Client)?;
(*connection).as_mut().set_server_name(&server_name)?;
let socket = Arc::new(crate::stream::socket::application::Single(socket));
let mut connection =
crate::stream::tls::S2nTlsConnection::from_connection(socket.clone(), connection)?;
let res = connection.negotiate(None).await;
let negotiate_end = env.clock().get_time();
env.endpoint_publisher_with_time(negotiate_end)
.on_stream_tls_connect(event::builder::StreamTlsConnect {
error: res.is_err(),
tcp_latency: kernel_start_time.saturating_duration_since(start),
tls_latency: negotiate_end.saturating_duration_since(kernel_start_time),
});
res?;
guard.reason = None;
crate::stream::tls::build_stream(
kernel_start_time,
addr,
socket,
connection,
env,
map,
s2n_quic_core::endpoint::Type::Client,
)?
.build()
}
}