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::{collections::HashMap, io, net::SocketAddr, sync::Arc, time::Duration};
8
9use crate::connection::Connection;
10use nxtquic_crypto::{
11    CryptoProvider, DirectionalKeys, HandshakeState, HeaderProtectionKey, Keys, PacketKey,
12    RustlsCryptoProvider,
13};
14use nxtquic_proto::{
15    ConnectionId, StreamId, Version,
16    frame::{AckFrame, CryptoFrame, Frame, StreamFrame},
17    packet::{
18        PacketType, decode_header,
19        initial::parse_initial,
20        protection::{apply_header_protection, remove_header_protection},
21    },
22    varint::VarInt,
23};
24
25/// Configuration for an endpoint.
26#[derive(Default, Clone, Debug)]
27pub struct EndpointConfig;
28
29/// Configuration for a server endpoint.
30#[derive(Clone)]
31pub struct ServerConfig {
32    rustls: Arc<rustls::ServerConfig>,
33    transport: TransportConfig,
34}
35
36impl std::fmt::Debug for ServerConfig {
37    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
38        formatter
39            .debug_struct("ServerConfig")
40            .field("transport", &self.transport)
41            .finish_non_exhaustive()
42    }
43}
44
45impl ServerConfig {
46    /// Starts building a server configuration.
47    pub fn builder() -> ServerConfigBuilder {
48        ServerConfigBuilder::default()
49    }
50
51    /// Returns the Rustls configuration used by this endpoint.
52    pub fn rustls(&self) -> &Arc<rustls::ServerConfig> {
53        &self.rustls
54    }
55
56    /// Returns the transport configuration.
57    pub fn transport(&self) -> &TransportConfig {
58        &self.transport
59    }
60}
61
62/// Builder for [`ServerConfig`].
63#[derive(Default)]
64pub struct ServerConfigBuilder {
65    rustls: Option<Arc<rustls::ServerConfig>>,
66    transport: TransportConfig,
67}
68
69impl ServerConfigBuilder {
70    /// Uses an already configured Rustls server configuration.
71    pub fn with_rustls(mut self, config: Arc<rustls::ServerConfig>) -> io::Result<Self> {
72        if config
73            .alpn_protocols
74            .iter()
75            .all(|protocol| protocol.as_slice() != b"h3")
76        {
77            return Err(io::Error::new(
78                io::ErrorKind::InvalidInput,
79                "a QUIC HTTP/3 server configuration must advertise ALPN h3",
80            ));
81        }
82        self.rustls = Some(config);
83        Ok(self)
84    }
85
86    /// Applies transport settings before building the server configuration.
87    pub fn with_transport_config(mut self, configure: impl FnOnce(&mut TransportConfig)) -> Self {
88        configure(&mut self.transport);
89        self
90    }
91
92    /// Builds the server configuration.
93    pub fn build(self) -> io::Result<ServerConfig> {
94        let rustls = self.rustls.ok_or_else(|| {
95            io::Error::new(
96                io::ErrorKind::InvalidInput,
97                "missing Rustls server configuration",
98            )
99        })?;
100        Ok(ServerConfig {
101            rustls,
102            transport: self.transport,
103        })
104    }
105}
106
107/// Transport settings shared by all server connections.
108#[derive(Default, Clone, Debug)]
109pub struct TransportConfig {
110    max_idle_timeout: Option<Duration>,
111}
112
113impl TransportConfig {
114    /// Sets the maximum time a connection may remain idle.
115    pub fn max_idle_timeout(&mut self, timeout: Option<Duration>) -> &mut Self {
116        self.max_idle_timeout = timeout;
117        self
118    }
119
120    /// Returns the configured idle timeout.
121    pub fn idle_timeout(&self) -> Option<Duration> {
122        self.max_idle_timeout
123    }
124}
125
126/// Configuration for a client endpoint.
127#[derive(Default, Clone, Debug)]
128pub struct ClientConfig;
129
130/// A QUIC endpoint.
131pub struct Endpoint {
132    config: EndpointConfig,
133    server_config: Option<ServerConfig>,
134    socket: Arc<tokio::net::UdpSocket>,
135    incoming: Arc<tokio::sync::Mutex<tokio::sync::mpsc::Receiver<Incoming>>>,
136}
137
138/// An ongoing connection attempt.
139pub struct Connecting;
140
141/// An incoming connection attempt.
142pub struct Incoming {
143    remote_addr: SocketAddr,
144    packet: Vec<u8>,
145    initial: DecryptedInitial,
146    socket: Arc<tokio::net::UdpSocket>,
147    server_config: Option<ServerConfig>,
148    datagrams: Option<tokio::sync::mpsc::Receiver<Vec<u8>>>,
149}
150
151/// A decrypted client Initial packet.  This is kept internal while the
152/// endpoint driver is completed, but it deliberately contains decoded frames
153/// rather than a synthetic "accepted" connection so callers cannot mistake an
154/// unprocessed UDP datagram for a completed handshake.
155#[derive(Debug)]
156struct DecryptedInitial {
157    version: Version,
158    dcid: ConnectionId,
159    scid: ConnectionId,
160    packet_number: u64,
161    frames: Vec<Frame>,
162}
163
164fn decrypt_client_initial(packet: &[u8]) -> io::Result<DecryptedInitial> {
165    let parsed =
166        parse_initial(packet).map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
167    let version = match parsed.version {
168        1 => Version::V1,
169        0x6b33_43cf => Version::V2,
170        version => {
171            return Err(io::Error::new(
172                io::ErrorKind::InvalidData,
173                format!("unsupported QUIC version {version:#010x}"),
174            ));
175        }
176    };
177
178    let provider = RustlsCryptoProvider::new(None, None);
179    let keys = provider.initial_keys(version, &parsed.dcid);
180    let mut protected = packet[..parsed.packet_end].to_vec();
181
182    let sample_offset = parsed.packet_number_offset + 4;
183    let sample: [u8; 16] = protected[sample_offset..sample_offset + 16]
184        .try_into()
185        .expect("Initial parser validates the header-protection sample");
186    let mask = keys
187        .client
188        .header_key
189        .protection_mask(&sample)
190        .map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
191    let packet_number_len =
192        remove_header_protection(&mut protected, &mask, parsed.packet_number_offset);
193    if packet_number_len == 0 || parsed.packet_number_offset + packet_number_len > parsed.packet_end
194    {
195        return Err(io::Error::new(
196            io::ErrorKind::InvalidData,
197            "invalid QUIC Initial packet number",
198        ));
199    }
200
201    let mut packet_number = 0_u64;
202    for byte in
203        &protected[parsed.packet_number_offset..parsed.packet_number_offset + packet_number_len]
204    {
205        packet_number = (packet_number << 8) | u64::from(*byte);
206    }
207    let header_end = parsed.packet_number_offset + packet_number_len;
208    let mut payload = protected[header_end..parsed.packet_end].to_vec();
209    keys.client
210        .packet_key
211        .decrypt(packet_number, &protected[..header_end], &mut payload)
212        .map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
213
214    let mut frames = Vec::new();
215    let mut payload = &payload[..];
216    while !payload.is_empty() {
217        let before = payload.len();
218        let frame = Frame::decode(&mut payload)
219            .map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
220        frames.push(frame);
221        if payload.len() == before {
222            return Err(io::Error::new(
223                io::ErrorKind::InvalidData,
224                "QUIC frame decoder made no progress",
225            ));
226        }
227    }
228
229    Ok(DecryptedInitial {
230        version,
231        dcid: parsed.dcid,
232        scid: parsed.scid,
233        packet_number,
234        frames,
235    })
236}
237
238fn encode_server_initial(
239    version: Version,
240    client_scid: ConnectionId,
241    server_scid: ConnectionId,
242    packet_number: u64,
243    crypto_data: Vec<u8>,
244    keys: nxtquic_crypto::InitialKeys,
245) -> io::Result<Vec<u8>> {
246    let mut plaintext = Vec::new();
247    Frame::Crypto(CryptoFrame {
248        offset: VarInt::ZERO,
249        data: crypto_data.into(),
250    })
251    .encode(&mut plaintext);
252
253    let mut header_prefix = Vec::new();
254    header_prefix.push(0xc0); // Initial, fixed bit, one-byte packet number
255    header_prefix.extend_from_slice(&version.as_u32().to_be_bytes());
256    header_prefix.push(client_scid.len() as u8);
257    header_prefix.extend_from_slice(client_scid.as_bytes());
258    header_prefix.push(server_scid.len() as u8);
259    header_prefix.extend_from_slice(server_scid.as_bytes());
260    header_prefix.push(0); // empty Retry token
261
262    // Every datagram containing an Initial packet must be at least 1200 bytes
263    // (RFC 9000 ยง14.1). Padding uses PADDING frames, which are zero bytes.
264    loop {
265        let encrypted_len = plaintext.len() + keys.server.packet_key.tag_len();
266        let mut candidate = header_prefix.clone();
267        VarInt::from_u64((1 + encrypted_len) as u64)
268            .expect("QUIC packet lengths fit in a VarInt")
269            .encode(&mut candidate);
270        if candidate.len() + 1 + encrypted_len >= 1200 {
271            header_prefix = candidate;
272            break;
273        }
274        plaintext.push(0);
275    }
276
277    let packet_number_len = 1;
278    let packet_number_offset = header_prefix.len();
279    header_prefix.push(packet_number as u8);
280    keys.server
281        .packet_key
282        .encrypt(packet_number, &header_prefix, &mut plaintext)
283        .map_err(|error| io::Error::new(io::ErrorKind::Other, error))?;
284    header_prefix.extend_from_slice(&plaintext);
285
286    let sample: [u8; 16] = header_prefix[packet_number_offset + 4..packet_number_offset + 20]
287        .try_into()
288        .expect("1200-byte Initial always contains a header-protection sample");
289    let mask = keys
290        .server
291        .header_key
292        .protection_mask(&sample)
293        .map_err(|error| io::Error::new(io::ErrorKind::Other, error))?;
294    apply_header_protection(
295        &mut header_prefix,
296        &mask,
297        packet_number_offset,
298        packet_number_len,
299    );
300    Ok(header_prefix)
301}
302
303impl std::fmt::Debug for Incoming {
304    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
305        formatter
306            .debug_struct("Incoming")
307            .field("remote_addr", &self.remote_addr)
308            .field("packet_len", &self.packet.len())
309            .finish()
310    }
311}
312
313impl Incoming {
314    /// Returns the peer address that sent the Initial packet.
315    pub fn remote_address(&self) -> SocketAddr {
316        self.remote_addr
317    }
318
319    /// Returns the first packet received for this connection attempt.
320    pub(crate) fn packet(&self) -> &[u8] {
321        &self.packet
322    }
323}
324
325impl Endpoint {
326    /// Creates a new endpoint.
327    /// Binds the endpoint to a local socket address.
328    pub async fn bind(addr: SocketAddr) -> io::Result<Self> {
329        Self::bind_inner(EndpointConfig::default(), None, addr).await
330    }
331
332    /// Binds a server endpoint using the supplied TLS and transport settings.
333    pub async fn server(server_config: ServerConfig, addr: SocketAddr) -> io::Result<Self> {
334        Self::bind_inner(EndpointConfig::default(), Some(server_config), addr).await
335    }
336
337    async fn bind_inner(
338        config: EndpointConfig,
339        server_config: Option<ServerConfig>,
340        addr: SocketAddr,
341    ) -> io::Result<Self> {
342        let socket = tokio::net::UdpSocket::bind(addr).await?;
343        let socket = Arc::new(socket);
344        let (incoming_tx, incoming_rx) = tokio::sync::mpsc::channel(128);
345        let receive_socket = Arc::clone(&socket);
346        let receive_server_config = server_config.clone();
347        tokio::spawn(async move {
348            let mut packet = vec![0_u8; 65_535];
349            let mut peers: HashMap<SocketAddr, tokio::sync::mpsc::Sender<Vec<u8>>> = HashMap::new();
350            loop {
351                let (length, remote_addr) = match receive_socket.recv_from(&mut packet).await {
352                    Ok(received) => received,
353                    Err(_) => break,
354                };
355
356                let datagram = &packet[..length];
357                if let Some(peer_tx) = peers.get(&remote_addr) {
358                    if peer_tx.send(datagram.to_vec()).await.is_err() {
359                        peers.remove(&remote_addr);
360                    }
361                    continue;
362                }
363                let Ok(header) = decode_header(datagram) else {
364                    continue;
365                };
366                if header.packet_type != PacketType::Initial {
367                    continue;
368                }
369                let Ok(initial) = decrypt_client_initial(datagram) else {
370                    continue;
371                };
372
373                let (peer_tx, peer_rx) = tokio::sync::mpsc::channel(256);
374                peers.insert(remote_addr, peer_tx);
375
376                if incoming_tx
377                    .send(Incoming {
378                        remote_addr,
379                        packet: datagram.to_vec(),
380                        initial,
381                        socket: Arc::clone(&receive_socket),
382                        server_config: receive_server_config.clone(),
383                        datagrams: Some(peer_rx),
384                    })
385                    .await
386                    .is_err()
387                {
388                    break;
389                }
390            }
391        });
392
393        Ok(Self {
394            config,
395            server_config,
396            socket,
397            incoming: Arc::new(tokio::sync::Mutex::new(incoming_rx)),
398        })
399    }
400
401    /// Returns the address the operating system assigned to this endpoint.
402    pub fn local_addr(&self) -> io::Result<SocketAddr> {
403        self.socket.local_addr()
404    }
405
406    /// Returns this endpoint's server configuration, if it is a server.
407    pub fn server_config(&self) -> Option<&ServerConfig> {
408        self.server_config.as_ref()
409    }
410
411    /// Connects to a remote endpoint.
412    pub async fn connect(
413        &self,
414        _addr: SocketAddr,
415        _server_name: &str,
416    ) -> std::io::Result<Connecting> {
417        Ok(Connecting)
418    }
419
420    /// Accepts an incoming connection.
421    pub async fn accept(&self) -> Option<Incoming> {
422        self.incoming.lock().await.recv().await
423    }
424}
425
426impl Connecting {
427    /// Waits for the connection attempt to complete.
428    pub async fn await_connection(self) -> std::io::Result<Connection> {
429        Ok(Connection::new())
430    }
431}
432
433impl Incoming {
434    /// Accepts the incoming connection.
435    pub async fn accept(self) -> std::io::Result<Connection> {
436        let config = self.server_config.ok_or_else(|| {
437            io::Error::new(
438                io::ErrorKind::InvalidInput,
439                "cannot accept a connection on an endpoint without ServerConfig",
440            )
441        })?;
442        let mut crypto = Vec::new();
443        for frame in &self.initial.frames {
444            if let Frame::Crypto(frame) = frame {
445                if frame.offset.into_inner() != crypto.len() as u64 {
446                    return Err(io::Error::new(
447                        io::ErrorKind::InvalidData,
448                        "Initial CRYPTO frames must be contiguous",
449                    ));
450                }
451                crypto.extend_from_slice(&frame.data);
452            }
453        }
454        if crypto.is_empty() {
455            return Err(io::Error::new(
456                io::ErrorKind::InvalidData,
457                "Initial packet contains no CRYPTO frame",
458            ));
459        }
460
461        let provider = RustlsCryptoProvider::new(None, Some(Arc::clone(config.rustls())));
462        
463        let mut transport_params = nxtquic_proto::TransportParameters::default();
464        transport_params.initial_max_data = VarInt::from_u64(10 * 1024 * 1024).unwrap();
465        transport_params.initial_max_stream_data_bidi_local = VarInt::from_u64(1024 * 1024).unwrap();
466        transport_params.initial_max_stream_data_bidi_remote = VarInt::from_u64(1024 * 1024).unwrap();
467        transport_params.initial_max_stream_data_uni = VarInt::from_u64(1024 * 1024).unwrap();
468        transport_params.initial_max_streams_bidi = VarInt::from_u64(100).unwrap();
469        transport_params.initial_max_streams_uni = VarInt::from_u64(100).unwrap();
470        if let Some(timeout) = config.transport().idle_timeout() {
471            transport_params.max_idle_timeout = VarInt::from_u64(timeout.as_millis() as u64).unwrap_or(VarInt::from_u32(0));
472        }
473        transport_params.max_udp_payload_size = VarInt::from_u64(1350).unwrap();
474        transport_params.active_connection_id_limit = VarInt::from_u64(2).unwrap();
475        
476        let server_scid = ConnectionId::from_slice(&rand::random::<[u8; 16]>());
477        transport_params.initial_source_connection_id = Some(server_scid.clone());
478        transport_params.original_destination_connection_id = Some(self.initial.dcid.clone());
479
480        let mut tp_encoded = Vec::new();
481        transport_params.encode(&mut tp_encoded);
482
483        let mut handshake = provider
484            .new_server_handshake(tp_encoded)
485            .map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
486        let output = handshake
487            .process(&crypto)
488            .map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
489        if output.data.is_empty() {
490            return Err(io::Error::new(
491                io::ErrorKind::InvalidData,
492                "TLS produced no Initial handshake data",
493            ));
494        }
495
496        let keys = provider.initial_keys(self.initial.version, &self.initial.dcid);
497        let response = encode_server_initial(
498            self.initial.version,
499            self.initial.scid,
500            server_scid,
501            0,
502            output.data,
503            keys,
504        )?;
505        self.socket.send_to(&response, self.remote_addr).await?;
506        let (outgoing_tx, outgoing_rx) = tokio::sync::mpsc::unbounded_channel();
507        let connection = Connection::with_outgoing(outgoing_tx);
508        if let Some(datagrams) = self.datagrams {
509            let driver_connection = connection.clone();
510            let driver_initial = self.initial;
511            let driver_socket = Arc::clone(&self.socket);
512            let driver_remote = self.remote_addr;
513            let driver_initial_keys = provider.initial_keys(driver_initial.version, &driver_initial.dcid);
514            let mut driver_handshake_keys = None;
515            let mut driver_one_rtt_keys = None;
516            while let Some((level, keys)) = handshake.next_keys() {
517                if level == nxtquic_crypto::Level::OneRtt {
518                    driver_one_rtt_keys = Some(keys);
519                } else {
520                    driver_handshake_keys = Some((level, keys));
521                }
522            }
523            tokio::spawn(async move {
524                drive_server_connection(
525                    driver_connection,
526                    driver_socket,
527                    driver_remote,
528                    datagrams,
529                    driver_initial,
530                    server_scid,
531                    driver_initial_keys,
532                    driver_handshake_keys,
533                    driver_one_rtt_keys,
534                    handshake,
535                    outgoing_rx,
536                )
537                .await;
538            });
539        }
540        Ok(connection)
541    }
542}
543
544struct LongPacket {
545    packet_type: PacketType,
546    dcid: ConnectionId,
547    scid: ConnectionId,
548    pn_offset: usize,
549    packet_end: usize,
550}
551
552fn parse_long_packet(datagram: &[u8]) -> io::Result<LongPacket> {
553    if datagram.len() < 7 || datagram[0] & 0xc0 != 0xc0 {
554        return Err(io::Error::new(io::ErrorKind::InvalidData, "invalid QUIC long header"));
555    }
556    let version = u32::from_be_bytes(datagram[1..5].try_into().unwrap());
557    if version == 0 {
558        return Err(io::Error::new(io::ErrorKind::InvalidData, "version negotiation is not a protected packet"));
559    }
560    let packet_type = match (datagram[0] >> 4) & 0x03 {
561        0 => PacketType::Initial,
562        1 => PacketType::ZeroRtt,
563        2 => PacketType::Handshake,
564        _ => PacketType::Retry,
565    };
566    let mut offset = 5;
567    let dcid_len = datagram[offset] as usize;
568    offset += 1;
569    let dcid_end = offset.checked_add(dcid_len).ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "invalid DCID"))?;
570    let dcid = ConnectionId::from_slice(datagram.get(offset..dcid_end).ok_or_else(|| io::Error::new(io::ErrorKind::UnexpectedEof, "truncated DCID"))?);
571    offset = dcid_end;
572    let scid_len = *datagram.get(offset).ok_or_else(|| io::Error::new(io::ErrorKind::UnexpectedEof, "missing SCID length"))? as usize;
573    offset += 1;
574    let scid_end = offset.checked_add(scid_len).ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "invalid SCID"))?;
575    let scid = ConnectionId::from_slice(datagram.get(offset..scid_end).ok_or_else(|| io::Error::new(io::ErrorKind::UnexpectedEof, "truncated SCID"))?);
576    offset = scid_end;
577    if packet_type == PacketType::Initial {
578        let mut rest = &datagram[offset..];
579        let token_len = VarInt::decode(&mut rest).map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "invalid Initial token length"))?.into_inner() as usize;
580        let consumed = datagram[offset..].len() - rest.len();
581        offset += consumed.checked_add(token_len).ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "invalid token"))?;
582    }
583    let mut rest = &datagram[offset..];
584    let payload_len = VarInt::decode(&mut rest).map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "invalid QUIC payload length"))?.into_inner() as usize;
585    offset += datagram[offset..].len() - rest.len();
586    let packet_end = offset.checked_add(payload_len).ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "invalid packet length"))?;
587    if packet_end > datagram.len() || packet_end < offset + 4 + 16 {
588        return Err(io::Error::new(io::ErrorKind::UnexpectedEof, "truncated protected QUIC packet"));
589    }
590    Ok(LongPacket { packet_type, dcid, scid, pn_offset: offset, packet_end })
591}
592
593fn decrypt_payload(mut packet: Vec<u8>, parsed: &LongPacket, keys: &DirectionalKeys) -> io::Result<(u64, Vec<u8>)> {
594    let sample: [u8; 16] = packet[parsed.pn_offset + 4..parsed.pn_offset + 20].try_into().map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "missing header-protection sample"))?;
595    let mask = keys.header_key.protection_mask(&sample).map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
596    let pn_len = remove_header_protection(&mut packet, &mask, parsed.pn_offset);
597    if pn_len == 0 || parsed.pn_offset + pn_len > parsed.packet_end {
598        return Err(io::Error::new(io::ErrorKind::InvalidData, "invalid packet number"));
599    }
600    let mut packet_number = 0;
601    for byte in &packet[parsed.pn_offset..parsed.pn_offset + pn_len] {
602        packet_number = (packet_number << 8) | u64::from(*byte);
603    }
604    let header_end = parsed.pn_offset + pn_len;
605    let mut payload = packet[header_end..parsed.packet_end].to_vec();
606    keys.packet_key.decrypt(packet_number, &packet[..header_end], &mut payload).map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
607    Ok((packet_number, payload))
608}
609
610async fn drive_server_connection(
611    connection: Connection,
612    socket: Arc<tokio::net::UdpSocket>,
613    remote_addr: SocketAddr,
614    mut datagrams: tokio::sync::mpsc::Receiver<Vec<u8>>,
615    initial: DecryptedInitial,
616    server_scid: ConnectionId,
617    initial_keys: nxtquic_crypto::InitialKeys,
618    mut handshake_keys: Option<(nxtquic_crypto::Level, Keys)>,
619    mut one_rtt_keys: Option<Keys>,
620    mut handshake: impl HandshakeState,
621    mut outgoing: tokio::sync::mpsc::UnboundedReceiver<crate::stream::WriteCommand>,
622) {
623    let mut largest_handshake = 0;
624    let mut application_packet_number = 0;
625    loop {
626        let Some(event) = (tokio::select! {
627            datagram = datagrams.recv() => datagram.map(Ok),
628            command = outgoing.recv() => command.map(Err),
629            else => None,
630        }) else { break };
631        let datagram = match event {
632            Err(command) => {
633                if let Some(keys) = one_rtt_keys.as_ref() {
634                    let crate::stream::WriteCommand::Data { stream_id, offset, data, fin } = command;
635                    let mut encoded = Vec::new();
636                    Frame::Stream(StreamFrame {
637                        stream_id: StreamId::from_u64(stream_id),
638                        offset: VarInt::from_u64(offset).unwrap(),
639                        length: Some(VarInt::from_u64(data.len() as u64).unwrap()),
640                        fin,
641                        data: bytes::Bytes::from(data),
642                    }).encode(&mut encoded);
643                    let _ = send_protected_short(&socket, remote_addr, initial.scid, application_packet_number, &encoded, &keys.local).await;
644                    application_packet_number += 1;
645                }
646                continue;
647            }
648            Ok(datagram) => datagram,
649        };
650        let Ok(header) = decode_header(&datagram) else { continue };
651        let (payload, packet_number, level) = match header.packet_type {
652            PacketType::Initial => {
653                let Ok(parsed) = parse_long_packet(&datagram) else { continue };
654                let Ok((pn, payload)) = decrypt_payload(datagram, &parsed, &initial_keys.client) else { continue };
655                (payload, pn, nxtquic_crypto::Level::Initial)
656            }
657            PacketType::Handshake => {
658                let Ok(parsed) = parse_long_packet(&datagram) else { continue };
659                let Some((_, keys)) = handshake_keys.as_ref() else { continue };
660                let Ok((pn, payload)) = decrypt_payload(datagram, &parsed, &keys.remote) else { continue };
661                largest_handshake = largest_handshake.max(pn);
662                (payload, pn, nxtquic_crypto::Level::Handshake)
663            }
664            PacketType::Short => {
665                let Some(keys) = one_rtt_keys.as_ref() else { continue };
666                if datagram.len() < 1 + server_scid.len() + 4 + 16 { continue; }
667                let parsed = LongPacket { packet_type: PacketType::Short, dcid: server_scid, scid: ConnectionId::from_slice(&[]), pn_offset: 1 + server_scid.len(), packet_end: datagram.len() };
668                let Ok((pn, payload)) = decrypt_payload(datagram, &parsed, &keys.remote) else { continue };
669                (payload, pn, nxtquic_crypto::Level::OneRtt)
670            }
671            _ => continue,
672        };
673        let mut frames = &payload[..];
674        while !frames.is_empty() {
675            let before = frames.len();
676            let Ok(frame) = Frame::decode(&mut frames) else { break };
677            match frame {
678                Frame::Crypto(crypto) if matches!(level, nxtquic_crypto::Level::Initial | nxtquic_crypto::Level::Handshake) => {
679                    let Ok(output) = handshake.process(&crypto.data) else { continue };
680                    while let Some((next_level, keys)) = handshake.next_keys() {
681                        if next_level == nxtquic_crypto::Level::OneRtt { one_rtt_keys = Some(keys); } else { handshake_keys = Some((next_level, keys)); }
682                    }
683                    if !output.data.is_empty() {
684                        let (packet_type, packet_number, send_keys) = match output.level {
685                            nxtquic_crypto::Level::Initial => (PacketType::Initial, 1, Some(&initial_keys.server)),
686                            nxtquic_crypto::Level::Handshake => (PacketType::Handshake, largest_handshake + 1, handshake_keys.as_ref().map(|(_, k)| &k.local)),
687                            _ => (PacketType::Handshake, largest_handshake + 1, None),
688                        };
689                        let _ = send_protected_long(&socket, remote_addr, packet_type, initial.scid, server_scid, packet_number, &output.data, send_keys).await;
690                    }
691                    if handshake.is_complete() {
692                        if let Some(keys) = one_rtt_keys.as_ref() {
693                            let mut handshake_done = Vec::new();
694                            Frame::HandshakeDone.encode(&mut handshake_done);
695                            let _ = send_protected_short(&socket, remote_addr, initial.scid, 0, &handshake_done, &keys.local).await;
696                        }
697                    }
698                }
699                Frame::Stream(stream) if level == nxtquic_crypto::Level::OneRtt => {
700                    let mut encoded = Vec::new();
701                    Frame::Stream(stream).encode(&mut encoded);
702                    let _ = connection.ingest_frames(&encoded).await;
703                }
704                _ => {}
705            }
706            if frames.len() >= before { break; }
707        }
708        let mut ack = Vec::new();
709        Frame::Ack(AckFrame {
710            largest_acknowledged: VarInt::from_u64(packet_number).unwrap(),
711            ack_delay: VarInt::ZERO,
712            first_ack_range: VarInt::ZERO,
713            ack_ranges: Vec::new(),
714            ecn_counts: None,
715        }).encode(&mut ack);
716        match level {
717            nxtquic_crypto::Level::Initial => {
718                let _ = send_protected_long(&socket, remote_addr, PacketType::Initial, initial.scid, server_scid, 2, &ack, Some(&initial_keys.server)).await;
719            }
720            nxtquic_crypto::Level::Handshake => {
721                if let Some((_, keys)) = handshake_keys.as_ref() {
722                    let _ = send_protected_long(&socket, remote_addr, PacketType::Handshake, initial.scid, server_scid, largest_handshake + 1, &ack, Some(&keys.local)).await;
723                }
724            }
725            nxtquic_crypto::Level::OneRtt => {
726                if let Some(keys) = one_rtt_keys.as_ref() {
727                    let _ = send_protected_short(&socket, remote_addr, initial.scid, application_packet_number, &ack, &keys.local).await;
728                    application_packet_number += 1;
729                }
730            }
731            _ => {}
732        }
733        let _ = packet_number;
734    }
735    connection.close().await;
736}
737
738async fn send_protected_long(
739    socket: &tokio::net::UdpSocket,
740    remote_addr: SocketAddr,
741    packet_type: PacketType,
742    dcid: ConnectionId,
743    scid: ConnectionId,
744    packet_number: u64,
745    payload: &[u8],
746    keys: Option<&DirectionalKeys>,
747) -> io::Result<()> {
748    let Some(keys) = keys else { return Ok(()); };
749    let type_bits = match packet_type { PacketType::Handshake => 0x20, PacketType::Initial => 0x00, _ => return Ok(()) };
750    let mut header = vec![0xC0 | type_bits, 0, 0, 0, 1, dcid.len() as u8];
751    header.extend_from_slice(dcid.as_bytes());
752    header.push(scid.len() as u8);
753    header.extend_from_slice(scid.as_bytes());
754    if packet_type == PacketType::Initial { VarInt::ZERO.encode(&mut header); }
755    let mut body = payload.to_vec();
756    while body.len() < 20 {
757        body.push(0);
758    }
759    let packet_len = 1 + body.len() + keys.packet_key.tag_len();
760    VarInt::from_u64(packet_len as u64).map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "packet too large"))?.encode(&mut header);
761    let pn_offset = header.len();
762    header.push(packet_number as u8);
763    keys.packet_key.encrypt(packet_number, &header, &mut body).map_err(|error| io::Error::new(io::ErrorKind::Other, error))?;
764    header.extend_from_slice(&body);
765    let sample: [u8; 16] = header[pn_offset + 4..pn_offset + 20].try_into().map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "packet too short for sample"))?;
766    let mask = keys.header_key.protection_mask(&sample).map_err(|error| io::Error::new(io::ErrorKind::Other, error))?;
767    apply_header_protection(&mut header, &mask, pn_offset, 1);
768    socket.send_to(&header, remote_addr).await.map(|_| ())
769}
770
771async fn send_protected_short(
772    socket: &tokio::net::UdpSocket,
773    remote_addr: SocketAddr,
774    dcid: ConnectionId,
775    packet_number: u64,
776    payload: &[u8],
777    keys: &DirectionalKeys,
778) -> io::Result<()> {
779    let mut header = vec![0x40];
780    header.extend_from_slice(dcid.as_bytes());
781    let pn_offset = header.len();
782    header.push(packet_number as u8);
783    let mut body = payload.to_vec();
784    while body.len() < 20 {
785        body.push(0);
786    }
787    keys.packet_key.encrypt(packet_number, &header, &mut body).map_err(|error| io::Error::new(io::ErrorKind::Other, error))?;
788    header.extend_from_slice(&body);
789    let sample: [u8; 16] = header[pn_offset + 4..pn_offset + 20].try_into().map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "packet too short for sample"))?;
790    let mask = keys.header_key.protection_mask(&sample).map_err(|error| io::Error::new(io::ErrorKind::Other, error))?;
791    apply_header_protection(&mut header, &mask, pn_offset, 1);
792    socket.send_to(&header, remote_addr).await.map(|_| ())
793}
794
795#[cfg(test)]
796mod packet_tests {
797    use super::*;
798    use bytes::Bytes;
799    use nxtquic_crypto::{CryptoProvider, RustlsCryptoProvider};
800    use nxtquic_proto::{
801        ConnectionId, Version,
802        frame::{CryptoFrame, Frame},
803        packet::protection::apply_header_protection,
804        varint::VarInt,
805    };
806
807    fn protected_client_initial() -> Vec<u8> {
808        let dcid = ConnectionId::from_slice(&[0x10; 8]);
809        let scid = [0x20; 8];
810        let provider = RustlsCryptoProvider::new(None, None);
811        let keys = provider.initial_keys(Version::V1, &dcid);
812
813        let mut plaintext = Vec::new();
814        Frame::Ping.encode(&mut plaintext);
815        Frame::Crypto(CryptoFrame {
816            offset: VarInt::ZERO,
817            data: Bytes::from_static(b"client hello bytes"),
818        })
819        .encode(&mut plaintext);
820
821        // The Initial Length field includes the packet number and AEAD tag.
822        let ciphertext_len = plaintext.len() + keys.client.packet_key.tag_len();
823        let mut header = vec![0xc0, 0, 0, 0, 1, dcid.len() as u8];
824        header.extend_from_slice(dcid.as_bytes());
825        header.push(scid.len() as u8);
826        header.extend_from_slice(&scid);
827        header.push(0); // empty token
828        VarInt::from_u64((1 + ciphertext_len) as u64)
829            .unwrap()
830            .encode(&mut header);
831        let packet_number_offset = header.len();
832        header.push(0); // packet number 0, one byte
833
834        keys.client
835            .packet_key
836            .encrypt(0, &header, &mut plaintext)
837            .unwrap();
838        header.extend_from_slice(&plaintext);
839        let sample: [u8; 16] = header[packet_number_offset + 4..packet_number_offset + 20]
840            .try_into()
841            .unwrap();
842        let mask = keys.client.header_key.protection_mask(&sample).unwrap();
843        apply_header_protection(&mut header, &mask, packet_number_offset, 1);
844        header
845    }
846
847    #[test]
848    fn decrypts_client_initial_and_decodes_crypto_frames() {
849        let packet = protected_client_initial();
850        let initial = decrypt_client_initial(&packet).unwrap();
851        assert_eq!(initial.packet_number, 0);
852        assert!(matches!(initial.frames.first(), Some(Frame::Ping)));
853        assert!(
854            matches!(initial.frames.get(1), Some(Frame::Crypto(frame)) if frame.data == Bytes::from_static(b"client hello bytes"))
855        );
856    }
857
858    #[test]
859    fn rejects_tampered_client_initial() {
860        let mut packet = protected_client_initial();
861        let last = packet.len() - 1;
862        packet[last] ^= 0x01;
863        assert!(decrypt_client_initial(&packet).is_err());
864    }
865}