use super::accept;
use crate::{
event,
event::EndpointPublisher,
path::secret,
stream::{
environment::{tokio::Environment, Environment as _},
TlsConnectionBuilder,
},
};
use s2n_quic_core::time::{Clock as _, Timestamp};
use std::{sync::Arc, time::Duration};
use tokio::net::TcpStream;
pub struct Builder {
rt: Arc<tokio::runtime::Runtime>,
config: Arc<dyn TlsConnectionBuilder>,
timeout: Duration,
}
impl Builder {
pub fn new(rt: Arc<tokio::runtime::Runtime>, config: Arc<dyn TlsConnectionBuilder>) -> Self {
Self {
rt,
config,
timeout: Duration::from_secs(1),
}
}
pub fn with_negotiate_timeout(mut self, timeout: Duration) -> Self {
self.timeout = timeout;
self
}
pub(crate) fn build<Sub>(
self,
sender: accept::Sender<Sub>,
env: Environment<Sub>,
map: secret::Map,
accept_flavor: accept::Flavor,
) -> TlsServer<Sub>
where
Sub: event::Subscriber + Clone,
{
TlsServer::new(
self.rt,
self.config,
sender,
env,
map,
accept_flavor,
self.timeout,
)
}
}
#[derive(Clone)]
pub struct TlsServer<Sub>
where
Sub: event::Subscriber + Clone,
{
rt: Option<Arc<tokio::runtime::Runtime>>,
config: Arc<dyn TlsConnectionBuilder>,
sender: accept::Sender<Sub>,
env: Environment<Sub>,
map: secret::Map,
accept_flavor: accept::Flavor,
timeout: std::time::Duration,
}
impl<Sub> Drop for TlsServer<Sub>
where
Sub: event::Subscriber + Clone,
{
fn drop(&mut self) {
#[expect(
clippy::unwrap_used,
reason = "rt is always Some until it is taken here during drop"
)]
if let Some(rt) = Arc::into_inner(self.rt.take().unwrap()) {
rt.shutdown_background();
}
}
}
impl<Sub> TlsServer<Sub>
where
Sub: event::Subscriber + Clone,
{
fn new(
rt: Arc<tokio::runtime::Runtime>,
config: Arc<dyn TlsConnectionBuilder>,
sender: accept::Sender<Sub>,
env: Environment<Sub>,
map: secret::Map,
accept_flavor: accept::Flavor,
timeout: Duration,
) -> Self {
TlsServer {
rt: Some(rt),
sender,
config,
env,
map,
accept_flavor,
timeout,
}
}
pub(crate) fn spawn(
&self,
socket: super::LazyBoundStream,
remote_address: s2n_quic_core::inet::SocketAddress,
buffer: crate::msg::recv::Message,
kernel_accept_time: Timestamp,
) {
match self.spawn_inner(socket, remote_address, buffer, kernel_accept_time) {
Ok(()) => {}
Err(error) => {
self.env
.endpoint_publisher()
.on_acceptor_tcp_tls_stream_rejected(
event::builder::AcceptorTcpTlsStreamRejected {
remote_address: &remote_address,
sojourn_time: self
.env
.clock()
.get_time()
.saturating_duration_since(kernel_accept_time),
error: &error.into(),
},
);
}
}
}
fn spawn_inner(
&self,
socket: super::LazyBoundStream,
remote_addr: s2n_quic_core::inet::SocketAddress,
buffer: crate::msg::recv::Message,
kernel_accept_time: Timestamp,
) -> Result<(), s2n_tls::error::Error> {
let conn = self.config.build_connection(s2n_tls::enums::Mode::Server)?;
let sender = self.sender.clone();
let env = self.env.clone();
let map = self.map.clone();
let flavor = self.accept_flavor;
let timeout = self.timeout;
#[expect(
clippy::unwrap_used,
reason = "rt is always Some until Self is dropped"
)]
let rt = self.rt.as_ref().unwrap();
rt.spawn(async move {
let fut = accept_conn(
socket,
remote_addr,
buffer,
conn,
sender,
&env,
map,
flavor,
kernel_accept_time,
);
let result = tokio::time::timeout(timeout, fut)
.await
.unwrap_or_else(|_| Err(std::io::Error::from(std::io::ErrorKind::TimedOut)));
if let Err(error) = result {
env.endpoint_publisher()
.on_acceptor_tcp_tls_stream_rejected(
event::builder::AcceptorTcpTlsStreamRejected {
remote_address: &remote_addr,
sojourn_time: env
.clock()
.get_time()
.saturating_duration_since(kernel_accept_time),
error: &error,
},
);
}
});
Ok(())
}
}
async fn accept_conn<Sub: event::Subscriber + Clone>(
socket: super::LazyBoundStream,
remote_addr: s2n_quic_core::inet::SocketAddress,
buffer: crate::msg::recv::Message,
conn: crate::stream::TlsConnection,
sender: accept::Sender<Sub>,
env: &Environment<Sub>,
map: secret::Map,
flavor: accept::Flavor,
kernel_accept_time: Timestamp,
) -> std::io::Result<()> {
let socket = match socket {
super::LazyBoundStream::Tokio(tcp_stream) => TcpStream::from_std(tcp_stream.into_std()?)?,
super::LazyBoundStream::Std(tcp_stream) => TcpStream::from_std(tcp_stream)?,
super::LazyBoundStream::TempEmpty => unreachable!(),
};
let socket = Arc::new(crate::stream::socket::application::Single(socket));
let mut connection =
crate::stream::tls::S2nTlsConnection::from_connection(socket.clone(), conn)?;
connection.negotiate(Some(buffer)).await?;
let mut stream_builder = crate::stream::tls::build_stream(
kernel_accept_time,
remote_addr.into(),
socket,
connection,
env,
&map,
s2n_quic_core::endpoint::Type::Server,
)?;
{
let remote_address: s2n_quic_core::inet::SocketAddress =
stream_builder.shared.remote_addr();
let remote_address = &remote_address;
stream_builder.app_queue_time = Some(env.clock().get_time());
env.endpoint_publisher()
.on_acceptor_tcp_tls_stream_enqueued(event::builder::AcceptorTcpTlsStreamEnqueued {
remote_address,
sojourn_time: env
.clock()
.get_time()
.saturating_duration_since(kernel_accept_time),
});
}
let res = match flavor {
accept::Flavor::Fifo => sender.send_back(stream_builder),
accept::Flavor::Lifo => sender.send_front(stream_builder),
};
match res {
Ok(prev) => {
if let Some(stream) = prev {
stream
.prune(event::builder::AcceptorStreamPruneReason::AcceptQueueCapacityExceeded);
}
}
Err(_err) => {
}
}
Ok(())
}
pub fn is_client_hello(buffer: &[u8]) -> Option<bool> {
const HANDSHAKE_TAG: u8 = 22;
const _: () = {
assert!(crate::packet::stream::Tag::IS_RECOVERY_PACKET & HANDSHAKE_TAG != 0);
};
match buffer.first().copied() {
Some(HANDSHAKE_TAG) => Some(true),
Some(_) => Some(false),
None => None,
}
}