1use 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#[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: 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 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 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 *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 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 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 *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 !this.frame.is_empty() {
192 futures::ready!(this.inner.as_mut().poll_flush(cx).map_err(std::io::Error::other))?;
194
195 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 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
243pub trait SegmenterExt: futures::Sink<Segment> {
245 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 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 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 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 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 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 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 assert_frame_segments(1, 1, &mut segments, &data).await?;
518
519 segments
521 .next()
522 .timeout(futures_time::time::Duration::from_millis(10))
523 .await
524 .expect_err("should time out");
525
526 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 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}