use std::net::SocketAddr;
use std::sync::Arc;
use bytes::Bytes;
use tokio::net::UdpSocket;
use tokio_quiche::quic::connect_with_config;
use tokio_quiche::settings::{CertificateKind, Hooks, QuicSettings, TlsCertificatePaths};
use tokio_quiche::socket::Socket;
use tokio_quiche::ConnectionParams;
use crate::buffer::PKT_BUF_LEN;
use crate::driver::{DriverBufferConfig, QuicheDriver, BYTE_CHANNEL_DEPTH};
use crate::stream::Connection;
use crate::Error;
const DEFAULT_ACCEPT_BIDI_CAP: usize = 128;
const DEFAULT_ACCEPT_UNI_CAP: usize = 128;
#[derive(Clone)]
pub struct H3QuicheClientConfig {
pub settings: QuicSettings,
pub hooks: Hooks,
pub cert_path: Option<String>,
pub key_path: Option<String>,
pub verify_peer: bool,
pub server_name: Option<String>,
pub recv_channel_depth: usize,
pub packet_buffer_size: usize,
pub max_buffered_send_bytes: Option<usize>,
}
impl Default for H3QuicheClientConfig {
fn default() -> Self {
Self {
settings: QuicSettings::default(),
hooks: Hooks::default(),
cert_path: None,
key_path: None,
verify_peer: true,
server_name: None,
recv_channel_depth: BYTE_CHANNEL_DEPTH,
packet_buffer_size: PKT_BUF_LEN,
max_buffered_send_bytes: None,
}
}
}
struct Inner {
server_addr: SocketAddr,
server_name: String,
config: H3QuicheClientConfig,
}
#[derive(Clone)]
pub struct H3QuicheConnector {
inner: Arc<Inner>,
}
impl H3QuicheConnector {
pub fn new(
server_addr: SocketAddr,
server_name: String,
config: H3QuicheClientConfig,
) -> Result<Self, Error> {
match (&config.cert_path, &config.key_path) {
(Some(cert), Some(key)) => {
crate::ensure_readable_file(cert, "client TLS certificate")?;
crate::ensure_readable_file(key, "client TLS private key")?;
}
(Some(_), None) => {
return Err("quiche-h3: client mTLS certificate set without a private key".into());
}
(None, Some(_)) => {
return Err("quiche-h3: client mTLS private key set without a certificate".into());
}
(None, None) => {}
}
crate::ensure_nonempty_alpn(&config.settings, "client config")?;
Ok(Self {
inner: Arc::new(Inner {
server_addr,
server_name,
config,
}),
})
}
pub async fn connect(&self) -> Result<Connection<Bytes>, Error> {
let inner = &self.inner;
let bind_addr = if inner.server_addr.is_ipv6() {
"[::]:0"
} else {
"0.0.0.0:0"
};
let udp = UdpSocket::bind(bind_addr)
.await
.map_err(|e| -> Error { Box::new(e) })?;
udp.connect(inner.server_addr)
.await
.map_err(|e| -> Error { Box::new(e) })?;
let socket = Socket::try_from(udp).map_err(|e| -> Error { Box::new(e) })?;
let mut settings = inner.config.settings.clone();
settings.verify_peer = inner.config.verify_peer;
let tls = match (&inner.config.cert_path, &inner.config.key_path) {
(Some(cert), Some(key)) => Some(TlsCertificatePaths {
cert,
private_key: key,
kind: CertificateKind::X509,
}),
_ => None,
};
let params = ConnectionParams::new_client(settings, tls, inner.config.hooks.clone());
let (driver, handles) = QuicheDriver::<Bytes>::with_buffers(
false,
DEFAULT_ACCEPT_BIDI_CAP,
DEFAULT_ACCEPT_UNI_CAP,
DriverBufferConfig {
recv_channel_depth: inner.config.recv_channel_depth,
packet_buffer_size: inner.config.packet_buffer_size,
max_buffered_send_bytes: inner.config.max_buffered_send_bytes,
},
);
match connect_with_config(socket, Some(&inner.server_name), ¶ms, driver).await {
Ok(_qconn) => {
handles
.into_established_connection()
.await
.map_err(|e| -> Error { Box::new(e) })
}
Err(e) => Err(e),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn addr() -> SocketAddr {
"127.0.0.1:4433".parse().unwrap()
}
#[test]
fn new_rejects_unreadable_cert() {
let config = H3QuicheClientConfig {
cert_path: Some("/nonexistent/quiche-h3/missing.crt".to_string()),
key_path: Some("/nonexistent/quiche-h3/missing.key".to_string()),
..H3QuicheClientConfig::default()
};
let err = match H3QuicheConnector::new(addr(), "localhost".to_string(), config) {
Ok(_) => panic!("unreadable client cert must be rejected"),
Err(e) => e,
};
assert!(err.to_string().contains("certificate"));
}
#[test]
fn new_accepts_config_without_mtls() {
let config = H3QuicheClientConfig::default();
let connector = H3QuicheConnector::new(addr(), "localhost".to_string(), config)
.expect("no-mTLS config is valid");
let _cloned = connector.clone();
}
#[test]
fn client_config_buffer_defaults_and_overrides() {
let def = H3QuicheClientConfig::default();
assert_eq!(def.recv_channel_depth, BYTE_CHANNEL_DEPTH);
assert_eq!(def.packet_buffer_size, PKT_BUF_LEN);
let custom = H3QuicheClientConfig {
recv_channel_depth: 32,
packet_buffer_size: 16384,
..H3QuicheClientConfig::default()
};
assert_eq!(custom.recv_channel_depth, 32);
assert_eq!(custom.packet_buffer_size, 16384);
}
#[test]
fn client_config_send_cap_defaults_none_and_overrides() {
let def = H3QuicheClientConfig::default();
assert_eq!(def.max_buffered_send_bytes, None);
let custom = H3QuicheClientConfig {
max_buffered_send_bytes: Some(512 * 1024),
..H3QuicheClientConfig::default()
};
assert_eq!(custom.max_buffered_send_bytes, Some(512 * 1024));
}
#[test]
fn new_rejects_partial_mtls() {
let cert_only = H3QuicheClientConfig {
cert_path: Some("/tmp/some.crt".to_string()),
key_path: None,
..H3QuicheClientConfig::default()
};
assert!(H3QuicheConnector::new(addr(), "localhost".to_string(), cert_only).is_err());
let key_only = H3QuicheClientConfig {
cert_path: None,
key_path: Some("/tmp/some.key".to_string()),
..H3QuicheClientConfig::default()
};
assert!(H3QuicheConnector::new(addr(), "localhost".to_string(), key_only).is_err());
}
#[test]
fn new_rejects_directory_as_cert() {
let dir = std::env::temp_dir();
let config = H3QuicheClientConfig {
cert_path: Some(dir.to_string_lossy().into_owned()),
key_path: Some(dir.to_string_lossy().into_owned()),
..H3QuicheClientConfig::default()
};
let err = match H3QuicheConnector::new(addr(), "localhost".to_string(), config) {
Ok(_) => panic!("a directory is not a valid cert file"),
Err(e) => e,
};
assert!(err.to_string().contains("regular file"));
}
#[test]
fn new_rejects_empty_alpn() {
let mut settings = QuicSettings::default();
settings.alpn = Vec::new();
let config = H3QuicheClientConfig {
settings,
..H3QuicheClientConfig::default()
};
let err = match H3QuicheConnector::new(addr(), "localhost".to_string(), config) {
Ok(_) => panic!("empty ALPN must be rejected"),
Err(e) => e,
};
assert!(err.to_string().contains("ALPN"));
}
}