Skip to main content

hopr_protocol_session/processing/
segmenter.rs

1//! This module defines a [`Segmenter`] adaptor for [`futures::Sink`].
2use std::{
3    collections::VecDeque,
4    pin::Pin,
5    task::{Context, Poll},
6};
7
8use tracing::instrument;
9
10use crate::{
11    protocol::{FrameId, Segment, SeqIndicator, SessionMessage},
12    utils::segment_into,
13};
14
15/// Segmenter is an adaptor to [`futures::Sink`] of [`Segment`] items
16/// that turns it into [`futures::io::AsyncWrite`].
17///
18/// Bytes written to the Segmenter are buffered up and/or chopped into [`Segment`]
19/// of at most `C` in payload size (the `data` member).
20///
21/// The bytes are written to the
22/// underlying Sink once more than `frame_size` is written (or unless flushed),
23/// so the Segmenter naturally acts as a buffered writer.
24/// Any unflushed bytes written to the Segmenter will be lost when it is closed.
25///
26/// The data are grouped into [`Frames`](crate::protocol::Frame) of the size given by the `frame_size`
27/// parameter, segments in each such group share the same [`FrameId`].
28/// This acts as a natural buffering feature of a Segmenter.
29///
30/// Segmenter can optionally send a [terminating](Segment::terminating) when `poll_close`
31/// is called.
32///
33/// Segmenter is essentially inverse of [`Reassembler`](super::reassembly::Reassembler).
34///
35/// Use [`SegmenterExt`] to turn a `Segment` sink into an `AsyncWrite` object using the `Segmenter`.
36#[must_use = "sinks do nothing unless polled"]
37#[pin_project::pin_project]
38pub struct Segmenter<const C: usize, S> {
39    #[pin]
40    inner: S,
41    state: State,
42    frame: Vec<u8>,
43    ready_segments: VecDeque<Segment>,
44    frame_size: usize,
45    frame_id: FrameId,
46    is_closed: bool,
47    send_terminating_segment: bool,
48    /// Datagram mode: emit exactly one frame per write (preserve datagram boundaries) instead of
49    /// coalescing writes up to `frame_size`. On stateless sessions the `hopr-transport-session`
50    /// crate enables this for the `NoDelay` (UDP-like) capability.
51    datagram: bool,
52}
53
54enum State {
55    BufferingFrame,
56    WritingFrame,
57}
58
59impl<const C: usize, S> Segmenter<C, S>
60where
61    S: futures::Sink<Segment>,
62    S::Error: std::error::Error + Send + Sync + 'static,
63{
64    fn new(inner: S, frame_size: usize, send_terminating_segment: bool, datagram: bool) -> Self {
65        // Clamp frame_size to [SESSION_MTU, SESSION_MTU * (SeqIndicator::MAX + 1)].
66        // Minimum is SESSION_MTU (= C - SEGMENT_OVERHEAD) so that a single frame fits in one
67        // HOPR packet (1 segment). Maximum is bounded by SeqIndicator capacity.
68        //
69        // In datagram mode `frame_size` no longer bounds a frame (each write is its own frame);
70        // an individual datagram may be up to the same SeqIndicator-bounded maximum. `frame_size`
71        // is then only a capacity hint for the frame buffer.
72        let frame_size = frame_size.clamp(
73            C - SessionMessage::<C>::SEGMENT_OVERHEAD,
74            (C - SessionMessage::<C>::SEGMENT_OVERHEAD) * (SeqIndicator::MAX + 1) as usize,
75        );
76
77        Self {
78            inner,
79            state: State::BufferingFrame,
80            frame: Vec::with_capacity(frame_size),
81            ready_segments: VecDeque::with_capacity(frame_size.div_ceil(C - SessionMessage::<C>::SEGMENT_OVERHEAD)),
82            frame_size,
83            frame_id: 1,
84            is_closed: false,
85            send_terminating_segment,
86            datagram,
87        }
88    }
89}
90
91impl<const C: usize, S> futures::io::AsyncWrite for Segmenter<C, S>
92where
93    S: futures::Sink<Segment>,
94    S::Error: std::error::Error + Send + Sync + 'static,
95{
96    #[instrument(name = "Segmenter::poll_write", level = "trace", skip(self, cx, buf), fields(frame_id = self.frame_id, buf_len = buf.len(), frame_size = self.frame.len(), ready_segments = self.ready_segments.len()), ret)]
97    fn poll_write(self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &[u8]) -> Poll<std::io::Result<usize>> {
98        if self.is_closed {
99            return Poll::Ready(Err(std::io::Error::new(
100                std::io::ErrorKind::BrokenPipe,
101                "segmenter closed",
102            )));
103        }
104
105        let mut this = self.project();
106        loop {
107            match this.state {
108                State::BufferingFrame => {
109                    // Datagram mode: emit exactly one frame per write so datagram boundaries are
110                    // preserved (the peer reads one datagram per read), regardless of `frame_size`.
111                    // Segment `buf` directly — there is nothing to accumulate, so the `frame`
112                    // buffer is left untouched. Segments drain on the next poll_write / poll_flush.
113                    if *this.datagram {
114                        if buf.is_empty() {
115                            return Poll::Ready(Ok(0));
116                        }
117                        segment_into(
118                            buf,
119                            C - SessionMessage::<C>::SEGMENT_OVERHEAD,
120                            *this.frame_id,
121                            this.ready_segments,
122                        )
123                        .map_err(std::io::Error::other)?;
124
125                        tracing::trace!(
126                            num_segments = this.ready_segments.len(),
127                            datagram_len = buf.len(),
128                            "datagram frame ready"
129                        );
130
131                        *this.frame_id += 1;
132                        *this.state = State::WritingFrame;
133
134                        return Poll::Ready(Ok(buf.len()));
135                    }
136
137                    // If there's space in the frame, keep writing to it
138                    if *this.frame_size > this.frame.len() {
139                        let to_write = buf.len().min(*this.frame_size - this.frame.len());
140                        this.frame.extend_from_slice(&buf[..to_write]);
141
142                        return Poll::Ready(Ok(to_write));
143                    } else {
144                        // No more space in the frame buffer, we need to segment it
145                        // and write segments to the downstream
146                        segment_into(
147                            this.frame.as_slice(),
148                            C - SessionMessage::<C>::SEGMENT_OVERHEAD,
149                            *this.frame_id,
150                            this.ready_segments,
151                        )
152                        .map_err(std::io::Error::other)?;
153
154                        tracing::trace!(num_segments = this.ready_segments.len(), "frame ready");
155
156                        this.frame.clear();
157                        *this.frame_id += 1;
158                        *this.state = State::WritingFrame;
159                    }
160                }
161                State::WritingFrame => {
162                    if !this.ready_segments.is_empty() {
163                        // Keep writing segments to downstream
164                        futures::ready!(this.inner.as_mut().poll_ready(cx).map_err(std::io::Error::other))?;
165
166                        let segment = this.ready_segments.pop_front().unwrap();
167                        tracing::trace!(seg_id = %segment.id(), "segment goes out");
168                        this.inner.as_mut().start_send(segment).map_err(std::io::Error::other)?;
169                    } else {
170                        // Once we're done, we can buffer another frame
171                        *this.state = State::BufferingFrame;
172                        tracing::trace!("all segments out");
173                    }
174                }
175            }
176        }
177    }
178
179    #[instrument(name = "Segmenter::poll_flush", level = "trace", skip(self, cx), fields(frame_id = self.frame_id, frame_size = self.frame.len(), ready_segments = self.ready_segments.len()), ret)]
180    fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
181        if self.is_closed {
182            return Poll::Ready(Err(std::io::Error::new(
183                std::io::ErrorKind::BrokenPipe,
184                "segmenter closed",
185            )));
186        }
187
188        let mut this = self.project();
189        loop {
190            // If there's any data in the unfinished frame, segment it
191            if !this.frame.is_empty() {
192                // Flush the downstream sink first
193                futures::ready!(this.inner.as_mut().poll_flush(cx).map_err(std::io::Error::other))?;
194
195                // Segment whatever data is in the frame
196                // At this point ready_segments must be empty,
197                // because poll_write always makes sure it is before returning Ready
198                segment_into(
199                    this.frame.as_slice(),
200                    C - SessionMessage::<C>::SEGMENT_OVERHEAD,
201                    *this.frame_id,
202                    this.ready_segments,
203                )
204                .map_err(std::io::Error::other)?;
205
206                tracing::trace!(num_segments = this.ready_segments.len(), "flushed frame ready");
207
208                this.frame.clear();
209                *this.frame_id += 1;
210            } else if !this.ready_segments.is_empty() {
211                futures::ready!(this.inner.as_mut().poll_ready(cx).map_err(std::io::Error::other))?;
212
213                let segment = this.ready_segments.pop_front().unwrap();
214                tracing::trace!(seg_id = %segment.id(), "segment flushing out");
215
216                this.inner.as_mut().start_send(segment).map_err(std::io::Error::other)?;
217            } else {
218                // Both buffers are empty, so only flush the downstream
219                futures::ready!(this.inner.as_mut().poll_flush(cx).map_err(std::io::Error::other))?;
220
221                tracing::trace!("all segments flushed out");
222                return Poll::Ready(Ok(()));
223            }
224        }
225    }
226
227    #[instrument(name = "Segmenter::poll_close", level = "trace", skip(self, cx), fields(frame_id = self.frame_id) , ret)]
228    fn poll_close(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
229        let mut this = self.project();
230
231        if *this.send_terminating_segment && !*this.is_closed {
232            futures::ready!(this.inner.as_mut().poll_ready(cx).map_err(std::io::Error::other))?;
233            let dummy = Segment::terminating(*this.frame_id);
234            this.inner.as_mut().start_send(dummy).map_err(std::io::Error::other)?;
235            tracing::trace!("sent terminating segment");
236        }
237
238        *this.is_closed = true;
239        this.inner.as_mut().poll_close(cx).map_err(std::io::Error::other)
240    }
241}
242
243/// Sink extension methods for segmenting binary data into a sink.
244pub trait SegmenterExt: futures::Sink<Segment> {
245    /// Attaches a [`Segmenter`] to the underlying sink.
246    fn segmenter<const C: usize>(self, frame_size: usize) -> Segmenter<C, Self>
247    where
248        Self: Sized,
249        Self::Error: std::error::Error + Send + Sync + 'static,
250    {
251        Segmenter::new(self, frame_size, false, false)
252    }
253
254    /// Attaches a [`Segmenter`] to the underlying sink.
255    /// The `Segmenter` also sends a [terminating](Segment::terminating) when closed.
256    ///
257    /// When `datagram` is set, each write is emitted as exactly one frame (datagram boundaries are
258    /// preserved) instead of being coalesced up to `frame_size`.
259    fn segmenter_with_terminating_segment<const C: usize>(self, frame_size: usize, datagram: bool) -> Segmenter<C, Self>
260    where
261        Self: Sized,
262        Self::Error: std::error::Error + Send + Sync + 'static,
263    {
264        Segmenter::new(self, frame_size, true, datagram)
265    }
266}
267
268impl<T: ?Sized> SegmenterExt for T where T: futures::Sink<Segment> {}
269
270#[cfg(test)]
271mod tests {
272    use anyhow::{Context, anyhow};
273    use futures::{AsyncWriteExt, Stream, StreamExt, pin_mut};
274    use futures_time::future::FutureExt;
275
276    use super::*;
277    use crate::{protocol::SeqNum, utils::segment};
278
279    const MTU: usize = 1000;
280    const SMTU: usize = MTU - SessionMessage::<MTU>::SEGMENT_OVERHEAD;
281    const FRAME_SIZE: usize = 1500;
282
283    const SEGMENTS_PER_FRAME: usize = FRAME_SIZE / MTU + 1;
284
285    async fn assert_frame_segments(
286        start_frame_id: FrameId,
287        num_frames: usize,
288        segments: &mut (impl Stream<Item = Segment> + Unpin),
289        data: &[u8],
290    ) -> anyhow::Result<()> {
291        for i in 0..num_frames * SEGMENTS_PER_FRAME {
292            let start_frame_id = start_frame_id as usize;
293            let frame_id = i / SEGMENTS_PER_FRAME + start_frame_id;
294            tracing::debug!("testing frame id {frame_id} {}", (i % SEGMENTS_PER_FRAME) as SeqNum);
295
296            let seg = segments
297                .next()
298                .timeout(futures_time::time::Duration::from_millis(500))
299                .await
300                .context(format!("assert_frame_segments {i}"))?
301                .ok_or(anyhow!("no more segments"))?;
302
303            assert_eq!(frame_id as FrameId, seg.frame_id);
304            assert_eq!((i % SEGMENTS_PER_FRAME) as SeqNum, seg.seq_idx);
305            assert_eq!((FRAME_SIZE / MTU + 1) as SeqNum, seg.seq_flags.seq_len());
306            if i % SEGMENTS_PER_FRAME == 0 {
307                assert_eq!(SMTU, seg.data.len());
308                assert_eq!(
309                    &data[(frame_id - start_frame_id) * FRAME_SIZE + i % SEGMENTS_PER_FRAME * SMTU
310                        ..(frame_id - start_frame_id) * FRAME_SIZE + i % SEGMENTS_PER_FRAME * SMTU + SMTU],
311                    seg.data.as_ref()
312                );
313            } else {
314                assert_eq!(FRAME_SIZE % SMTU, seg.data.len());
315                assert_eq!(
316                    &data[(frame_id - start_frame_id) * FRAME_SIZE + i % SEGMENTS_PER_FRAME * SMTU
317                        ..(frame_id - start_frame_id) * FRAME_SIZE + i % SEGMENTS_PER_FRAME * SMTU + FRAME_SIZE % SMTU],
318                    seg.data.as_ref()
319                );
320            }
321        }
322
323        Ok(())
324    }
325
326    #[tokio::test]
327    async fn segmenter_should_not_segment_small_data_unless_flushed() -> anyhow::Result<()> {
328        let (segments_tx, segments) = futures::channel::mpsc::unbounded();
329        let mut writer = segments_tx.segmenter::<MTU>(FRAME_SIZE);
330
331        writer.write_all(b"test").await?;
332
333        pin_mut!(segments);
334        segments
335            .next()
336            .timeout(futures_time::time::Duration::from_millis(10))
337            .await
338            .expect_err("should time out");
339
340        writer.flush().await?;
341
342        let seg = segments.next().await.ok_or(anyhow!("no more segments"))?;
343        assert_eq!(1, seg.frame_id);
344        assert_eq!(1, seg.seq_flags.seq_len());
345        assert_eq!(0, seg.seq_idx);
346        assert_eq!(b"test", seg.data.as_ref());
347
348        Ok(())
349    }
350
351    #[tokio::test]
352    async fn datagram_mode_emits_one_frame_per_write() -> anyhow::Result<()> {
353        let (segments_tx, segments) = futures::channel::mpsc::unbounded();
354        let mut writer = segments_tx.segmenter_with_terminating_segment::<MTU>(FRAME_SIZE, true);
355        pin_mut!(segments);
356
357        // Two small datagrams that would coalesce into a single frame in byte-stream mode. In
358        // datagram mode each write is its own frame (distinct frame_id), so the boundary is kept.
359        writer.write_all(b"hello").await?;
360        writer.write_all(b"world").await?;
361        writer.flush().await?;
362
363        let first = segments.next().await.ok_or(anyhow!("no first datagram"))?;
364        assert_eq!(1, first.frame_id);
365        assert_eq!(0, first.seq_idx);
366        assert_eq!(1, first.seq_flags.seq_len(), "single-segment datagram");
367        assert_eq!(b"hello", first.data.as_ref());
368
369        let second = segments.next().await.ok_or(anyhow!("no second datagram"))?;
370        assert_eq!(2, second.frame_id, "each write is a distinct frame");
371        assert_eq!(0, second.seq_idx);
372        assert_eq!(1, second.seq_flags.seq_len());
373        assert_eq!(b"world", second.data.as_ref());
374
375        Ok(())
376    }
377
378    #[tokio::test]
379    async fn datagram_mode_keeps_a_multi_segment_datagram_in_one_frame() -> anyhow::Result<()> {
380        let (segments_tx, segments) = futures::channel::mpsc::unbounded();
381        let mut writer = segments_tx.segmenter_with_terminating_segment::<MTU>(FRAME_SIZE, true);
382        pin_mut!(segments);
383
384        // A datagram larger than one segment (and larger than frame_size) stays a single frame: all
385        // segments share one frame_id/seq_len with contiguous indices, matching a direct segment().
386        let datagram = vec![0xABu8; SMTU * 2 + 10];
387        writer.write_all(&datagram).await?;
388        writer.flush().await?;
389
390        for expected in segment(&datagram, SMTU, 1)? {
391            let seg = segments
392                .next()
393                .timeout(futures_time::time::Duration::from_millis(500))
394                .await
395                .context("multi-segment datagram")?
396                .ok_or(anyhow!("no more segments"))?;
397            assert_eq!(1, seg.frame_id, "one frame_id for the whole datagram");
398            assert_eq!(expected.seq_idx, seg.seq_idx);
399            assert_eq!(expected.seq_flags.seq_len(), seg.seq_flags.seq_len());
400            assert_eq!(expected.data.as_ref(), seg.data.as_ref());
401        }
402
403        Ok(())
404    }
405
406    #[parameterized::parameterized(num_frames = { 1, 3, 5, 11 })]
407    #[parameterized_macro(tokio::test)]
408    async fn segmenter_should_segment_complete_frames(num_frames: usize) -> anyhow::Result<()> {
409        let (segments_tx, segments) = futures::channel::mpsc::unbounded();
410        let mut writer = segments_tx.segmenter::<MTU>(FRAME_SIZE);
411
412        let mut all_data = Vec::new();
413        for _ in 0..num_frames {
414            let data = hopr_types::crypto_random::random_bytes::<FRAME_SIZE>();
415            writer.write_all(&data).await?;
416            all_data.extend(data);
417        }
418
419        writer.flush().await?;
420
421        pin_mut!(segments);
422        assert_frame_segments(1, num_frames, &mut segments, &all_data).await?;
423
424        writer.close().await?;
425
426        assert_eq!(None, segments.next().await);
427        Ok(())
428    }
429
430    #[tokio::test]
431    async fn segmenter_full_frame_segmentation_must_be_consistent_with_segment_function() -> anyhow::Result<()> {
432        let (segments_tx, segments) = futures::channel::mpsc::unbounded();
433        let mut writer = segments_tx.segmenter::<MTU>(FRAME_SIZE);
434
435        let data = hopr_types::crypto_random::random_bytes::<FRAME_SIZE>();
436
437        writer.write_all(&data).await?;
438        writer.flush().await?;
439        writer.close().await?;
440
441        // Segmenter already takes into account the SessionMessage overhead
442        let expected = segment(data, SMTU, 1)?;
443        let actual = segments.collect::<Vec<_>>().await;
444
445        assert_eq!(expected, actual);
446
447        Ok(())
448    }
449
450    #[test_log::test(tokio::test)]
451    async fn segmenter_full_frame_segmentation_must_also_include_terminating_segment() -> anyhow::Result<()> {
452        let (segments_tx, segments) = futures::channel::mpsc::unbounded();
453        let mut writer = segments_tx.segmenter_with_terminating_segment::<MTU>(FRAME_SIZE, false);
454
455        let data = hopr_types::crypto_random::random_bytes::<FRAME_SIZE>();
456
457        writer.write_all(&data).await?;
458        writer.flush().await?;
459        writer.close().await?;
460
461        // Segmenter already takes into account the SessionMessage overhead
462        let mut expected = segment(data, SMTU, 1)?;
463        expected.push(Segment::terminating(2));
464        let actual = segments.collect::<Vec<_>>().await;
465
466        assert_eq!(expected, actual);
467
468        Ok(())
469    }
470
471    #[test_log::test(tokio::test)]
472    async fn segmenter_should_segment_complete_frame_with_misaligned_mtu() -> anyhow::Result<()> {
473        let (segments_tx, segments) = futures::channel::mpsc::unbounded();
474        let mut writer = segments_tx.segmenter::<MTU>(FRAME_SIZE);
475
476        // Make sure the FRAME_SIZE is not a multiple of MTU
477        assert_ne!(0, FRAME_SIZE % MTU);
478
479        let data = hopr_types::crypto_random::random_bytes::<FRAME_SIZE>();
480        writer.write_all(&data).await?;
481        writer.flush().await?;
482        writer.close().await?;
483
484        pin_mut!(segments);
485
486        for i in 0..(FRAME_SIZE / MTU) {
487            let seg = segments.next().await.ok_or(anyhow!("no more segments"))?;
488            assert_eq!(1, seg.frame_id);
489            assert_eq!(i as SeqNum, seg.seq_idx);
490            assert_eq!(((FRAME_SIZE / SMTU) + 1) as SeqNum, seg.seq_flags.seq_len());
491            assert_eq!(SMTU, seg.data.len());
492            assert_eq!(&data[i * SMTU..i * SMTU + SMTU], seg.data.as_ref());
493        }
494
495        let seg = segments.next().await.ok_or(anyhow!("no more segments"))?;
496        assert_eq!(1, seg.frame_id);
497        assert_eq!((FRAME_SIZE / SMTU) as SeqNum, seg.seq_idx);
498        assert_eq!(((FRAME_SIZE / SMTU) + 1) as SeqNum, seg.seq_flags.seq_len());
499        assert_eq!(FRAME_SIZE % SMTU, seg.data.len());
500        assert_eq!(&data[FRAME_SIZE - FRAME_SIZE % SMTU..], seg.data.as_ref());
501
502        assert_eq!(None, segments.next().await);
503        Ok(())
504    }
505
506    #[test_log::test(tokio::test)]
507    async fn segmenter_should_segment_multiple_complete_frames_and_incomplete_frame_on_flush() -> anyhow::Result<()> {
508        let (segments_tx, segments) = futures::channel::mpsc::unbounded();
509        let mut writer = segments_tx.segmenter::<MTU>(FRAME_SIZE);
510
511        let data = hopr_types::crypto_random::random_bytes::<{ FRAME_SIZE + 4 }>();
512        writer.write_all(&data).await?;
513
514        pin_mut!(segments);
515
516        // The first frame should come out even without a flush
517        assert_frame_segments(1, 1, &mut segments, &data).await?;
518
519        // And no more segment comes out for the remaining bytes
520        segments
521            .next()
522            .timeout(futures_time::time::Duration::from_millis(10))
523            .await
524            .expect_err("should time out");
525
526        // ... until it is flushed
527        writer.flush().await?;
528
529        let seg = segments
530            .next()
531            .timeout(futures_time::time::Duration::from_millis(500))
532            .await?
533            .ok_or(anyhow!("no more segments"))?;
534        assert_eq!(2, seg.frame_id);
535        assert_eq!(0, seg.seq_idx);
536        assert_eq!(1, seg.seq_flags.seq_len());
537        assert_eq!(4, seg.data.len());
538        assert_eq!(&data[FRAME_SIZE..], seg.data.as_ref());
539
540        // The next full frame should come out normally after a flush
541        let data = hopr_types::crypto_random::random_bytes::<FRAME_SIZE>();
542        writer.write_all(&data).await?;
543        writer.flush().await?;
544
545        assert_frame_segments(3, 1, &mut segments, &data).await?;
546
547        Ok(())
548    }
549
550    #[test_log::test(tokio::test)]
551    async fn segmenter_should_work_with_buffering_backend() -> anyhow::Result<()> {
552        let (tx, rx) = futures::channel::mpsc::channel(5);
553        let mut writer = tx.segmenter::<MTU>(FRAME_SIZE);
554
555        let data = hopr_types::crypto_random::random_bytes::<{ 10 * FRAME_SIZE }>();
556
557        let jh_recv = tokio::task::spawn(
558            rx.collect::<Vec<_>>()
559                .delay(futures_time::time::Duration::from_millis(200)),
560        );
561        let jh_send = tokio::task::spawn(async move {
562            writer.write_all(&data).await?;
563            writer.flush().await?;
564            writer.close().await?;
565            Ok::<_, std::io::Error>(())
566        });
567
568        let (segments, send_res) = futures::future::try_join(jh_recv, jh_send).await?;
569        send_res?;
570
571        assert_frame_segments(1, 10, &mut futures::stream::iter(segments), &data).await
572    }
573}