Skip to main content

nxtquic_api/
endpoint.rs

1//! Endpoint configuration and socket ownership.
2//!
3//! This module owns the UDP socket used by an endpoint. Packet driving is
4//! intentionally kept separate from socket binding so a server cannot report
5//! that it is listening before the operating system has accepted the bind.
6
7use 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/// Configuration for an endpoint.
37#[derive(Default, Clone, Debug)]
38pub struct EndpointConfig;
39
40/// Configuration for a server endpoint.
41#[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    /// Starts building a server configuration.
67    pub fn builder() -> ServerConfigBuilder {
68        ServerConfigBuilder::default()
69    }
70
71    /// Returns the Rustls configuration used by this endpoint.
72    pub fn rustls(&self) -> &Arc<rustls::ServerConfig> {
73        &self.rustls
74    }
75
76    /// Returns the transport configuration.
77    pub fn transport(&self) -> &TransportConfig {
78        &self.transport
79    }
80
81    /// Returns the maximum allowed concurrent connections.
82    pub fn max_concurrent_connections(&self) -> Option<usize> {
83        self.max_concurrent_connections
84    }
85
86    /// Returns whether active connection migration is enabled.
87    pub fn migration_enabled(&self) -> bool {
88        self.migration_enabled
89    }
90
91    /// Returns the advertised preferred address, if any.
92    pub fn preferred_address(&self) -> Option<SocketAddr> {
93        self.preferred_address
94    }
95}
96
97/// Builder for [`ServerConfig`].
98#[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    /// Uses an already configured Rustls server configuration.
109    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    /// Applies transport settings before building the server configuration.
125    pub fn with_transport_config(mut self, configure: impl FnOnce(&mut TransportConfig)) -> Self {
126        configure(&mut self.transport);
127        self
128    }
129
130    /// Sets the maximum number of concurrent active connections.
131    pub fn with_max_concurrent_connections(mut self, max: usize) -> Self {
132        self.max_concurrent_connections = Some(max);
133        self
134    }
135
136    /// Enables or disables connection migration support.
137    pub fn with_migration(mut self, enable: bool) -> Self {
138        self.migration_enabled = enable;
139        self
140    }
141
142    /// Sets the server preferred address (RFC 9000 §18.2).
143    pub fn with_preferred_address(mut self, addr: SocketAddr) -> Self {
144        self.preferred_address = Some(addr);
145        self
146    }
147
148    /// Builds the server configuration.
149    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/// Transport settings shared by all connections (RFC 9000 §18.2).
167#[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    /// Sets the maximum time a connection may remain idle before closing (RFC 9000 §10.1).
212    pub fn max_idle_timeout(&mut self, timeout: Option<Duration>) -> &mut Self {
213        self.max_idle_timeout = timeout;
214        self
215    }
216
217    /// Returns the configured idle timeout.
218    pub fn idle_timeout(&self) -> Option<Duration> {
219        self.max_idle_timeout
220    }
221
222    /// Sets the interval at which PING frames are sent to keep the connection alive (RFC 9000 §10.1.2).
223    pub fn keep_alive_interval(&mut self, interval: Option<Duration>) -> &mut Self {
224        self.keep_alive_interval = interval;
225        self
226    }
227
228    /// Returns the keep-alive interval.
229    pub fn keep_alive(&self) -> Option<Duration> {
230        self.keep_alive_interval
231    }
232
233    /// Sets the datagram receive buffer capacity in bytes (RFC 9221).
234    pub fn datagram_receive_buffer_size(&mut self, size: usize) -> &mut Self {
235        self.datagram_receive_buffer_size = size;
236        self
237    }
238
239    /// Returns the datagram receive buffer size.
240    pub fn datagram_buffer_size(&self) -> usize {
241        self.datagram_receive_buffer_size
242    }
243
244    /// Sets the initial connection-level flow control data limit in bytes (RFC 9000 §18.2).
245    pub fn initial_max_data(&mut self, value: u64) -> &mut Self {
246        self.initial_max_data = value;
247        self
248    }
249
250    /// Returns initial connection-level data limit.
251    pub fn get_initial_max_data(&self) -> u64 {
252        self.initial_max_data
253    }
254
255    /// Sets the initial flow control limit for locally initiated bidirectional streams.
256    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    /// Sets the initial flow control limit for peer-initiated bidirectional streams.
262    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    /// Sets the initial flow control limit for peer-initiated unidirectional streams.
268    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    /// Sets the initial maximum number of bidirectional streams the peer may open.
274    pub fn initial_max_streams_bidi(&mut self, value: u64) -> &mut Self {
275        self.initial_max_streams_bidi = value;
276        self
277    }
278
279    /// Sets the initial maximum number of unidirectional streams the peer may open.
280    pub fn initial_max_streams_uni(&mut self, value: u64) -> &mut Self {
281        self.initial_max_streams_uni = value;
282        self
283    }
284
285    /// Sets the maximum UDP payload size that can be received.
286    pub fn max_udp_payload_size(&mut self, value: u64) -> &mut Self {
287        self.max_udp_payload_size = value;
288        self
289    }
290
291    /// Sets the ACK delay exponent.
292    pub fn ack_delay_exponent(&mut self, value: u8) -> &mut Self {
293        self.ack_delay_exponent = value;
294        self
295    }
296
297    /// Sets the maximum ACK delay.
298    pub fn max_ack_delay(&mut self, value: Duration) -> &mut Self {
299        self.max_ack_delay = value;
300        self
301    }
302
303    /// Enables or disables active connection migration.
304    pub fn disable_active_migration(&mut self, disable: bool) -> &mut Self {
305        self.disable_active_migration = disable;
306        self
307    }
308
309    /// Sets the active connection ID limit advertised to the peer.
310    pub fn active_connection_id_limit(&mut self, limit: u64) -> &mut Self {
311        self.active_connection_id_limit = limit;
312        self
313    }
314
315    /// Sets the maximum concurrent unidirectional streams limit.
316    pub fn max_concurrent_uni_streams(&mut self, count: u64) -> &mut Self {
317        self.max_concurrent_uni_streams = count;
318        self
319    }
320
321    /// Sets the maximum concurrent bidirectional streams limit.
322    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/// Configuration for a client endpoint.
329#[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    /// Starts building a client configuration.
354    pub fn builder() -> ClientConfigBuilder {
355        ClientConfigBuilder::default()
356    }
357
358    /// Returns the Rustls client configuration.
359    pub fn rustls(&self) -> Option<&Arc<rustls::ClientConfig>> {
360        self.rustls.as_ref()
361    }
362
363    /// Returns the client transport configuration.
364    pub fn transport(&self) -> &TransportConfig {
365        &self.transport
366    }
367}
368
369/// Builder for [`ClientConfig`].
370#[derive(Default)]
371pub struct ClientConfigBuilder {
372    rustls: Option<Arc<rustls::ClientConfig>>,
373    transport: TransportConfig,
374}
375
376impl ClientConfigBuilder {
377    /// Uses an existing Rustls client configuration.
378    pub fn with_rustls(mut self, config: Arc<rustls::ClientConfig>) -> Self {
379        self.rustls = Some(config);
380        self
381    }
382
383    /// Configures transport parameters.
384    pub fn with_transport_config(mut self, configure: impl FnOnce(&mut TransportConfig)) -> Self {
385        configure(&mut self.transport);
386        self
387    }
388
389    /// Configures standard TLS 1.3 client defaults with ALPN `h3`.
390    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    /// Builds the client configuration.
401    pub fn build(self) -> ClientConfig {
402        ClientConfig {
403            rustls: self.rustls,
404            transport: self.transport,
405        }
406    }
407}
408
409/// Statistics about an endpoint.
410#[derive(Clone, Debug, Default)]
411pub struct EndpointStats {
412    /// Total incoming and outgoing connections initiated.
413    pub total_connections: u64,
414    /// Number of active connections.
415    pub active_connections: u64,
416    /// Number of failed handshakes.
417    pub handshake_failures: u64,
418    /// Total UDP packets received.
419    pub packets_received: u64,
420    /// Total UDP packets sent.
421    pub packets_sent: u64,
422}
423
424/// A QUIC endpoint.
425pub 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
435/// An ongoing connection attempt.
436pub struct Connecting;
437
438/// An incoming connection attempt.
439pub 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/// A decrypted client Initial packet.
449#[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); // Initial, fixed bit, one-byte packet number
550    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); // empty Retry token
556
557    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    /// Returns the peer address that sent the Initial packet.
608    pub fn remote_address(&self) -> SocketAddr {
609        self.remote_addr
610    }
611
612    /// Returns the local IP address on which the Initial packet arrived, if known.
613    pub fn local_ip(&self) -> Option<IpAddr> {
614        self.socket.local_addr().ok().map(|a| a.ip())
615    }
616
617    /// Cancels this incoming connection attempt without responding to the client.
618    pub fn cancel(&mut self) {
619        self.datagrams = None;
620    }
621
622    /// Rejects the incoming connection by closing the handshake with an error code (RFC 9000 §10.2).
623    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    /// Returns the first packet received for this connection attempt.
631    pub(crate) fn packet(&self) -> &[u8] {
632        &self.packet
633    }
634}
635
636impl Endpoint {
637    /// Creates a new endpoint bound to a local socket address.
638    pub async fn bind(addr: SocketAddr) -> io::Result<Self> {
639        Self::bind_inner(EndpointConfig::default(), None, addr).await
640    }
641
642    /// Binds a server endpoint using the supplied TLS and transport settings.
643    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    /// Returns the address the operating system assigned to this endpoint.
732    pub fn local_addr(&self) -> io::Result<SocketAddr> {
733        self.socket.local_addr()
734    }
735
736    /// Returns this endpoint's server configuration, if configured.
737    pub async fn server_config(&self) -> Option<ServerConfig> {
738        self.server_config.lock().await.clone()
739    }
740
741    /// Hot-reloads the server configuration and TLS certificates.
742    pub async fn set_server_config(&self, config: ServerConfig) {
743        *self.server_config.lock().await = Some(config);
744    }
745
746    /// Returns snapshot statistics for this endpoint.
747    pub async fn stats(&self) -> EndpointStats {
748        self.stats.lock().await.clone()
749    }
750
751    /// Gracefully closes the endpoint and all active connections.
752    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    /// Awaits until all active connections on this endpoint have closed.
758    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    /// Connects to a remote endpoint.
769    pub async fn connect(
770        &self,
771        _addr: SocketAddr,
772        _server_name: &str,
773    ) -> std::io::Result<Connecting> {
774        Ok(Connecting)
775    }
776
777    /// Connects to a remote endpoint using custom client configuration.
778    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    /// Connects to a remote endpoint with 0-RTT early data support (RFC 9001 §4.2).
788    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    /// Accepts an incoming connection.
797    pub async fn accept(&self) -> Option<Incoming> {
798        self.incoming.lock().await.recv().await
799    }
800}
801
802impl Connecting {
803    /// Waits for the connection attempt to complete.
804    pub async fn await_connection(self) -> std::io::Result<Connection> {
805        Ok(Connection::new())
806    }
807
808    /// Splits the connecting future into early 0-RTT data and fully confirmed connection.
809    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    /// Accepts the incoming connection.
817    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        // The Initial Length field includes the packet number and AEAD tag.
1445        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); // empty token
1451        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); // packet number 0, one byte
1456
1457        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}