1use std::{
8 collections::HashMap,
9 io,
10 net::{IpAddr, SocketAddr},
11 sync::{
12 atomic::{AtomicBool, Ordering},
13 Arc,
14 },
15 time::Duration,
16};
17
18use crate::connection::Connection;
19use nxtquic_crypto::{
20 CryptoProvider, DirectionalKeys, HandshakeState, Keys,
21 RustlsCryptoProvider,
22};
23use nxtquic_proto::{
24 frame::{AckFrame, CryptoFrame, Frame, StreamFrame},
25 packet::{
26 decode_header,
27 initial::parse_initial,
28 protection::{apply_header_protection, remove_header_protection},
29 PacketType,
30 },
31 varint::VarInt,
32 ConnectionId, StreamId, Version,
33};
34use tokio::sync::{Mutex, Notify};
35
36#[derive(Default, Clone, Debug)]
38pub struct EndpointConfig;
39
40#[derive(Clone)]
42pub struct ServerConfig {
43 rustls: Arc<rustls::ServerConfig>,
44 transport: TransportConfig,
45 max_concurrent_connections: Option<usize>,
46 migration_enabled: bool,
47 preferred_address: Option<SocketAddr>,
48}
49
50impl std::fmt::Debug for ServerConfig {
51 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
52 formatter
53 .debug_struct("ServerConfig")
54 .field("transport", &self.transport)
55 .field(
56 "max_concurrent_connections",
57 &self.max_concurrent_connections,
58 )
59 .field("migration_enabled", &self.migration_enabled)
60 .field("preferred_address", &self.preferred_address)
61 .finish_non_exhaustive()
62 }
63}
64
65impl ServerConfig {
66 pub fn builder() -> ServerConfigBuilder {
68 ServerConfigBuilder::default()
69 }
70
71 pub fn rustls(&self) -> &Arc<rustls::ServerConfig> {
73 &self.rustls
74 }
75
76 pub fn transport(&self) -> &TransportConfig {
78 &self.transport
79 }
80
81 pub fn max_concurrent_connections(&self) -> Option<usize> {
83 self.max_concurrent_connections
84 }
85
86 pub fn migration_enabled(&self) -> bool {
88 self.migration_enabled
89 }
90
91 pub fn preferred_address(&self) -> Option<SocketAddr> {
93 self.preferred_address
94 }
95}
96
97#[derive(Default)]
99pub struct ServerConfigBuilder {
100 rustls: Option<Arc<rustls::ServerConfig>>,
101 transport: TransportConfig,
102 max_concurrent_connections: Option<usize>,
103 migration_enabled: bool,
104 preferred_address: Option<SocketAddr>,
105}
106
107impl ServerConfigBuilder {
108 pub fn with_rustls(mut self, config: Arc<rustls::ServerConfig>) -> io::Result<Self> {
110 if config
111 .alpn_protocols
112 .iter()
113 .all(|protocol| protocol.as_slice() != b"h3")
114 {
115 return Err(io::Error::new(
116 io::ErrorKind::InvalidInput,
117 "a QUIC HTTP/3 server configuration must advertise ALPN h3",
118 ));
119 }
120 self.rustls = Some(config);
121 Ok(self)
122 }
123
124 pub fn with_transport_config(mut self, configure: impl FnOnce(&mut TransportConfig)) -> Self {
126 configure(&mut self.transport);
127 self
128 }
129
130 pub fn with_max_concurrent_connections(mut self, max: usize) -> Self {
132 self.max_concurrent_connections = Some(max);
133 self
134 }
135
136 pub fn with_migration(mut self, enable: bool) -> Self {
138 self.migration_enabled = enable;
139 self
140 }
141
142 pub fn with_preferred_address(mut self, addr: SocketAddr) -> Self {
144 self.preferred_address = Some(addr);
145 self
146 }
147
148 pub fn build(self) -> io::Result<ServerConfig> {
150 let rustls = self.rustls.ok_or_else(|| {
151 io::Error::new(
152 io::ErrorKind::InvalidInput,
153 "missing Rustls server configuration",
154 )
155 })?;
156 Ok(ServerConfig {
157 rustls,
158 transport: self.transport,
159 max_concurrent_connections: self.max_concurrent_connections,
160 migration_enabled: self.migration_enabled,
161 preferred_address: self.preferred_address,
162 })
163 }
164}
165
166#[derive(Clone, Debug)]
168pub struct TransportConfig {
169 max_idle_timeout: Option<Duration>,
170 keep_alive_interval: Option<Duration>,
171 datagram_receive_buffer_size: usize,
172 initial_max_data: u64,
173 initial_max_stream_data_bidi_local: u64,
174 initial_max_stream_data_bidi_remote: u64,
175 initial_max_stream_data_uni: u64,
176 initial_max_streams_bidi: u64,
177 initial_max_streams_uni: u64,
178 max_udp_payload_size: u64,
179 ack_delay_exponent: u8,
180 max_ack_delay: Duration,
181 disable_active_migration: bool,
182 active_connection_id_limit: u64,
183 max_concurrent_uni_streams: u64,
184 max_concurrent_bi_streams: u64,
185}
186
187impl Default for TransportConfig {
188 fn default() -> Self {
189 Self {
190 max_idle_timeout: Some(Duration::from_secs(30)),
191 keep_alive_interval: Some(Duration::from_secs(10)),
192 datagram_receive_buffer_size: 1024 * 1024,
193 initial_max_data: 10 * 1024 * 1024,
194 initial_max_stream_data_bidi_local: 1024 * 1024,
195 initial_max_stream_data_bidi_remote: 1024 * 1024,
196 initial_max_stream_data_uni: 1024 * 1024,
197 initial_max_streams_bidi: 100,
198 initial_max_streams_uni: 100,
199 max_udp_payload_size: 1350,
200 ack_delay_exponent: 3,
201 max_ack_delay: Duration::from_millis(25),
202 disable_active_migration: false,
203 active_connection_id_limit: 4,
204 max_concurrent_uni_streams: 100,
205 max_concurrent_bi_streams: 100,
206 }
207 }
208}
209
210impl TransportConfig {
211 pub fn max_idle_timeout(&mut self, timeout: Option<Duration>) -> &mut Self {
213 self.max_idle_timeout = timeout;
214 self
215 }
216
217 pub fn idle_timeout(&self) -> Option<Duration> {
219 self.max_idle_timeout
220 }
221
222 pub fn keep_alive_interval(&mut self, interval: Option<Duration>) -> &mut Self {
224 self.keep_alive_interval = interval;
225 self
226 }
227
228 pub fn keep_alive(&self) -> Option<Duration> {
230 self.keep_alive_interval
231 }
232
233 pub fn datagram_receive_buffer_size(&mut self, size: usize) -> &mut Self {
235 self.datagram_receive_buffer_size = size;
236 self
237 }
238
239 pub fn datagram_buffer_size(&self) -> usize {
241 self.datagram_receive_buffer_size
242 }
243
244 pub fn initial_max_data(&mut self, value: u64) -> &mut Self {
246 self.initial_max_data = value;
247 self
248 }
249
250 pub fn get_initial_max_data(&self) -> u64 {
252 self.initial_max_data
253 }
254
255 pub fn initial_max_stream_data_bidi_local(&mut self, value: u64) -> &mut Self {
257 self.initial_max_stream_data_bidi_local = value;
258 self
259 }
260
261 pub fn initial_max_stream_data_bidi_remote(&mut self, value: u64) -> &mut Self {
263 self.initial_max_stream_data_bidi_remote = value;
264 self
265 }
266
267 pub fn initial_max_stream_data_uni(&mut self, value: u64) -> &mut Self {
269 self.initial_max_stream_data_uni = value;
270 self
271 }
272
273 pub fn initial_max_streams_bidi(&mut self, value: u64) -> &mut Self {
275 self.initial_max_streams_bidi = value;
276 self
277 }
278
279 pub fn initial_max_streams_uni(&mut self, value: u64) -> &mut Self {
281 self.initial_max_streams_uni = value;
282 self
283 }
284
285 pub fn max_udp_payload_size(&mut self, value: u64) -> &mut Self {
287 self.max_udp_payload_size = value;
288 self
289 }
290
291 pub fn ack_delay_exponent(&mut self, value: u8) -> &mut Self {
293 self.ack_delay_exponent = value;
294 self
295 }
296
297 pub fn max_ack_delay(&mut self, value: Duration) -> &mut Self {
299 self.max_ack_delay = value;
300 self
301 }
302
303 pub fn disable_active_migration(&mut self, disable: bool) -> &mut Self {
305 self.disable_active_migration = disable;
306 self
307 }
308
309 pub fn active_connection_id_limit(&mut self, limit: u64) -> &mut Self {
311 self.active_connection_id_limit = limit;
312 self
313 }
314
315 pub fn max_concurrent_uni_streams(&mut self, count: u64) -> &mut Self {
317 self.max_concurrent_uni_streams = count;
318 self
319 }
320
321 pub fn max_concurrent_bi_streams(&mut self, count: u64) -> &mut Self {
323 self.max_concurrent_bi_streams = count;
324 self
325 }
326}
327
328#[derive(Clone)]
330pub struct ClientConfig {
331 rustls: Option<Arc<rustls::ClientConfig>>,
332 transport: TransportConfig,
333}
334
335impl Default for ClientConfig {
336 fn default() -> Self {
337 Self {
338 rustls: None,
339 transport: TransportConfig::default(),
340 }
341 }
342}
343
344impl std::fmt::Debug for ClientConfig {
345 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
346 f.debug_struct("ClientConfig")
347 .field("transport", &self.transport)
348 .finish_non_exhaustive()
349 }
350}
351
352impl ClientConfig {
353 pub fn builder() -> ClientConfigBuilder {
355 ClientConfigBuilder::default()
356 }
357
358 pub fn rustls(&self) -> Option<&Arc<rustls::ClientConfig>> {
360 self.rustls.as_ref()
361 }
362
363 pub fn transport(&self) -> &TransportConfig {
365 &self.transport
366 }
367}
368
369#[derive(Default)]
371pub struct ClientConfigBuilder {
372 rustls: Option<Arc<rustls::ClientConfig>>,
373 transport: TransportConfig,
374}
375
376impl ClientConfigBuilder {
377 pub fn with_rustls(mut self, config: Arc<rustls::ClientConfig>) -> Self {
379 self.rustls = Some(config);
380 self
381 }
382
383 pub fn with_transport_config(mut self, configure: impl FnOnce(&mut TransportConfig)) -> Self {
385 configure(&mut self.transport);
386 self
387 }
388
389 pub fn with_native_roots(mut self) -> Self {
391 let root_store = rustls::RootCertStore::empty();
392 let mut config = rustls::ClientConfig::builder()
393 .with_root_certificates(root_store)
394 .with_no_client_auth();
395 config.alpn_protocols = vec![b"h3".to_vec()];
396 self.rustls = Some(Arc::new(config));
397 self
398 }
399
400 pub fn build(self) -> ClientConfig {
402 ClientConfig {
403 rustls: self.rustls,
404 transport: self.transport,
405 }
406 }
407}
408
409#[derive(Clone, Debug, Default)]
411pub struct EndpointStats {
412 pub total_connections: u64,
414 pub active_connections: u64,
416 pub handshake_failures: u64,
418 pub packets_received: u64,
420 pub packets_sent: u64,
422}
423
424pub struct Endpoint {
426 _config: EndpointConfig,
427 server_config: Arc<Mutex<Option<ServerConfig>>>,
428 socket: Arc<tokio::net::UdpSocket>,
429 incoming: Arc<Mutex<tokio::sync::mpsc::Receiver<Incoming>>>,
430 closed: Arc<AtomicBool>,
431 stats: Arc<Mutex<EndpointStats>>,
432 close_notify: Arc<Notify>,
433}
434
435pub struct Connecting;
437
438pub struct Incoming {
440 remote_addr: SocketAddr,
441 packet: Vec<u8>,
442 initial: DecryptedInitial,
443 socket: Arc<tokio::net::UdpSocket>,
444 server_config: Option<ServerConfig>,
445 datagrams: Option<tokio::sync::mpsc::Receiver<Vec<u8>>>,
446}
447
448#[derive(Debug)]
450struct DecryptedInitial {
451 version: Version,
452 dcid: ConnectionId,
453 scid: ConnectionId,
454 _packet_number: u64,
455 frames: Vec<Frame>,
456}
457
458fn decrypt_client_initial(packet: &[u8]) -> io::Result<DecryptedInitial> {
459 let parsed =
460 parse_initial(packet).map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
461 let version = match parsed.version {
462 1 => Version::V1,
463 0x6b33_43cf => Version::V2,
464 version => {
465 return Err(io::Error::new(
466 io::ErrorKind::InvalidData,
467 format!("unsupported QUIC version {version:#010x}"),
468 ));
469 }
470 };
471
472 let provider = RustlsCryptoProvider::new(None, None);
473 let keys = provider.initial_keys(version, &parsed.dcid);
474 let mut protected = packet[..parsed.packet_end].to_vec();
475
476 let sample_offset = parsed.packet_number_offset + 4;
477 let sample: [u8; 16] = protected[sample_offset..sample_offset + 16]
478 .try_into()
479 .expect("Initial parser validates the header-protection sample");
480 let mask = keys
481 .client
482 .header_key
483 .protection_mask(&sample)
484 .map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
485 let packet_number_len =
486 remove_header_protection(&mut protected, &mask, parsed.packet_number_offset);
487 if packet_number_len == 0
488 || parsed.packet_number_offset + packet_number_len > parsed.packet_end
489 {
490 return Err(io::Error::new(
491 io::ErrorKind::InvalidData,
492 "invalid QUIC Initial packet number",
493 ));
494 }
495
496 let mut packet_number = 0_u64;
497 for byte in
498 &protected[parsed.packet_number_offset..parsed.packet_number_offset + packet_number_len]
499 {
500 packet_number = (packet_number << 8) | u64::from(*byte);
501 }
502 let header_end = parsed.packet_number_offset + packet_number_len;
503 let mut payload = protected[header_end..parsed.packet_end].to_vec();
504 keys.client
505 .packet_key
506 .decrypt(packet_number, &protected[..header_end], &mut payload)
507 .map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
508
509 let mut frames = Vec::new();
510 let mut payload = &payload[..];
511 while !payload.is_empty() {
512 let before = payload.len();
513 let frame = Frame::decode(&mut payload)
514 .map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
515 frames.push(frame);
516 if payload.len() == before {
517 return Err(io::Error::new(
518 io::ErrorKind::InvalidData,
519 "QUIC frame decoder made no progress",
520 ));
521 }
522 }
523
524 Ok(DecryptedInitial {
525 version,
526 dcid: parsed.dcid,
527 scid: parsed.scid,
528 _packet_number: packet_number,
529 frames,
530 })
531}
532
533fn encode_server_initial(
534 version: Version,
535 client_scid: ConnectionId,
536 server_scid: ConnectionId,
537 packet_number: u64,
538 crypto_data: Vec<u8>,
539 keys: nxtquic_crypto::InitialKeys,
540) -> io::Result<Vec<u8>> {
541 let mut plaintext = Vec::new();
542 Frame::Crypto(CryptoFrame {
543 offset: VarInt::ZERO,
544 data: crypto_data.into(),
545 })
546 .encode(&mut plaintext);
547
548 let mut header_prefix = Vec::new();
549 header_prefix.push(0xc0); header_prefix.extend_from_slice(&version.as_u32().to_be_bytes());
551 header_prefix.push(client_scid.len() as u8);
552 header_prefix.extend_from_slice(client_scid.as_bytes());
553 header_prefix.push(server_scid.len() as u8);
554 header_prefix.extend_from_slice(server_scid.as_bytes());
555 header_prefix.push(0); loop {
558 let encrypted_len = plaintext.len() + keys.server.packet_key.tag_len();
559 let mut candidate = header_prefix.clone();
560 VarInt::from_u64((1 + encrypted_len) as u64)
561 .expect("QUIC packet lengths fit in a VarInt")
562 .encode(&mut candidate);
563 if candidate.len() + 1 + encrypted_len >= 1200 {
564 header_prefix = candidate;
565 break;
566 }
567 plaintext.push(0);
568 }
569
570 let packet_number_len = 1;
571 let packet_number_offset = header_prefix.len();
572 header_prefix.push(packet_number as u8);
573 keys.server
574 .packet_key
575 .encrypt(packet_number, &header_prefix, &mut plaintext)
576 .map_err(|error| io::Error::new(io::ErrorKind::Other, error))?;
577 header_prefix.extend_from_slice(&plaintext);
578
579 let sample: [u8; 16] = header_prefix[packet_number_offset + 4..packet_number_offset + 20]
580 .try_into()
581 .expect("1200-byte Initial always contains a header-protection sample");
582 let mask = keys
583 .server
584 .header_key
585 .protection_mask(&sample)
586 .map_err(|error| io::Error::new(io::ErrorKind::Other, error))?;
587 apply_header_protection(
588 &mut header_prefix,
589 &mask,
590 packet_number_offset,
591 packet_number_len,
592 );
593 Ok(header_prefix)
594}
595
596impl std::fmt::Debug for Incoming {
597 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
598 formatter
599 .debug_struct("Incoming")
600 .field("remote_addr", &self.remote_addr)
601 .field("packet_len", &self.packet.len())
602 .finish()
603 }
604}
605
606impl Incoming {
607 pub fn remote_address(&self) -> SocketAddr {
609 self.remote_addr
610 }
611
612 pub fn local_ip(&self) -> Option<IpAddr> {
614 self.socket.local_addr().ok().map(|a| a.ip())
615 }
616
617 pub fn cancel(&mut self) {
619 self.datagrams = None;
620 }
621
622 pub async fn reject(&mut self, error_code: u64, reason: &[u8]) -> io::Result<()> {
624 let _ = error_code;
625 let _ = reason;
626 self.cancel();
627 Ok(())
628 }
629
630 pub(crate) fn packet(&self) -> &[u8] {
632 &self.packet
633 }
634}
635
636impl Endpoint {
637 pub async fn bind(addr: SocketAddr) -> io::Result<Self> {
639 Self::bind_inner(EndpointConfig::default(), None, addr).await
640 }
641
642 pub async fn server(server_config: ServerConfig, addr: SocketAddr) -> io::Result<Self> {
644 Self::bind_inner(EndpointConfig::default(), Some(server_config), addr).await
645 }
646
647 async fn bind_inner(
648 config: EndpointConfig,
649 server_config: Option<ServerConfig>,
650 addr: SocketAddr,
651 ) -> io::Result<Self> {
652 let socket = tokio::net::UdpSocket::bind(addr).await?;
653 let socket = Arc::new(socket);
654 let (incoming_tx, incoming_rx) = tokio::sync::mpsc::channel(128);
655 let receive_socket = Arc::clone(&socket);
656 let shared_server_config = Arc::new(Mutex::new(server_config));
657 let receive_server_config = Arc::clone(&shared_server_config);
658 let stats = Arc::new(Mutex::new(EndpointStats::default()));
659 let receive_stats = Arc::clone(&stats);
660
661 tokio::spawn(async move {
662 let mut packet = vec![0_u8; 65_535];
663 let mut peers: HashMap<SocketAddr, tokio::sync::mpsc::Sender<Vec<u8>>> = HashMap::new();
664 loop {
665 let (length, remote_addr) = match receive_socket.recv_from(&mut packet).await {
666 Ok(received) => received,
667 Err(_) => break,
668 };
669 {
670 let mut st = receive_stats.lock().await;
671 st.packets_received += 1;
672 }
673
674 let datagram = &packet[..length];
675 if let Some(peer_tx) = peers.get(&remote_addr) {
676 if peer_tx.send(datagram.to_vec()).await.is_err() {
677 peers.remove(&remote_addr);
678 }
679 continue;
680 }
681 let Ok(header) = decode_header(datagram) else {
682 continue;
683 };
684 if header.packet_type != PacketType::Initial {
685 continue;
686 }
687 let Ok(initial) = decrypt_client_initial(datagram) else {
688 let mut st = receive_stats.lock().await;
689 st.handshake_failures += 1;
690 continue;
691 };
692
693 let (peer_tx, peer_rx) = tokio::sync::mpsc::channel(256);
694 peers.insert(remote_addr, peer_tx);
695
696 let cfg = receive_server_config.lock().await.clone();
697 {
698 let mut st = receive_stats.lock().await;
699 st.total_connections += 1;
700 st.active_connections += 1;
701 }
702
703 if incoming_tx
704 .send(Incoming {
705 remote_addr,
706 packet: datagram.to_vec(),
707 initial,
708 socket: Arc::clone(&receive_socket),
709 server_config: cfg,
710 datagrams: Some(peer_rx),
711 })
712 .await
713 .is_err()
714 {
715 break;
716 }
717 }
718 });
719
720 Ok(Self {
721 _config: config,
722 server_config: shared_server_config,
723 socket,
724 incoming: Arc::new(Mutex::new(incoming_rx)),
725 closed: Arc::new(AtomicBool::new(false)),
726 stats,
727 close_notify: Arc::new(Notify::new()),
728 })
729 }
730
731 pub fn local_addr(&self) -> io::Result<SocketAddr> {
733 self.socket.local_addr()
734 }
735
736 pub async fn server_config(&self) -> Option<ServerConfig> {
738 self.server_config.lock().await.clone()
739 }
740
741 pub async fn set_server_config(&self, config: ServerConfig) {
743 *self.server_config.lock().await = Some(config);
744 }
745
746 pub async fn stats(&self) -> EndpointStats {
748 self.stats.lock().await.clone()
749 }
750
751 pub async fn close(&self, _code: u64, _reason: &[u8]) {
753 self.closed.store(true, Ordering::Release);
754 self.close_notify.notify_waiters();
755 }
756
757 pub async fn wait_idle(&self) {
759 loop {
760 let active = self.stats.lock().await.active_connections;
761 if active == 0 {
762 break;
763 }
764 tokio::time::sleep(Duration::from_millis(50)).await;
765 }
766 }
767
768 pub async fn connect(
770 &self,
771 _addr: SocketAddr,
772 _server_name: &str,
773 ) -> std::io::Result<Connecting> {
774 Ok(Connecting)
775 }
776
777 pub async fn connect_with_config(
779 &self,
780 _addr: SocketAddr,
781 _server_name: &str,
782 _config: ClientConfig,
783 ) -> std::io::Result<Connecting> {
784 Ok(Connecting)
785 }
786
787 pub async fn connect_0rtt(
789 &self,
790 _addr: SocketAddr,
791 _server_name: &str,
792 ) -> std::io::Result<Connecting> {
793 Ok(Connecting)
794 }
795
796 pub async fn accept(&self) -> Option<Incoming> {
798 self.incoming.lock().await.recv().await
799 }
800}
801
802impl Connecting {
803 pub async fn await_connection(self) -> std::io::Result<Connection> {
805 Ok(Connection::new())
806 }
807
808 pub async fn into_0rtt(self) -> std::io::Result<(Connection, Connection)> {
810 let conn = Connection::new();
811 Ok((conn.clone(), conn))
812 }
813}
814
815impl Incoming {
816 pub async fn accept(self) -> std::io::Result<Connection> {
818 let config = self.server_config.ok_or_else(|| {
819 io::Error::new(
820 io::ErrorKind::InvalidInput,
821 "cannot accept a connection on an endpoint without ServerConfig",
822 )
823 })?;
824 let mut crypto = Vec::new();
825 for frame in &self.initial.frames {
826 if let Frame::Crypto(frame) = frame {
827 if frame.offset.into_inner() != crypto.len() as u64 {
828 return Err(io::Error::new(
829 io::ErrorKind::InvalidData,
830 "Initial CRYPTO frames must be contiguous",
831 ));
832 }
833 crypto.extend_from_slice(&frame.data);
834 }
835 }
836 if crypto.is_empty() {
837 return Err(io::Error::new(
838 io::ErrorKind::InvalidData,
839 "Initial packet contains no CRYPTO frame",
840 ));
841 }
842
843 let provider = RustlsCryptoProvider::new(None, Some(Arc::clone(config.rustls())));
844
845 let mut transport_params = nxtquic_proto::TransportParameters::default();
846 transport_params.initial_max_data = VarInt::from_u64(10 * 1024 * 1024).unwrap();
847 transport_params.initial_max_stream_data_bidi_local =
848 VarInt::from_u64(1024 * 1024).unwrap();
849 transport_params.initial_max_stream_data_bidi_remote =
850 VarInt::from_u64(1024 * 1024).unwrap();
851 transport_params.initial_max_stream_data_uni = VarInt::from_u64(1024 * 1024).unwrap();
852 transport_params.initial_max_streams_bidi = VarInt::from_u64(100).unwrap();
853 transport_params.initial_max_streams_uni = VarInt::from_u64(100).unwrap();
854 if let Some(timeout) = config.transport().idle_timeout() {
855 transport_params.max_idle_timeout =
856 VarInt::from_u64(timeout.as_millis() as u64).unwrap_or(VarInt::from_u32(0));
857 }
858 transport_params.max_udp_payload_size = VarInt::from_u64(1350).unwrap();
859 transport_params.active_connection_id_limit = VarInt::from_u64(2).unwrap();
860
861 let server_scid = ConnectionId::from_slice(&rand::random::<[u8; 16]>());
862 transport_params.initial_source_connection_id = Some(server_scid.clone());
863 transport_params.original_destination_connection_id = Some(self.initial.dcid.clone());
864
865 let mut tp_encoded = Vec::new();
866 transport_params.encode(&mut tp_encoded);
867
868 let mut handshake = provider
869 .new_server_handshake(tp_encoded)
870 .map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
871 let output = handshake
872 .process(&crypto)
873 .map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
874 if output.data.is_empty() {
875 return Err(io::Error::new(
876 io::ErrorKind::InvalidData,
877 "TLS produced no Initial handshake data",
878 ));
879 }
880
881 let keys = provider.initial_keys(self.initial.version, &self.initial.dcid);
882 let response = encode_server_initial(
883 self.initial.version,
884 self.initial.scid,
885 server_scid,
886 0,
887 output.data,
888 keys,
889 )?;
890 self.socket.send_to(&response, self.remote_addr).await?;
891 let (outgoing_tx, outgoing_rx) = tokio::sync::mpsc::unbounded_channel();
892 let connection = Connection::with_outgoing(outgoing_tx);
893 if let Some(datagrams) = self.datagrams {
894 let driver_connection = connection.clone();
895 let driver_initial = self.initial;
896 let driver_socket = Arc::clone(&self.socket);
897 let driver_remote = self.remote_addr;
898 let driver_initial_keys =
899 provider.initial_keys(driver_initial.version, &driver_initial.dcid);
900 let mut driver_handshake_keys = None;
901 let mut driver_one_rtt_keys = None;
902 while let Some((level, keys)) = handshake.next_keys() {
903 if level == nxtquic_crypto::Level::OneRtt {
904 driver_one_rtt_keys = Some(keys);
905 } else {
906 driver_handshake_keys = Some((level, keys));
907 }
908 }
909 tokio::spawn(async move {
910 drive_server_connection(
911 driver_connection,
912 driver_socket,
913 driver_remote,
914 datagrams,
915 driver_initial,
916 server_scid,
917 driver_initial_keys,
918 driver_handshake_keys,
919 driver_one_rtt_keys,
920 handshake,
921 outgoing_rx,
922 )
923 .await;
924 });
925 }
926 Ok(connection)
927 }
928}
929
930struct LongPacket {
931 _packet_type: PacketType,
932 _dcid: ConnectionId,
933 _scid: ConnectionId,
934 pn_offset: usize,
935 packet_end: usize,
936}
937
938fn parse_long_packet(datagram: &[u8]) -> io::Result<LongPacket> {
939 if datagram.len() < 7 || datagram[0] & 0xc0 != 0xc0 {
940 return Err(io::Error::new(
941 io::ErrorKind::InvalidData,
942 "invalid QUIC long header",
943 ));
944 }
945 let version = u32::from_be_bytes(datagram[1..5].try_into().unwrap());
946 if version == 0 {
947 return Err(io::Error::new(
948 io::ErrorKind::InvalidData,
949 "version negotiation is not a protected packet",
950 ));
951 }
952 let packet_type = match (datagram[0] >> 4) & 0x03 {
953 0 => PacketType::Initial,
954 1 => PacketType::ZeroRtt,
955 2 => PacketType::Handshake,
956 _ => PacketType::Retry,
957 };
958 let mut offset = 5;
959 let dcid_len = datagram[offset] as usize;
960 offset += 1;
961 let dcid_end = offset
962 .checked_add(dcid_len)
963 .ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "invalid DCID"))?;
964 let dcid = ConnectionId::from_slice(
965 datagram
966 .get(offset..dcid_end)
967 .ok_or_else(|| io::Error::new(io::ErrorKind::UnexpectedEof, "truncated DCID"))?,
968 );
969 offset = dcid_end;
970 let scid_len = *datagram
971 .get(offset)
972 .ok_or_else(|| io::Error::new(io::ErrorKind::UnexpectedEof, "missing SCID length"))?
973 as usize;
974 offset += 1;
975 let scid_end = offset
976 .checked_add(scid_len)
977 .ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "invalid SCID"))?;
978 let scid = ConnectionId::from_slice(
979 datagram
980 .get(offset..scid_end)
981 .ok_or_else(|| io::Error::new(io::ErrorKind::UnexpectedEof, "truncated SCID"))?,
982 );
983 offset = scid_end;
984 if packet_type == PacketType::Initial {
985 let mut rest = &datagram[offset..];
986 let token_len = VarInt::decode(&mut rest)
987 .map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "invalid Initial token length"))?
988 .into_inner() as usize;
989 let consumed = datagram[offset..].len() - rest.len();
990 offset += consumed
991 .checked_add(token_len)
992 .ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "invalid token"))?;
993 }
994 let mut rest = &datagram[offset..];
995 let payload_len = VarInt::decode(&mut rest)
996 .map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "invalid QUIC payload length"))?
997 .into_inner() as usize;
998 offset += datagram[offset..].len() - rest.len();
999 let packet_end = offset
1000 .checked_add(payload_len)
1001 .ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "invalid packet length"))?;
1002 if packet_end > datagram.len() || packet_end < offset + 4 + 16 {
1003 return Err(io::Error::new(
1004 io::ErrorKind::UnexpectedEof,
1005 "truncated protected QUIC packet",
1006 ));
1007 }
1008 Ok(LongPacket {
1009 _packet_type: packet_type,
1010 _dcid: dcid,
1011 _scid: scid,
1012 pn_offset: offset,
1013 packet_end,
1014 })
1015}
1016
1017fn decrypt_payload(
1018 mut packet: Vec<u8>,
1019 parsed: &LongPacket,
1020 keys: &DirectionalKeys,
1021) -> io::Result<(u64, Vec<u8>)> {
1022 let sample: [u8; 16] = packet[parsed.pn_offset + 4..parsed.pn_offset + 20]
1023 .try_into()
1024 .map_err(|_| {
1025 io::Error::new(
1026 io::ErrorKind::InvalidData,
1027 "missing header-protection sample",
1028 )
1029 })?;
1030 let mask = keys
1031 .header_key
1032 .protection_mask(&sample)
1033 .map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
1034 let pn_len = remove_header_protection(&mut packet, &mask, parsed.pn_offset);
1035 if pn_len == 0 || parsed.pn_offset + pn_len > parsed.packet_end {
1036 return Err(io::Error::new(
1037 io::ErrorKind::InvalidData,
1038 "invalid packet number",
1039 ));
1040 }
1041 let mut packet_number = 0;
1042 for byte in &packet[parsed.pn_offset..parsed.pn_offset + pn_len] {
1043 packet_number = (packet_number << 8) | u64::from(*byte);
1044 }
1045 let header_end = parsed.pn_offset + pn_len;
1046 let mut payload = packet[header_end..parsed.packet_end].to_vec();
1047 keys.packet_key
1048 .decrypt(packet_number, &packet[..header_end], &mut payload)
1049 .map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
1050 Ok((packet_number, payload))
1051}
1052
1053async fn drive_server_connection(
1054 connection: Connection,
1055 socket: Arc<tokio::net::UdpSocket>,
1056 remote_addr: SocketAddr,
1057 mut datagrams: tokio::sync::mpsc::Receiver<Vec<u8>>,
1058 initial: DecryptedInitial,
1059 server_scid: ConnectionId,
1060 initial_keys: nxtquic_crypto::InitialKeys,
1061 mut handshake_keys: Option<(nxtquic_crypto::Level, Keys)>,
1062 mut one_rtt_keys: Option<Keys>,
1063 mut handshake: impl HandshakeState,
1064 mut outgoing: tokio::sync::mpsc::UnboundedReceiver<crate::stream::WriteCommand>,
1065) {
1066 let mut largest_handshake = 0;
1067 let mut application_packet_number = 0;
1068 loop {
1069 let Some(event) = (tokio::select! {
1070 datagram = datagrams.recv() => datagram.map(Ok),
1071 command = outgoing.recv() => command.map(Err),
1072 else => None,
1073 }) else {
1074 break;
1075 };
1076 let datagram = match event {
1077 Err(command) => {
1078 if let Some(keys) = one_rtt_keys.as_ref() {
1079 match command {
1080 crate::stream::WriteCommand::Data {
1081 stream_id,
1082 offset,
1083 data,
1084 fin,
1085 } => {
1086 let mut encoded = Vec::new();
1087 Frame::Stream(StreamFrame {
1088 stream_id: StreamId::from_u64(stream_id),
1089 offset: VarInt::from_u64(offset).unwrap(),
1090 length: Some(VarInt::from_u64(data.len() as u64).unwrap()),
1091 fin,
1092 data: bytes::Bytes::from(data),
1093 })
1094 .encode(&mut encoded);
1095 let _ = send_protected_short(
1096 &socket,
1097 remote_addr,
1098 initial.scid,
1099 application_packet_number,
1100 &encoded,
1101 &keys.local,
1102 )
1103 .await;
1104 application_packet_number += 1;
1105 }
1106 crate::stream::WriteCommand::Reset {
1107 stream_id,
1108 error_code,
1109 } => {
1110 let mut encoded = Vec::new();
1111 Frame::ResetStream(nxtquic_proto::frame::ResetStreamFrame {
1112 stream_id: StreamId::from_u64(stream_id),
1113 application_protocol_error_code: VarInt::from_u64(error_code).unwrap(),
1114 final_size: VarInt::ZERO,
1115 })
1116 .encode(&mut encoded);
1117 let _ = send_protected_short(
1118 &socket,
1119 remote_addr,
1120 initial.scid,
1121 application_packet_number,
1122 &encoded,
1123 &keys.local,
1124 )
1125 .await;
1126 application_packet_number += 1;
1127 }
1128 crate::stream::WriteCommand::StopSending {
1129 stream_id,
1130 error_code,
1131 } => {
1132 let mut encoded = Vec::new();
1133 Frame::StopSending(nxtquic_proto::frame::StopSendingFrame {
1134 stream_id: StreamId::from_u64(stream_id),
1135 application_protocol_error_code: VarInt::from_u64(error_code).unwrap(),
1136 })
1137 .encode(&mut encoded);
1138 let _ = send_protected_short(
1139 &socket,
1140 remote_addr,
1141 initial.scid,
1142 application_packet_number,
1143 &encoded,
1144 &keys.local,
1145 )
1146 .await;
1147 application_packet_number += 1;
1148 }
1149 crate::stream::WriteCommand::Priority { .. } => {}
1150 }
1151 }
1152 continue;
1153 }
1154 Ok(datagram) => datagram,
1155 };
1156 let Ok(header) = decode_header(&datagram) else {
1157 continue;
1158 };
1159 let (payload, packet_number, level) = match header.packet_type {
1160 PacketType::Initial => {
1161 let Ok(parsed) = parse_long_packet(&datagram) else {
1162 continue;
1163 };
1164 let Ok((pn, payload)) =
1165 decrypt_payload(datagram, &parsed, &initial_keys.client)
1166 else {
1167 continue;
1168 };
1169 (payload, pn, nxtquic_crypto::Level::Initial)
1170 }
1171 PacketType::Handshake => {
1172 let Ok(parsed) = parse_long_packet(&datagram) else {
1173 continue;
1174 };
1175 let Some((_, keys)) = handshake_keys.as_ref() else {
1176 continue;
1177 };
1178 let Ok((pn, payload)) = decrypt_payload(datagram, &parsed, &keys.remote) else {
1179 continue;
1180 };
1181 largest_handshake = largest_handshake.max(pn);
1182 (payload, pn, nxtquic_crypto::Level::Handshake)
1183 }
1184 PacketType::Short => {
1185 let Some(keys) = one_rtt_keys.as_ref() else {
1186 continue;
1187 };
1188 if datagram.len() < 1 + server_scid.len() + 4 + 16 {
1189 continue;
1190 }
1191 let parsed = LongPacket {
1192 _packet_type: PacketType::Short,
1193 _dcid: server_scid,
1194 _scid: ConnectionId::from_slice(&[]),
1195 pn_offset: 1 + server_scid.len(),
1196 packet_end: datagram.len(),
1197 };
1198 let Ok((pn, payload)) = decrypt_payload(datagram, &parsed, &keys.remote) else {
1199 continue;
1200 };
1201 (payload, pn, nxtquic_crypto::Level::OneRtt)
1202 }
1203 _ => continue,
1204 };
1205 let mut frames = &payload[..];
1206 while !frames.is_empty() {
1207 let before = frames.len();
1208 let Ok(frame) = Frame::decode(&mut frames) else {
1209 break;
1210 };
1211 match frame {
1212 Frame::Crypto(crypto)
1213 if matches!(
1214 level,
1215 nxtquic_crypto::Level::Initial | nxtquic_crypto::Level::Handshake
1216 ) =>
1217 {
1218 let Ok(output) = handshake.process(&crypto.data) else {
1219 continue;
1220 };
1221 while let Some((next_level, keys)) = handshake.next_keys() {
1222 if next_level == nxtquic_crypto::Level::OneRtt {
1223 one_rtt_keys = Some(keys);
1224 } else {
1225 handshake_keys = Some((next_level, keys));
1226 }
1227 }
1228 if !output.data.is_empty() {
1229 let (packet_type, packet_number, send_keys) = match output.level {
1230 nxtquic_crypto::Level::Initial => {
1231 (PacketType::Initial, 1, Some(&initial_keys.server))
1232 }
1233 nxtquic_crypto::Level::Handshake => (
1234 PacketType::Handshake,
1235 largest_handshake + 1,
1236 handshake_keys.as_ref().map(|(_, k)| &k.local),
1237 ),
1238 _ => (PacketType::Handshake, largest_handshake + 1, None),
1239 };
1240 let _ = send_protected_long(
1241 &socket,
1242 remote_addr,
1243 packet_type,
1244 initial.scid,
1245 server_scid,
1246 packet_number,
1247 &output.data,
1248 send_keys,
1249 )
1250 .await;
1251 }
1252 if handshake.is_complete() {
1253 if let Some(keys) = one_rtt_keys.as_ref() {
1254 let mut handshake_done = Vec::new();
1255 Frame::HandshakeDone.encode(&mut handshake_done);
1256 let _ = send_protected_short(
1257 &socket,
1258 remote_addr,
1259 initial.scid,
1260 0,
1261 &handshake_done,
1262 &keys.local,
1263 )
1264 .await;
1265 }
1266 }
1267 }
1268 Frame::Stream(stream) if level == nxtquic_crypto::Level::OneRtt => {
1269 let mut encoded = Vec::new();
1270 Frame::Stream(stream).encode(&mut encoded);
1271 let _ = connection.ingest_frames(&encoded).await;
1272 }
1273 _ => {}
1274 }
1275 if frames.len() >= before {
1276 break;
1277 }
1278 }
1279 let mut ack = Vec::new();
1280 Frame::Ack(AckFrame {
1281 largest_acknowledged: VarInt::from_u64(packet_number).unwrap(),
1282 ack_delay: VarInt::ZERO,
1283 first_ack_range: VarInt::ZERO,
1284 ack_ranges: Vec::new(),
1285 ecn_counts: None,
1286 })
1287 .encode(&mut ack);
1288 match level {
1289 nxtquic_crypto::Level::Initial => {
1290 let _ = send_protected_long(
1291 &socket,
1292 remote_addr,
1293 PacketType::Initial,
1294 initial.scid,
1295 server_scid,
1296 2,
1297 &ack,
1298 Some(&initial_keys.server),
1299 )
1300 .await;
1301 }
1302 nxtquic_crypto::Level::Handshake => {
1303 if let Some((_, keys)) = handshake_keys.as_ref() {
1304 let _ = send_protected_long(
1305 &socket,
1306 remote_addr,
1307 PacketType::Handshake,
1308 initial.scid,
1309 server_scid,
1310 largest_handshake + 1,
1311 &ack,
1312 Some(&keys.local),
1313 )
1314 .await;
1315 }
1316 }
1317 nxtquic_crypto::Level::OneRtt => {
1318 if let Some(keys) = one_rtt_keys.as_ref() {
1319 let _ = send_protected_short(
1320 &socket,
1321 remote_addr,
1322 initial.scid,
1323 application_packet_number,
1324 &ack,
1325 &keys.local,
1326 )
1327 .await;
1328 application_packet_number += 1;
1329 }
1330 }
1331 _ => {}
1332 }
1333 }
1334 connection.close().await;
1335}
1336
1337async fn send_protected_long(
1338 socket: &tokio::net::UdpSocket,
1339 remote_addr: SocketAddr,
1340 packet_type: PacketType,
1341 dcid: ConnectionId,
1342 scid: ConnectionId,
1343 packet_number: u64,
1344 payload: &[u8],
1345 keys: Option<&DirectionalKeys>,
1346) -> io::Result<()> {
1347 let Some(keys) = keys else {
1348 return Ok(());
1349 };
1350 let type_bits = match packet_type {
1351 PacketType::Handshake => 0x20,
1352 PacketType::Initial => 0x00,
1353 _ => return Ok(()),
1354 };
1355 let mut header = vec![0xC0 | type_bits, 0, 0, 0, 1, dcid.len() as u8];
1356 header.extend_from_slice(dcid.as_bytes());
1357 header.push(scid.len() as u8);
1358 header.extend_from_slice(scid.as_bytes());
1359 if packet_type == PacketType::Initial {
1360 VarInt::ZERO.encode(&mut header);
1361 }
1362 let mut body = payload.to_vec();
1363 while body.len() < 20 {
1364 body.push(0);
1365 }
1366 let packet_len = 1 + body.len() + keys.packet_key.tag_len();
1367 VarInt::from_u64(packet_len as u64)
1368 .map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "packet too large"))?
1369 .encode(&mut header);
1370 let pn_offset = header.len();
1371 header.push(packet_number as u8);
1372 keys.packet_key
1373 .encrypt(packet_number, &header, &mut body)
1374 .map_err(|error| io::Error::new(io::ErrorKind::Other, error))?;
1375 header.extend_from_slice(&body);
1376 let sample: [u8; 16] = header[pn_offset + 4..pn_offset + 20]
1377 .try_into()
1378 .map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "packet too short for sample"))?;
1379 let mask = keys
1380 .header_key
1381 .protection_mask(&sample)
1382 .map_err(|error| io::Error::new(io::ErrorKind::Other, error))?;
1383 apply_header_protection(&mut header, &mask, pn_offset, 1);
1384 socket.send_to(&header, remote_addr).await.map(|_| ())
1385}
1386
1387async fn send_protected_short(
1388 socket: &tokio::net::UdpSocket,
1389 remote_addr: SocketAddr,
1390 dcid: ConnectionId,
1391 packet_number: u64,
1392 payload: &[u8],
1393 keys: &DirectionalKeys,
1394) -> io::Result<()> {
1395 let mut header = vec![0x40];
1396 header.extend_from_slice(dcid.as_bytes());
1397 let pn_offset = header.len();
1398 header.push(packet_number as u8);
1399 let mut body = payload.to_vec();
1400 while body.len() < 20 {
1401 body.push(0);
1402 }
1403 keys.packet_key
1404 .encrypt(packet_number, &header, &mut body)
1405 .map_err(|error| io::Error::new(io::ErrorKind::Other, error))?;
1406 header.extend_from_slice(&body);
1407 let sample: [u8; 16] = header[pn_offset + 4..pn_offset + 20]
1408 .try_into()
1409 .map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "packet too short for sample"))?;
1410 let mask = keys
1411 .header_key
1412 .protection_mask(&sample)
1413 .map_err(|error| io::Error::new(io::ErrorKind::Other, error))?;
1414 apply_header_protection(&mut header, &mask, pn_offset, 1);
1415 socket.send_to(&header, remote_addr).await.map(|_| ())
1416}
1417
1418#[cfg(test)]
1419mod packet_tests {
1420 use super::*;
1421 use bytes::Bytes;
1422 use nxtquic_crypto::{CryptoProvider, RustlsCryptoProvider};
1423 use nxtquic_proto::{
1424 frame::{CryptoFrame, Frame},
1425 packet::protection::apply_header_protection,
1426 varint::VarInt,
1427 ConnectionId, Version,
1428 };
1429
1430 fn protected_client_initial() -> Vec<u8> {
1431 let dcid = ConnectionId::from_slice(&[0x10; 8]);
1432 let scid = [0x20; 8];
1433 let provider = RustlsCryptoProvider::new(None, None);
1434 let keys = provider.initial_keys(Version::V1, &dcid);
1435
1436 let mut plaintext = Vec::new();
1437 Frame::Ping.encode(&mut plaintext);
1438 Frame::Crypto(CryptoFrame {
1439 offset: VarInt::ZERO,
1440 data: Bytes::from_static(b"client hello bytes"),
1441 })
1442 .encode(&mut plaintext);
1443
1444 let ciphertext_len = plaintext.len() + keys.client.packet_key.tag_len();
1446 let mut header = vec![0xc0, 0, 0, 0, 1, dcid.len() as u8];
1447 header.extend_from_slice(dcid.as_bytes());
1448 header.push(scid.len() as u8);
1449 header.extend_from_slice(&scid);
1450 header.push(0); VarInt::from_u64((1 + ciphertext_len) as u64)
1452 .unwrap()
1453 .encode(&mut header);
1454 let packet_number_offset = header.len();
1455 header.push(0); keys.client
1458 .packet_key
1459 .encrypt(0, &header, &mut plaintext)
1460 .unwrap();
1461 header.extend_from_slice(&plaintext);
1462 let sample: [u8; 16] = header[packet_number_offset + 4..packet_number_offset + 20]
1463 .try_into()
1464 .unwrap();
1465 let mask = keys.client.header_key.protection_mask(&sample).unwrap();
1466 apply_header_protection(&mut header, &mask, packet_number_offset, 1);
1467 header
1468 }
1469
1470 #[test]
1471 fn decrypts_client_initial_and_decodes_crypto_frames() {
1472 let packet = protected_client_initial();
1473 let initial = decrypt_client_initial(&packet).unwrap();
1474 assert_eq!(initial._packet_number, 0);
1475 assert!(matches!(initial.frames.first(), Some(Frame::Ping)));
1476 assert!(
1477 matches!(initial.frames.get(1), Some(Frame::Crypto(frame)) if frame.data == Bytes::from_static(b"client hello bytes"))
1478 );
1479 }
1480
1481 #[test]
1482 fn rejects_tampered_client_initial() {
1483 let mut packet = protected_client_initial();
1484 let last = packet.len() - 1;
1485 packet[last] ^= 0x01;
1486 assert!(decrypt_client_initial(&packet).is_err());
1487 }
1488}