use super::io::{self, HandshakeFailed};
use crate::{path::secret, psk::io::HandshakeReason};
use s2n_quic::{
provider::{event::Subscriber as Sub, tls::Provider as Prov},
server::Name,
};
use std::{net::SocketAddr, sync::Arc};
use tokio::runtime::Runtime;
use tokio_util::sync::DropGuard;
mod builder;
pub use crate::path::secret::HandshakeKind;
pub use builder::Builder;
#[derive(Clone)]
pub struct Provider {
state: Arc<State>,
}
struct State {
runtime: Option<(Arc<Runtime>, DropGuard)>,
map: secret::Map,
client: io::Client,
local_addr: SocketAddr,
}
fn make_runtime() -> std::io::Result<(Arc<Runtime>, DropGuard)> {
let runtime = Arc::new(
tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()?,
);
let token = tokio_util::sync::CancellationToken::new();
let cancelled = token.clone().cancelled_owned();
let rt = runtime.clone();
std::thread::Builder::new()
.name(String::from("hs-client"))
.spawn(move || {
rt.block_on(cancelled);
})?;
Ok((runtime, token.drop_guard()))
}
impl State {
fn new_runtime<
Provider: Prov + Send + Sync + 'static,
Subscriber: Sub + Send + Sync + 'static,
Event: s2n_quic::provider::event::Subscriber,
>(
addr: SocketAddr,
map: secret::Map,
tls_materials_provider: Provider,
subscriber: Subscriber,
builder: Builder<Event>,
) -> io::Result<Self> {
let (runtime, rt_guard) = make_runtime()?;
let guard = runtime.enter();
let client = io::Client::bind::<Provider, Subscriber, Event>(
addr,
map.clone(),
tls_materials_provider,
subscriber,
builder,
)?;
drop(guard);
Ok(Self {
map,
runtime: Some((runtime, rt_guard)),
local_addr: client.local_addr()?,
client,
})
}
}
impl Provider {
pub fn builder() -> Builder<impl s2n_quic::provider::event::Subscriber> {
Builder::default()
}
pub fn new<
Provider: Prov + Send + Sync + 'static,
Subscriber: Sub + Send + Sync + 'static,
Event: s2n_quic::provider::event::Subscriber,
>(
addr: SocketAddr,
map: secret::Map,
tls_materials_provider: Provider,
subscriber: Subscriber,
builder: Builder<Event>,
server_name: Name,
) -> io::Result<Self> {
let state = State::new_runtime(
addr,
map.clone(),
tls_materials_provider,
subscriber,
builder,
)?;
let state = Arc::new(state);
let weak = Arc::downgrade(&state);
map.register_request_handshake(Box::new(move |peer, reason| {
let state = weak.upgrade()?;
let runtime = state.runtime.as_ref().map(|v| &v.0)?;
let client = state.client.clone();
let server_name = server_name.clone();
Some(runtime.spawn(async move {
if let Err(HandshakeFailed { .. }) = client.connect(peer, reason, server_name).await
{
}
}))
}));
Ok(Self { state })
}
#[inline]
pub async fn handshake_with(
&self,
peer: SocketAddr,
server_name: Name,
) -> std::io::Result<HandshakeKind> {
let (_peer, kind) = self.handshake_with_entry(peer, server_name).await?;
Ok(kind)
}
#[inline]
#[doc(hidden)]
pub async fn handshake_with_entry(
&self,
peer: SocketAddr,
server_name: Name,
) -> std::io::Result<(secret::map::Peer, HandshakeKind)> {
if let Some(peer) = self.state.map.get_tracked(peer) {
return Ok((peer, HandshakeKind::Cached));
}
if self.state.runtime.is_some() {
let _ = self.background_handshake_with(peer, server_name.clone());
}
let state = self.state.clone();
if let Some((runtime, _)) = self.state.runtime.as_ref() {
runtime
.spawn(async move {
state
.client
.connect(peer, HandshakeReason::User, server_name)
.await
})
.await??;
} else {
state
.client
.connect(peer, HandshakeReason::User, server_name)
.await?;
}
let peer = self.state.map.get_untracked(peer).ok_or_else(|| {
std::io::Error::new(
std::io::ErrorKind::NotFound,
format!("handshake failed to exchange credentials for {peer}"),
)
})?;
Ok((peer, HandshakeKind::Fresh))
}
#[inline]
#[expect(
clippy::panic,
clippy::panic_in_result_fn,
reason = "the panic is only reachable in the deterministic testing configuration where no runtime is present"
)]
pub fn background_handshake_with(
&self,
peer: SocketAddr,
server_name: Name,
) -> std::io::Result<HandshakeKind> {
if self.state.map.contains(&peer) {
return Ok(HandshakeKind::Cached);
}
let client = self.state.client.clone();
if let Some((runtime, _)) = self.state.runtime.as_ref() {
runtime.spawn(async move {
if let Err(HandshakeFailed { .. }) = client
.connect(peer, HandshakeReason::User, server_name)
.await
{
}
});
} else {
panic!("background_handshake_with not supported with deterministic testing");
}
Ok(HandshakeKind::Fresh)
}
#[inline]
#[expect(
clippy::panic,
clippy::panic_in_result_fn,
reason = "the panic is only reachable in the deterministic testing configuration where no runtime is present"
)]
pub fn blocking_handshake_with(
&self,
peer: SocketAddr,
server_name: Name,
) -> std::io::Result<HandshakeKind> {
if self.state.map.contains(&peer) {
return Ok(HandshakeKind::Cached);
}
let fut = self
.state
.client
.connect(peer, HandshakeReason::User, server_name);
if let Some((runtime, _)) = self.state.runtime.as_ref() {
runtime.block_on(fut)?
} else {
panic!("blocking_handshake_with not supported with deterministic testing");
}
debug_assert!(self.state.map.contains(&peer));
Ok(HandshakeKind::Fresh)
}
#[inline]
#[doc(hidden)]
pub async fn unconditionally_handshake_with_entry(
&self,
peer: SocketAddr,
server_name: Name,
) -> std::io::Result<secret::map::Peer> {
let state = self.state.clone();
if let Some((runtime, _)) = self.state.runtime.as_ref() {
runtime
.spawn(async move {
state
.client
.connect(peer, HandshakeReason::User, server_name)
.await
})
.await??;
} else {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"missing runtime for handshake client",
));
}
let peer = self.state.map.get_untracked(peer).ok_or_else(|| {
std::io::Error::new(
std::io::ErrorKind::NotFound,
format!("handshake failed to exchange credentials for {peer}"),
)
})?;
Ok(peer)
}
#[inline]
pub fn local_addr(&self) -> std::io::Result<SocketAddr> {
Ok(self.state.local_addr)
}
pub fn map(&self) -> &secret::Map {
&self.state.map
}
}