use std::num::NonZeroUsize;
use std::sync::Arc;
use bytes::Bytes;
use futures::future::BoxFuture;
use futures::stream::FuturesUnordered;
use futures::{FutureExt, StreamExt};
use tokio::net::UdpSocket;
use tokio_quiche::metrics::DefaultMetrics;
use tokio_quiche::settings::{CertificateKind, Hooks, QuicSettings, TlsCertificatePaths};
use tokio_quiche::{ConnectionParams, QuicConnectionStream};
use crate::buffer::PKT_BUF_LEN;
use crate::driver::{DriverBufferConfig, QuicheDriver, BYTE_CHANNEL_DEPTH};
use crate::endpoint::{EndpointShared, H3QuicheEndpoint};
use crate::stream::Connection;
use crate::Error;
pub const DEFAULT_MAX_IN_FLIGHT_HANDSHAKES: usize = 256;
const DEFAULT_ACCEPT_BIDI_CAP: usize = 128;
const DEFAULT_ACCEPT_UNI_CAP: usize = 128;
#[derive(Clone)]
pub struct H3QuicheServerConfig {
pub cert_path: String,
pub key_path: String,
pub settings: QuicSettings,
pub hooks: Hooks,
pub accept_bidi_cap: usize,
pub accept_uni_cap: usize,
pub max_in_flight_handshakes: NonZeroUsize,
pub recv_channel_depth: usize,
pub packet_buffer_size: usize,
pub max_buffered_send_bytes: Option<usize>,
}
impl Default for H3QuicheServerConfig {
fn default() -> Self {
Self {
cert_path: String::new(),
key_path: String::new(),
settings: QuicSettings::default(),
hooks: Hooks::default(),
accept_bidi_cap: DEFAULT_ACCEPT_BIDI_CAP,
accept_uni_cap: DEFAULT_ACCEPT_UNI_CAP,
max_in_flight_handshakes: NonZeroUsize::new(DEFAULT_MAX_IN_FLIGHT_HANDSHAKES)
.expect("DEFAULT_MAX_IN_FLIGHT_HANDSHAKES is non-zero"),
recv_channel_depth: BYTE_CHANNEL_DEPTH,
packet_buffer_size: PKT_BUF_LEN,
max_buffered_send_bytes: None,
}
}
}
pub struct H3QuicheAcceptor {
stream: QuicConnectionStream<DefaultMetrics>,
handshakes: FuturesUnordered<BoxFuture<'static, Result<Connection<Bytes>, Error>>>,
max_in_flight_handshakes: NonZeroUsize,
accept_bidi_cap: usize,
accept_uni_cap: usize,
buffers: DriverBufferConfig,
incoming_done: bool,
shared: Arc<EndpointShared>,
}
impl H3QuicheAcceptor {
pub fn bind(
sockets: impl IntoIterator<Item = UdpSocket>,
config: &H3QuicheServerConfig,
) -> Result<Vec<Self>, Error> {
let sockets: Vec<UdpSocket> = sockets.into_iter().collect();
if sockets.is_empty() {
return Err("quiche-h3: H3QuicheAcceptor::bind requires at least one socket".into());
}
crate::ensure_readable_file(&config.cert_path, "TLS certificate")?;
crate::ensure_readable_file(&config.key_path, "TLS private key")?;
crate::ensure_nonempty_alpn(&config.settings, "server config")?;
let params = ConnectionParams::new_server(
config.settings.clone(),
TlsCertificatePaths {
cert: &config.cert_path,
private_key: &config.key_path,
kind: CertificateKind::X509,
},
config.hooks.clone(),
);
let streams = tokio_quiche::listen(sockets, params, DefaultMetrics)
.map_err(|e| -> Error { Box::new(e) })?;
let shared = EndpointShared::new();
Ok(streams
.into_iter()
.map(|stream| Self {
stream,
handshakes: FuturesUnordered::new(),
max_in_flight_handshakes: config.max_in_flight_handshakes,
accept_bidi_cap: config.accept_bidi_cap,
accept_uni_cap: config.accept_uni_cap,
buffers: DriverBufferConfig {
recv_channel_depth: config.recv_channel_depth,
packet_buffer_size: config.packet_buffer_size,
max_buffered_send_bytes: config.max_buffered_send_bytes,
},
incoming_done: false,
shared: Arc::clone(&shared),
})
.collect())
}
pub fn endpoint(&self) -> H3QuicheEndpoint {
H3QuicheEndpoint::new(Arc::clone(&self.shared))
}
pub async fn accept(&mut self) -> Result<Option<Connection<Bytes>>, Error> {
loop {
let accept_wake = self.shared.accept_wake.notified();
tokio::pin!(accept_wake);
accept_wake.as_mut().enable();
let closing = self.shared.is_closing();
let admission_done = self.incoming_done || closing;
if admission_done && self.handshakes.is_empty() {
return Ok(None);
}
tokio::select! {
biased;
_ = &mut accept_wake, if !closing => {
continue;
}
Some(res) = self.handshakes.next(), if !self.handshakes.is_empty() => {
match res {
Ok(conn) => {
if closing || self.shared.is_closing() {
drop(conn);
continue;
}
return Ok(Some(conn));
}
Err(_e) => {
#[cfg(feature = "tracing")]
tracing::debug!(
error = %_e,
"quiche-h3: connection setup failed before handshake"
);
continue;
}
}
}
iqc = self.stream.next(),
if !admission_done
&& self.handshakes.len() < self.max_in_flight_handshakes.get() =>
{
let Some(item) = iqc else {
self.incoming_done = true;
continue;
};
match item {
Ok(iqc) => {
let (mut driver, handles) = QuicheDriver::<Bytes>::with_buffers(
true,
self.accept_bidi_cap,
self.accept_uni_cap,
self.buffers,
);
match crate::endpoint::try_register(&self.shared, &handles.cmd_tx) {
None => {
continue;
}
Some(reg) => {
driver.set_conn_registration(reg);
let _qconn = iqc.start(driver);
self.handshakes.push(
handles
.into_established_connection()
.map(|res| res.map_err(|e| -> Error { Box::new(e) }))
.boxed(),
);
}
}
}
Err(_e) => {
#[cfg(feature = "tracing")]
tracing::debug!(
error = %_e,
"quiche-h3: rejected initial connection packet; listener continues"
);
continue;
}
}
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn bind_rejects_empty_socket_set() {
let config = H3QuicheServerConfig::default();
match H3QuicheAcceptor::bind(Vec::<UdpSocket>::new(), &config) {
Ok(_) => panic!("empty socket set must be rejected"),
Err(e) => assert!(e.to_string().contains("at least one socket")),
}
}
#[tokio::test]
async fn bind_rejects_missing_cert() {
let sock = UdpSocket::bind("127.0.0.1:0").await.expect("bind udp");
let config = H3QuicheServerConfig {
cert_path: "/nonexistent/quiche-h3/missing.crt".to_string(),
key_path: "/nonexistent/quiche-h3/missing.key".to_string(),
..H3QuicheServerConfig::default()
};
let err = match H3QuicheAcceptor::bind([sock], &config) {
Ok(_) => panic!("missing cert path must be rejected"),
Err(e) => e,
};
assert!(err.to_string().contains("certificate"));
}
#[test]
fn default_handshake_cap_is_256() {
assert_eq!(DEFAULT_MAX_IN_FLIGHT_HANDSHAKES, 256);
assert_eq!(
H3QuicheServerConfig::default()
.max_in_flight_handshakes
.get(),
256
);
}
#[test]
fn server_config_buffer_defaults_and_overrides() {
let def = H3QuicheServerConfig::default();
assert_eq!(def.recv_channel_depth, BYTE_CHANNEL_DEPTH);
assert_eq!(def.packet_buffer_size, PKT_BUF_LEN);
let custom = H3QuicheServerConfig {
recv_channel_depth: 16,
packet_buffer_size: 8192,
..H3QuicheServerConfig::default()
};
assert_eq!(custom.recv_channel_depth, 16);
assert_eq!(custom.packet_buffer_size, 8192);
}
#[test]
fn server_config_send_cap_defaults_none_and_overrides() {
let def = H3QuicheServerConfig::default();
assert_eq!(def.max_buffered_send_bytes, None);
let custom = H3QuicheServerConfig {
max_buffered_send_bytes: Some(1 << 20),
..H3QuicheServerConfig::default()
};
assert_eq!(custom.max_buffered_send_bytes, Some(1 << 20));
}
}