use crate::error::{Error, Result};
use crate::network::Message;
use quinn::{ClientConfig, Connection, Endpoint, ServerConfig};
use rustls::pki_types::{CertificateDer, PrivateKeyDer, PrivatePkcs8KeyDer};
use std::collections::HashMap;
use std::net::SocketAddr;
use std::sync::Arc;
use std::time::Duration;
#[derive(Debug, Clone)]
pub struct QuicConfig {
pub bind_addr: String,
pub port: u16,
pub keep_alive: Duration,
pub idle_timeout: Duration,
pub max_concurrent_streams: u32,
pub max_connections: usize,
}
impl Default for QuicConfig {
fn default() -> Self {
Self {
bind_addr: "0.0.0.0".to_string(),
port: 19081,
keep_alive: Duration::from_secs(10),
idle_timeout: Duration::from_secs(30),
max_concurrent_streams: 100,
max_connections: 1000,
}
}
}
impl QuicConfig {
pub fn iot_mode() -> Self {
Self {
bind_addr: "0.0.0.0".to_string(),
port: 19081,
keep_alive: Duration::from_secs(30),
idle_timeout: Duration::from_secs(60),
max_concurrent_streams: 10,
max_connections: 100,
}
}
pub fn production() -> Self {
Self {
bind_addr: "0.0.0.0".to_string(),
port: 19081,
keep_alive: Duration::from_secs(5),
idle_timeout: Duration::from_secs(30),
max_concurrent_streams: 1000,
max_connections: 10000,
}
}
}
pub struct QuicServer {
config: QuicConfig,
endpoint: Option<Endpoint>,
connections: HashMap<SocketAddr, Connection>,
node_id: String,
}
impl QuicServer {
pub fn new(config: QuicConfig, node_id: String) -> Self {
Self {
config,
endpoint: None,
connections: HashMap::new(),
node_id,
}
}
pub async fn start(&mut self) -> Result<()> {
let addr: SocketAddr = format!("{}:{}", self.config.bind_addr, self.config.port)
.parse()
.map_err(|e| Error::network(format!("Invalid address: {}", e)))?;
let (server_config, _cert) = self.generate_server_config()?;
let endpoint = Endpoint::server(server_config, addr)
.map_err(|e| Error::network(format!("Failed to create QUIC endpoint: {}", e)))?;
log::info!("QUIC server started on {} (node: {})", addr, self.node_id);
self.endpoint = Some(endpoint);
Ok(())
}
pub async fn accept(&mut self) -> Result<Option<SocketAddr>> {
let endpoint = self.endpoint.as_ref().ok_or(Error::NotInitialized)?;
if let Some(incoming) = endpoint.accept().await {
match incoming.await {
Ok(connection) => {
let remote = connection.remote_address();
log::debug!("QUIC connection accepted from {}", remote);
self.connections.insert(remote, connection);
return Ok(Some(remote));
}
Err(e) => {
log::warn!("Failed to accept QUIC connection: {}", e);
}
}
}
Ok(None)
}
pub async fn connect(&mut self, addr: &SocketAddr) -> Result<()> {
let endpoint = self.endpoint.as_ref().ok_or(Error::NotInitialized)?;
let client_config = self.generate_client_config()?;
let connection = endpoint
.connect_with(client_config, *addr, "localhost")
.map_err(|e| Error::network(format!("Failed to initiate connection: {}", e)))?
.await
.map_err(|e| Error::network(format!("Connection failed: {}", e)))?;
log::debug!("QUIC connection established to {}", addr);
self.connections.insert(*addr, connection);
Ok(())
}
pub async fn send(&mut self, addr: &SocketAddr, message: &Message) -> Result<()> {
let connection = self
.connections
.get(addr)
.ok_or_else(|| Error::network(format!("No connection to {}", addr)))?;
let payload = serde_json::to_vec(message)?;
let mut send_stream = connection
.open_uni()
.await
.map_err(|e| Error::network(format!("Failed to open stream: {}", e)))?;
let len = payload.len() as u32;
send_stream
.write_all(&len.to_be_bytes())
.await
.map_err(|e| Error::network(format!("Failed to write length: {}", e)))?;
send_stream
.write_all(&payload)
.await
.map_err(|e| Error::network(format!("Failed to write payload: {}", e)))?;
send_stream
.finish()
.map_err(|e| Error::network(format!("Failed to finish stream: {}", e)))?;
log::trace!("Sent message to {}: {:?}", addr, message);
Ok(())
}
pub async fn recv(&mut self) -> Result<Option<(SocketAddr, Message)>> {
for (addr, connection) in &self.connections {
match connection.accept_uni().await {
Ok(mut recv_stream) => {
let mut len_buf = [0u8; 4];
if recv_stream.read_exact(&mut len_buf).await.is_err() {
continue;
}
let len = u32::from_be_bytes(len_buf) as usize;
const MAX_MESSAGE_SIZE: usize = 1024 * 1024;
if len > MAX_MESSAGE_SIZE {
log::warn!("Rejecting oversized QUIC message: {} bytes from {}", len, addr);
continue;
}
let mut payload = vec![0u8; len];
if recv_stream.read_exact(&mut payload).await.is_err() {
continue;
}
match serde_json::from_slice::<Message>(&payload) {
Ok(message) => {
log::trace!("Received message from {}: {:?}", addr, message);
return Ok(Some((*addr, message)));
}
Err(e) => {
log::warn!("Failed to deserialize message from {}: {}", addr, e);
}
}
}
Err(e) => {
log::trace!("No incoming stream from {}: {}", addr, e);
continue;
}
}
}
Ok(None)
}
pub fn disconnect(&mut self, addr: &SocketAddr) {
if let Some(connection) = self.connections.remove(addr) {
connection.close(0u32.into(), b"disconnected");
log::debug!("Disconnected from {}", addr);
}
}
pub fn connected_peers(&self) -> Vec<SocketAddr> {
self.connections.keys().copied().collect()
}
pub fn connection_count(&self) -> usize {
self.connections.len()
}
pub fn is_connected(&self, addr: &SocketAddr) -> bool {
self.connections.contains_key(addr)
}
pub fn stop(&mut self) {
for (addr, connection) in self.connections.drain() {
connection.close(0u32.into(), b"server shutdown");
log::debug!("Closed connection to {}", addr);
}
if let Some(endpoint) = self.endpoint.take() {
endpoint.close(0u32.into(), b"server shutdown");
}
log::info!("QUIC server stopped");
}
fn generate_server_config(&self) -> Result<(ServerConfig, CertificateDer<'static>)> {
let cert = rcgen::generate_simple_self_signed(vec!["localhost".into()])
.map_err(|e| Error::crypto(format!("Failed to generate certificate: {}", e)))?;
let cert_der = CertificateDer::from(cert.cert.der().to_vec());
let key_der = PrivateKeyDer::Pkcs8(PrivatePkcs8KeyDer::from(cert.key_pair.serialize_der()));
let mut server_crypto = rustls::ServerConfig::builder()
.with_no_client_auth()
.with_single_cert(vec![cert_der.clone()], key_der)
.map_err(|e| Error::crypto(format!("TLS config error: {}", e)))?;
server_crypto.alpn_protocols = vec![b"aingle".to_vec()];
let mut server_config = ServerConfig::with_crypto(Arc::new(
quinn::crypto::rustls::QuicServerConfig::try_from(server_crypto)
.map_err(|e| Error::crypto(format!("QUIC crypto error: {}", e)))?,
));
let mut transport = quinn::TransportConfig::default();
transport.keep_alive_interval(Some(self.config.keep_alive));
transport.max_idle_timeout(Some(
self.config
.idle_timeout
.try_into()
.map_err(|e| Error::network(format!("Invalid timeout: {}", e)))?,
));
transport.max_concurrent_uni_streams(self.config.max_concurrent_streams.into());
transport.max_concurrent_bidi_streams(self.config.max_concurrent_streams.into());
server_config.transport_config(Arc::new(transport));
Ok((server_config, cert_der))
}
fn generate_client_config(&self) -> Result<ClientConfig> {
let mut root_store = rustls::RootCertStore::empty();
log::warn!("QUIC client using permissive certificate validation — pin peer certs in production");
let crypto = rustls::ClientConfig::builder()
.dangerous()
.with_custom_certificate_verifier(Arc::new(LoggingCertVerifier))
.with_no_client_auth();
let mut client_config = ClientConfig::new(Arc::new(
quinn::crypto::rustls::QuicClientConfig::try_from(crypto)
.map_err(|e| Error::crypto(format!("QUIC crypto error: {}", e)))?,
));
let mut transport = quinn::TransportConfig::default();
transport.keep_alive_interval(Some(self.config.keep_alive));
transport.max_idle_timeout(Some(
self.config
.idle_timeout
.try_into()
.map_err(|e| Error::network(format!("Invalid timeout: {}", e)))?,
));
client_config.transport_config(Arc::new(transport));
Ok(client_config)
}
}
#[derive(Debug)]
struct LoggingCertVerifier;
impl rustls::client::danger::ServerCertVerifier for LoggingCertVerifier {
fn verify_server_cert(
&self,
end_entity: &CertificateDer<'_>,
_intermediates: &[CertificateDer<'_>],
server_name: &rustls::pki_types::ServerName<'_>,
_ocsp_response: &[u8],
_now: rustls::pki_types::UnixTime,
) -> std::result::Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
let fingerprint = blake3::hash(end_entity.as_ref());
log::info!(
"QUIC peer cert fingerprint for {:?}: {}",
server_name,
hex::encode(fingerprint.as_bytes())
);
Ok(rustls::client::danger::ServerCertVerified::assertion())
}
fn verify_tls12_signature(
&self,
message: &[u8],
cert: &CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> std::result::Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
rustls::crypto::verify_tls12_signature(
message,
cert,
dss,
&rustls::crypto::ring::default_provider().signature_verification_algorithms,
)
}
fn verify_tls13_signature(
&self,
message: &[u8],
cert: &CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> std::result::Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
rustls::crypto::verify_tls13_signature(
message,
cert,
dss,
&rustls::crypto::ring::default_provider().signature_verification_algorithms,
)
}
fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
rustls::crypto::ring::default_provider()
.signature_verification_algorithms
.supported_schemes()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_quic_config_default() {
let config = QuicConfig::default();
assert_eq!(config.port, 19081);
assert_eq!(config.max_concurrent_streams, 100);
}
#[test]
fn test_quic_config_iot_mode() {
let config = QuicConfig::iot_mode();
assert_eq!(config.max_concurrent_streams, 10);
assert_eq!(config.max_connections, 100);
}
#[test]
fn test_quic_config_production() {
let config = QuicConfig::production();
assert_eq!(config.max_concurrent_streams, 1000);
assert_eq!(config.max_connections, 10000);
}
#[test]
fn test_quic_server_new() {
let config = QuicConfig::default();
let server = QuicServer::new(config, "test-node".to_string());
assert_eq!(server.connection_count(), 0);
assert!(server.connected_peers().is_empty());
}
#[test]
fn test_quic_server_is_connected() {
let config = QuicConfig::default();
let server = QuicServer::new(config, "test-node".to_string());
let addr: SocketAddr = "127.0.0.1:19081".parse().unwrap();
assert!(!server.is_connected(&addr));
}
}