#[cfg(any(feature = "ring", feature = "aws-lc-rs"))]
use std::sync::Arc;
use std::{num::NonZeroUsize, time::Duration};
use futures::{StreamExt, future::BoxFuture, stream::FuturesUnordered};
#[cfg(any(feature = "ring", feature = "aws-lc-rs"))]
use rustls::pki_types::{CertificateDer, PrivateKeyDer};
use url::Url;
#[cfg(any(feature = "ring", feature = "aws-lc-rs"))]
use crate::{CongestionControl, crypto};
use crate::{Connect, ServerError, Session, Settings};
#[cfg(any(feature = "ring", feature = "aws-lc-rs"))]
pub struct ServerBuilder {
provider: crypto::Provider,
addr: std::net::SocketAddr,
transport: quinn::TransportConfig,
handshake_timeout: Option<Duration>,
max_pending_handshakes: NonZeroUsize,
}
#[cfg(any(feature = "ring", feature = "aws-lc-rs"))]
impl Default for ServerBuilder {
fn default() -> Self {
Self::new()
}
}
#[cfg(any(feature = "ring", feature = "aws-lc-rs"))]
impl ServerBuilder {
pub fn new() -> Self {
Self {
provider: crypto::default_provider(),
addr: std::net::SocketAddr::from(([0_u16; 8], 443)),
transport: quinn::TransportConfig::default(),
handshake_timeout: None,
max_pending_handshakes: NonZeroUsize::MAX,
}
}
pub fn with_addr(self, addr: std::net::SocketAddr) -> Self {
Self { addr, ..self }
}
pub fn with_congestion_control(mut self, algorithm: CongestionControl) -> Self {
match algorithm {
CongestionControl::LowLatency => self.transport.congestion_controller_factory(
Arc::new(quinn::congestion::NewRenoConfig::default()),
),
CongestionControl::Throughput => self
.transport
.congestion_controller_factory(Arc::new(quinn::congestion::BbrConfig::default())),
CongestionControl::Default => self
.transport
.congestion_controller_factory(Arc::new(quinn::congestion::CubicConfig::default())),
};
self
}
pub fn with_transport_config(mut self, transport: quinn::TransportConfig) -> Self {
self.transport = transport;
self
}
pub fn with_handshake_timeout(mut self, timeout: Duration) -> Self {
self.handshake_timeout = Some(timeout);
self
}
pub fn with_max_pending_handshakes(mut self, limit: NonZeroUsize) -> Self {
self.max_pending_handshakes = limit;
self
}
pub fn with_certificate(
self,
chain: Vec<CertificateDer<'static>>,
key: PrivateKeyDer<'static>,
) -> Result<Server, ServerError> {
let mut config = rustls::ServerConfig::builder_with_provider(self.provider.clone())
.with_protocol_versions(&[&rustls::version::TLS13])?
.with_no_client_auth()
.with_single_cert(chain, key)?;
config.alpn_protocols = vec![crate::ALPN.as_bytes().to_vec()];
let config: quinn::crypto::rustls::QuicServerConfig = config
.try_into()
.map_err(|_| ServerError::InvalidCryptoConfiguration)?;
let mut config = quinn::ServerConfig::with_crypto(Arc::new(config));
config.transport_config(Arc::new(self.transport));
let server = quinn::Endpoint::server(config, self.addr)
.map_err(|e| ServerError::IoError(e.into()))?;
Ok(Server::with_accept_limits(
server,
self.handshake_timeout,
self.max_pending_handshakes,
))
}
pub fn with_certificates(
self,
certificates: Vec<(String, Vec<CertificateDer<'static>>, PrivateKeyDer<'static>)>,
) -> Result<Server, ServerError> {
let mut resolver = rustls::server::ResolvesServerCertUsingSni::new();
for (name, chain, key) in certificates {
let certified = rustls::sign::CertifiedKey::from_der(chain, key, &self.provider)?;
resolver.add(&name, certified)?;
}
let mut config = rustls::ServerConfig::builder_with_provider(self.provider.clone())
.with_protocol_versions(&[&rustls::version::TLS13])?
.with_no_client_auth()
.with_cert_resolver(Arc::new(resolver));
config.alpn_protocols = vec![crate::ALPN.as_bytes().to_vec()];
let config: quinn::crypto::rustls::QuicServerConfig = config
.try_into()
.map_err(|_| ServerError::InvalidCryptoConfiguration)?;
let mut config = quinn::ServerConfig::with_crypto(Arc::new(config));
config.transport_config(Arc::new(self.transport));
let server = quinn::Endpoint::server(config, self.addr)
.map_err(|e| ServerError::IoError(e.into()))?;
Ok(Server::with_accept_limits(
server,
self.handshake_timeout,
self.max_pending_handshakes,
))
}
}
pub struct Server {
endpoint: quinn::Endpoint,
accept: FuturesUnordered<BoxFuture<'static, Result<Request, ServerError>>>,
handshake_timeout: Option<Duration>,
max_pending_handshakes: NonZeroUsize,
}
impl Server {
pub fn new(endpoint: quinn::Endpoint) -> Self {
Self {
endpoint,
accept: Default::default(),
handshake_timeout: None,
max_pending_handshakes: NonZeroUsize::MAX,
}
}
fn with_accept_limits(
endpoint: quinn::Endpoint,
handshake_timeout: Option<Duration>,
max_pending_handshakes: NonZeroUsize,
) -> Self {
Self {
endpoint,
accept: Default::default(),
handshake_timeout,
max_pending_handshakes,
}
}
pub fn local_addr(&self) -> Result<std::net::SocketAddr, std::io::Error> {
self.endpoint.local_addr()
}
pub async fn accept(&mut self) -> Option<Result<Request, ServerError>> {
loop {
tokio::select! {
res = self.endpoint.accept(), if self.accept.len() < self.max_pending_handshakes.get() => {
let conn = res?;
let timeout = self.handshake_timeout;
self.accept.push(Box::pin(async move {
let handshake = async move {
let conn = conn.await?;
Request::accept(conn).await
};
match timeout {
Some(timeout) => tokio::time::timeout(timeout, handshake)
.await
.map_err(|_| ServerError::HandshakeTimeout)?,
None => handshake.await,
}
}));
}
Some(res) = self.accept.next() => {
return Some(res);
}
}
}
}
}
pub struct Request {
conn: Option<quinn::Connection>,
settings: Option<Settings>,
connect: Option<Connect>,
url: Url,
}
impl Request {
pub async fn accept(conn: quinn::Connection) -> Result<Self, ServerError> {
let settings = Settings::connect(&conn, false).await?;
let connect = Connect::accept(&conn).await?;
let url = connect.url().clone();
Ok(Self {
conn: Some(conn),
settings: Some(settings),
connect: Some(connect),
url,
})
}
pub fn url(&self) -> &Url {
&self.url
}
pub async fn ok(mut self) -> Result<Session, ServerError> {
let mut connect = self
.connect
.take()
.ok_or(ServerError::RequestAlreadyCompleted)?;
connect.respond(http::StatusCode::OK).await?;
let conn = self
.conn
.take()
.ok_or(ServerError::RequestAlreadyCompleted)?;
let settings = self
.settings
.take()
.ok_or(ServerError::RequestAlreadyCompleted)?;
Ok(Session::new(conn, settings, connect))
}
pub async fn close(mut self, status: http::StatusCode) -> Result<(), ServerError> {
let mut connect = self
.connect
.take()
.ok_or(ServerError::RequestAlreadyCompleted)?;
connect.reject(status).await?;
Ok(())
}
}
impl Drop for Request {
fn drop(&mut self) {
let Some(mut connect) = self.connect.take() else {
return;
};
let conn = self.conn.take();
let settings = self.settings.take();
let url = self.url.clone();
let Ok(runtime) = tokio::runtime::Handle::try_current() else {
tracing::error!(
%url,
"dropped unanswered WebTransport request outside a Tokio runtime; \
unable to send the automatic rejection"
);
return;
};
runtime.spawn(async move {
if let Err(error) = connect
.reject(http::StatusCode::INTERNAL_SERVER_ERROR)
.await
{
tracing::warn!(
%url,
%error,
"failed to send automatic rejection for dropped WebTransport request"
);
} else {
tracing::warn!(
%url,
status = http::StatusCode::INTERNAL_SERVER_ERROR.as_u16(),
"automatically rejected dropped unanswered WebTransport request"
);
}
drop((conn, settings));
});
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn builder_records_handshake_admission_limits() {
let limit = NonZeroUsize::new(17).unwrap();
let builder = ServerBuilder::new()
.with_handshake_timeout(Duration::from_secs(9))
.with_max_pending_handshakes(limit);
assert_eq!(builder.handshake_timeout, Some(Duration::from_secs(9)));
assert_eq!(builder.max_pending_handshakes, limit);
}
}