Skip to main content

capnp/
serialize.rs

1// Copyright (c) 2013-2015 Sandstorm Development Group, Inc. and contributors
2// Licensed under the MIT License:
3//
4// Permission is hereby granted, free of charge, to any person obtaining a copy
5// of this software and associated documentation files (the "Software"), to deal
6// in the Software without restriction, including without limitation the rights
7// to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
8// copies of the Software, and to permit persons to whom the Software is
9// furnished to do so, subject to the following conditions:
10//
11// The above copyright notice and this permission notice shall be included in
12// all copies or substantial portions of the Software.
13//
14// THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
15// IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
16// FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
17// AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
18// LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
19// OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN
20// THE SOFTWARE.
21
22//! Reading and writing of messages using the
23//! [standard stream framing](https://capnproto.org/encoding.html#serialization-over-a-stream).
24//!
25//! Each message is preceded by a segment table indicating the size of its segments.
26
27pub(crate) mod no_alloc_buffer_segments;
28pub use no_alloc_buffer_segments::{
29    NoAllocBufferSegments, NoAllocSegmentTableInfo, NoAllocSliceSegments,
30};
31
32use crate::io::{Read, Write};
33
34use crate::message;
35use crate::private::units::BYTES_PER_WORD;
36use crate::Result;
37use crate::{Error, ErrorKind};
38
39pub const SEGMENTS_COUNT_LIMIT: usize = 512;
40
41/// Segments read from a single flat slice of words.
42#[cfg(feature = "alloc")]
43type SliceSegments<'a> = BufferSegments<&'a [u8]>;
44
45/// Reads a serialized message (including a segment table) from a flat slice of bytes, without copying.
46///
47/// The slice is allowed to extend beyond the end of the message. On success, updates `slice` to point
48/// to the remaining bytes beyond the end of the message.
49///
50/// ALIGNMENT: If the "unaligned" feature is enabled, then there are no alignment requirements on `slice`.
51/// Otherwise, `slice` must be 8-byte aligned (attempts to read the message will trigger errors).
52#[cfg(feature = "alloc")]
53pub fn read_message_from_flat_slice<'a>(
54    slice: &mut &'a [u8],
55    options: message::ReaderOptions,
56) -> Result<message::Reader<BufferSegments<&'a [u8]>>> {
57    let all_bytes = *slice;
58    let mut bytes = *slice;
59    let orig_bytes_len = bytes.len();
60    let Some(segment_lengths_builder) = read_segment_table(&mut bytes, options)? else {
61        return Err(Error::from_kind(ErrorKind::EmptySlice));
62    };
63    let segment_table_bytes_len = orig_bytes_len - bytes.len();
64    assert_eq!(segment_table_bytes_len % BYTES_PER_WORD, 0);
65    let num_words = segment_lengths_builder.total_words();
66    let body_bytes = &all_bytes[segment_table_bytes_len..];
67    if num_words > (body_bytes.len() / BYTES_PER_WORD) {
68        Err(Error::from_kind(ErrorKind::MessageEndsPrematurely(
69            num_words,
70            body_bytes.len() / BYTES_PER_WORD,
71        )))
72    } else {
73        *slice = &body_bytes[(num_words * BYTES_PER_WORD)..];
74        Ok(message::Reader::new(
75            segment_lengths_builder.into_slice_segments(all_bytes, segment_table_bytes_len),
76            options,
77        ))
78    }
79}
80
81/// Reads a serialized message (including a segment table) from a flat slice of bytes, without copying.
82///
83/// The slice is allowed to extend beyond the end of the message. On success, updates `slice` to point
84/// to the remaining bytes beyond the end of the message.
85///
86/// Unlike [read_message_from_flat_slice], it does not do heap allocation.
87///
88/// ALIGNMENT: If the "unaligned" feature is enabled, then there are no alignment requirements on `slice`.
89/// Otherwise, `slice` must be 8-byte aligned (attempts to read the message will trigger errors).
90pub fn read_message_from_flat_slice_no_alloc<'a>(
91    slice: &mut &'a [u8],
92    options: message::ReaderOptions,
93) -> Result<message::Reader<NoAllocSliceSegments<'a>>> {
94    let segments = NoAllocSliceSegments::from_slice(slice, options)?;
95
96    Ok(message::Reader::new(segments, options))
97}
98
99/// Segments read from a buffer, useful for when you have the message in a buffer and don't want the extra
100/// copy performed by [`read_message`].
101#[cfg(feature = "alloc")]
102pub struct BufferSegments<T> {
103    buffer: T,
104
105    // Number of bytes in the segment table.
106    segment_table_bytes_len: usize,
107
108    // Each pair represents a segment inside of `buffer`:
109    // (starting index (in words), ending index (in words)),
110    // where the indices are relative to the end of the segment table.
111    segment_indices: alloc::vec::Vec<(usize, usize)>,
112}
113
114#[cfg(feature = "alloc")]
115impl<T: core::ops::Deref<Target = [u8]>> BufferSegments<T> {
116    /// Reads a serialized message (including a segment table) from a buffer and takes ownership, without copying.
117    /// The buffer is allowed to be longer than the message. Provide this to `Reader::new` with options that make
118    /// sense for your use case. Very long lived mmaps may need unlimited traversal limit.
119    ///
120    /// ALIGNMENT: If the "unaligned" feature is enabled, then there are no alignment requirements on `buffer`.
121    /// Otherwise, `buffer` must be 8-byte aligned (attempts to read the message will trigger errors).
122    pub fn new(buffer: T, options: message::ReaderOptions) -> Result<Self> {
123        let mut segment_bytes = &*buffer;
124
125        let Some(segment_table) = read_segment_table(&mut segment_bytes, options)? else {
126            return Err(Error::from_kind(ErrorKind::EmptyBuffer));
127        };
128        let segment_table_bytes_len = buffer.len() - segment_bytes.len();
129
130        if segment_table.total_words() * 8 > segment_bytes.len() {
131            return Err(Error::from_kind(ErrorKind::MessageEndsPrematurely(
132                segment_table.total_words(),
133                segment_bytes.len() / 8,
134            )));
135        }
136
137        let segment_indices = segment_table.to_segment_indices();
138        Ok(Self {
139            buffer,
140            segment_table_bytes_len,
141            segment_indices,
142        })
143    }
144
145    pub fn into_buffer(self) -> T {
146        self.buffer
147    }
148}
149
150#[cfg(feature = "alloc")]
151impl<T: core::ops::Deref<Target = [u8]>> message::ReaderSegments for BufferSegments<T> {
152    fn get_segment(&self, id: u32) -> Option<&[u8]> {
153        let id_usize = id as usize;
154        if id_usize < self.segment_indices.len() {
155            let (a, b) = self.segment_indices[id_usize];
156            Some(
157                &self.buffer[(self.segment_table_bytes_len + a * BYTES_PER_WORD)
158                    ..(self.segment_table_bytes_len + b * BYTES_PER_WORD)],
159            )
160        } else {
161            None
162        }
163    }
164
165    fn len(&self) -> usize {
166        self.segment_indices.len()
167    }
168}
169
170/// Owned memory containing a message's segments sequentialized in a single contiguous buffer.
171/// The segments are guaranteed to be 8-byte aligned.
172#[cfg(feature = "alloc")]
173pub struct OwnedSegments {
174    // Each pair represents a segment inside of `owned_space`.
175    // (starting index (in words), ending index (in words))
176    segment_indices: alloc::vec::Vec<(usize, usize)>,
177
178    owned_space: alloc::vec::Vec<crate::Word>,
179}
180
181#[cfg(feature = "alloc")]
182impl core::ops::Deref for OwnedSegments {
183    type Target = [u8];
184    fn deref(&self) -> &[u8] {
185        crate::Word::words_to_bytes(&self.owned_space[..])
186    }
187}
188
189#[cfg(feature = "alloc")]
190impl core::ops::DerefMut for OwnedSegments {
191    fn deref_mut(&mut self) -> &mut [u8] {
192        crate::Word::words_to_bytes_mut(&mut self.owned_space[..])
193    }
194}
195
196#[cfg(feature = "alloc")]
197impl crate::message::ReaderSegments for OwnedSegments {
198    fn get_segment(&self, id: u32) -> Option<&[u8]> {
199        let id_usize = id as usize;
200        if id_usize < self.segment_indices.len() {
201            let (a, b) = self.segment_indices[id_usize];
202            Some(&self[(a * BYTES_PER_WORD)..(b * BYTES_PER_WORD)])
203        } else {
204            None
205        }
206    }
207
208    fn len(&self) -> usize {
209        self.segment_indices.len()
210    }
211}
212
213#[cfg(feature = "alloc")]
214/// Helper object for constructing an `OwnedSegments` or a `SliceSegments`.
215pub struct SegmentLengthsBuilder {
216    segment_indices: alloc::vec::Vec<(usize, usize)>,
217    total_words: usize,
218}
219
220#[cfg(feature = "alloc")]
221impl SegmentLengthsBuilder {
222    /// Creates a new `SegmentLengthsBuilder`, initializing the segment_indices vector with
223    /// `Vec::with_capacitiy(capacity)`. `capacity` should equal the number of times that `push_segment()`
224    /// is expected to be called.
225    pub fn with_capacity(capacity: usize) -> Self {
226        Self {
227            segment_indices: alloc::vec::Vec::with_capacity(capacity),
228            total_words: 0,
229        }
230    }
231
232    /// Pushes a new segment length. The `n`th time (starting at 0) this is called specifies the length of
233    /// the segment with ID `n`. If the segment overflows the total word count, then this returns
234    /// a MessageSizeOverflow error.
235    pub fn try_push_segment(&mut self, length_in_words: usize) -> Result<()> {
236        let new_total_words = self
237            .total_words
238            .checked_add(length_in_words)
239            .ok_or_else(|| Error::from_kind(ErrorKind::MessageSizeOverflow))?;
240        self.segment_indices
241            .push((self.total_words, new_total_words));
242        self.total_words = new_total_words;
243        Ok(())
244    }
245
246    /// Constructs an `OwnedSegments`, allocating a single buffer of 8-byte aligned memory to hold
247    /// all segments.
248    pub fn into_owned_segments(self) -> OwnedSegments {
249        let owned_space = crate::Word::allocate_zeroed_vec(self.total_words);
250        OwnedSegments {
251            segment_indices: self.segment_indices,
252            owned_space,
253        }
254    }
255
256    /// Constructs a `SliceSegments`.
257    /// `slice` contains the full message (including the segment header).
258    pub fn into_slice_segments(
259        self,
260        slice: &[u8],
261        segment_table_bytes_len: usize,
262    ) -> SliceSegments<'_> {
263        assert!(self.total_words * BYTES_PER_WORD <= slice.len());
264        BufferSegments {
265            buffer: slice,
266            segment_table_bytes_len,
267            segment_indices: self.segment_indices,
268        }
269    }
270
271    /// Returns the sum of the lengths of the segments pushed so far.
272    pub fn total_words(&self) -> usize {
273        self.total_words
274    }
275
276    /// Returns the vector of segment indices. Each entry is a pair (start_word_index, end_word_index).
277    /// This method primarily exists to enable testing.
278    pub fn to_segment_indices(self) -> alloc::vec::Vec<(usize, usize)> {
279        self.segment_indices
280    }
281}
282
283/// Reads a serialized message from a stream with the provided options.
284///
285/// For optimal performance, `read` should be a buffered reader type.
286#[cfg(feature = "alloc")]
287pub fn read_message<R>(
288    mut read: R,
289    options: message::ReaderOptions,
290) -> Result<message::Reader<OwnedSegments>>
291where
292    R: Read,
293{
294    let Some(owned_segments_builder) = read_segment_table(&mut read, options)? else {
295        return Err(Error::from_kind(ErrorKind::PrematureEndOfFile));
296    };
297    read_segments(
298        &mut read,
299        owned_segments_builder.into_owned_segments(),
300        options,
301    )
302}
303
304/// Like `read_message()`, but returns None instead of an error if there are zero bytes left in
305/// `read`.
306///
307/// This is useful for reading a stream containing an unknown number of messages -- you
308/// call this function until it returns None.
309#[cfg(feature = "alloc")]
310pub fn try_read_message<R>(
311    mut read: R,
312    options: message::ReaderOptions,
313) -> Result<Option<message::Reader<OwnedSegments>>>
314where
315    R: Read,
316{
317    let Some(owned_segments_builder) = read_segment_table(&mut read, options)? else {
318        return Ok(None);
319    };
320    Ok(Some(read_segments(
321        &mut read,
322        owned_segments_builder.into_owned_segments(),
323        options,
324    )?))
325}
326
327/// Like `try_read_message()`, but does not allocate any memory.
328///
329/// Stores the message in `buffer`. Returns a `BufferNotLargeEnough`
330/// error if the buffer is not large enough.
331/// ALIGNMENT: If the "unaligned" feature is enabled, then there are no alignment requirements on `buffer`.
332/// Otherwise, `buffer` must be 8-byte aligned (attempts to read the message will trigger errors).
333pub fn try_read_message_no_alloc<R>(
334    mut read: R,
335    buffer: &mut [u8],
336    options: message::ReaderOptions,
337) -> Result<Option<message::Reader<NoAllocBufferSegments<&[u8]>>>>
338where
339    R: Read,
340{
341    if !cfg!(feature = "unaligned") && buffer.as_ptr() as usize % BYTES_PER_WORD != 0 {
342        return Err(Error::from_kind(ErrorKind::UnalignedSegment));
343    }
344
345    if buffer.len() < 8 {
346        return Err(Error::from_kind(ErrorKind::BufferNotLargeEnough));
347    }
348
349    // read the first Word, which contains segment_count and the 1st segment length
350    {
351        let n = read.read(&mut buffer[0..8])?;
352        if n == 0 {
353            // Clean EOF on message boundary
354            return Ok(None);
355        } else if n < 8 {
356            read.read_exact(&mut buffer[n..8])?;
357        }
358    }
359
360    let segment_count =
361        u32::from_le_bytes(buffer[0..4].try_into().unwrap()).wrapping_add(1) as usize;
362
363    if segment_count >= SEGMENTS_COUNT_LIMIT || segment_count == 0 {
364        return Err(Error::from_kind(ErrorKind::InvalidNumberOfSegments(
365            segment_count,
366        )));
367    }
368
369    let mut total_body_words: usize = u32::from_le_bytes(buffer[4..8].try_into().unwrap()) as usize;
370    let mut num_segment_counts_read = 1;
371    while num_segment_counts_read < segment_count {
372        let start = (num_segment_counts_read + 1) * 4;
373        let end = start + 8;
374        if buffer.len() < end {
375            return Err(Error::from_kind(ErrorKind::BufferNotLargeEnough));
376        }
377
378        read.read_exact(&mut buffer[start..end])?;
379
380        total_body_words = total_body_words
381            .checked_add(
382                u32::from_le_bytes(buffer[start..(start + 4)].try_into().unwrap()) as usize,
383            )
384            .ok_or_else(|| Error::from_kind(ErrorKind::MessageSizeOverflow))?;
385
386        num_segment_counts_read += 1;
387        if num_segment_counts_read < segment_count {
388            total_body_words = total_body_words
389                .checked_add(
390                    u32::from_le_bytes(buffer[(start + 4)..end].try_into().unwrap()) as usize,
391                )
392                .ok_or_else(|| Error::from_kind(ErrorKind::MessageSizeOverflow))?;
393        }
394        num_segment_counts_read += 1;
395    }
396
397    if let Some(limit) = options.traversal_limit_in_words {
398        if total_body_words > limit {
399            return Err(Error::from_kind(ErrorKind::MessageTooLarge(
400                total_body_words,
401            )));
402        }
403    }
404
405    let start = (num_segment_counts_read + 1) * 4;
406    let end = start + (total_body_words * 8);
407    if buffer.len() < end {
408        return Err(Error::from_kind(ErrorKind::BufferNotLargeEnough));
409    }
410    read.read_exact(&mut buffer[start..end])?;
411
412    let info = no_alloc_buffer_segments::NoAllocSegmentTableInfo {
413        segments_count: segment_count,
414        segment_table_length_bytes: (num_segment_counts_read + 1) * 4,
415        total_segments_length_bytes: total_body_words * 8,
416    };
417
418    let segments = NoAllocSliceSegments::from_segment_table_info(buffer, info);
419    Ok(Some(crate::message::Reader::new(segments, options)))
420}
421
422/// Like `read_message()`, but does not allocate.
423///
424/// Stores the message in `buffer`. Returns a `BufferNotLargeEnough`
425/// error if the buffer is not large enough.
426/// ALIGNMENT: If the "unaligned" feature is enabled, then there are no alignment requirements on `buffer`.
427/// Otherwise, `buffer` must be 8-byte aligned (attempts to read the message will trigger errors).
428pub fn read_message_no_alloc<R>(
429    read: R,
430    buffer: &mut [u8],
431    options: message::ReaderOptions,
432) -> Result<message::Reader<NoAllocBufferSegments<&[u8]>>>
433where
434    R: Read,
435{
436    match try_read_message_no_alloc(read, buffer, options)? {
437        Some(m) => Ok(m),
438        None => Err(Error::from_kind(ErrorKind::PrematureEndOfFile)),
439    }
440}
441
442/// Reads a segment table from `read` and returns the total number of words across all
443/// segments, as well as the segment offsets.
444///
445/// The segment table format for streams is defined in the Cap'n Proto
446/// [encoding spec](https://capnproto.org/encoding.html)
447#[cfg(feature = "alloc")]
448fn read_segment_table<R>(
449    read: &mut R,
450    options: message::ReaderOptions,
451) -> Result<Option<SegmentLengthsBuilder>>
452where
453    R: Read,
454{
455    // read the first Word, which contains segment_count and the 1st segment length
456    let mut buf: [u8; 8] = [0; 8];
457    {
458        let n = read.read(&mut buf[..])?;
459        if n == 0 {
460            // Clean EOF on message boundary
461            return Ok(None);
462        } else if n < 8 {
463            read.read_exact(&mut buf[n..])?;
464        }
465    }
466
467    let segment_count = u32::from_le_bytes(buf[0..4].try_into().unwrap()).wrapping_add(1) as usize;
468
469    if segment_count >= SEGMENTS_COUNT_LIMIT || segment_count == 0 {
470        return Err(Error::from_kind(ErrorKind::InvalidNumberOfSegments(
471            segment_count,
472        )));
473    }
474
475    let mut segment_lengths_builder = SegmentLengthsBuilder::with_capacity(segment_count);
476    segment_lengths_builder
477        .try_push_segment(u32::from_le_bytes(buf[4..8].try_into().unwrap()) as usize)?;
478    if segment_count > 1 {
479        if segment_count < 4 {
480            read.read_exact(&mut buf)?;
481            for idx in 0..(segment_count - 1) {
482                let segment_len =
483                    u32::from_le_bytes(buf[(idx * 4)..(idx + 1) * 4].try_into().unwrap()) as usize;
484                segment_lengths_builder.try_push_segment(segment_len)?;
485            }
486        } else {
487            let mut segment_sizes = vec![0u8; (segment_count & !1) * 4];
488            read.read_exact(&mut segment_sizes[..])?;
489            for idx in 0..(segment_count - 1) {
490                let segment_len =
491                    u32::from_le_bytes(segment_sizes[(idx * 4)..(idx + 1) * 4].try_into().unwrap())
492                        as usize;
493                segment_lengths_builder.try_push_segment(segment_len)?;
494            }
495        }
496    }
497
498    // Don't accept a message which the receiver couldn't possibly traverse without hitting the
499    // traversal limit. Without this check, a malicious client could transmit a very large segment
500    // size to make the receiver allocate excessive space and possibly crash.
501    if let Some(limit) = options.traversal_limit_in_words {
502        if segment_lengths_builder.total_words() > limit {
503            return Err(Error::from_kind(ErrorKind::MessageTooLarge(
504                segment_lengths_builder.total_words(),
505            )));
506        }
507    }
508
509    Ok(Some(segment_lengths_builder))
510}
511
512#[cfg(feature = "alloc")]
513/// Reads segments from `read`.
514fn read_segments<R>(
515    read: &mut R,
516    mut owned_segments: OwnedSegments,
517    options: message::ReaderOptions,
518) -> Result<message::Reader<OwnedSegments>>
519where
520    R: Read,
521{
522    read.read_exact(&mut owned_segments[..])?;
523    Ok(crate::message::Reader::new(owned_segments, options))
524}
525
526/// Constructs a flat vector containing the entire message, including a segment header.
527#[cfg(feature = "alloc")]
528pub fn write_message_to_words<A>(message: &message::Builder<A>) -> alloc::vec::Vec<u8>
529where
530    A: message::Allocator,
531{
532    flatten_segments(&*message.get_segments_for_output())
533}
534
535/// Like `write_message_to_words()`, but takes a `ReaderSegments`, allowing it to be
536/// used on `message::Reader` objects (via `into_segments()`).
537#[cfg(feature = "alloc")]
538pub fn write_message_segments_to_words<R>(message: &R) -> alloc::vec::Vec<u8>
539where
540    R: message::ReaderSegments,
541{
542    flatten_segments(message)
543}
544
545#[cfg(feature = "alloc")]
546fn flatten_segments<R: message::ReaderSegments + ?Sized>(segments: &R) -> alloc::vec::Vec<u8> {
547    let word_count = compute_serialized_size(segments);
548    let segment_count: u32 = segments.len().try_into().unwrap();
549    let table_size = segment_count / 2 + 1;
550    let mut result = alloc::vec::Vec::with_capacity(word_count);
551    result.resize(table_size as usize * BYTES_PER_WORD, 0);
552    {
553        let mut bytes = &mut result[..];
554        write_segment_table_internal(&mut bytes, segments).expect("Failed to write segment table.");
555    }
556    for i in 0..segment_count {
557        let segment = segments.get_segment(i).unwrap();
558        result.extend(segment);
559    }
560    debug_assert!(
561        result.capacity() == word_count,
562        "unexpected result capacity growth"
563    );
564    result
565}
566
567/// Writes the provided message to `write`.
568///
569/// For optimal performance, `write` should be a buffered writer. `flush()` will not be called on
570/// the writer.
571///
572/// The only source of errors from this function are `write.write_all()` calls. If you pass in
573/// a writer that never returns an error, then this function will never return an error.
574pub fn write_message<W, A>(mut write: W, message: &message::Builder<A>) -> Result<()>
575where
576    W: Write,
577    A: message::Allocator,
578{
579    let segments = message.get_segments_for_output();
580    write_segment_table(&mut write, &segments)?;
581    write_segments(&mut write, &segments)
582}
583
584/// Like `write_message()`, but takes a `ReaderSegments`, allowing it to be
585/// used on `message::Reader` objects (via `into_segments()`).
586pub fn write_message_segments<W, R>(mut write: W, segments: &R) -> Result<()>
587where
588    W: Write,
589    R: message::ReaderSegments,
590{
591    write_segment_table_internal(&mut write, segments)?;
592    write_segments(&mut write, segments)
593}
594
595fn write_segment_table<W>(write: &mut W, segments: &[&[u8]]) -> Result<()>
596where
597    W: Write,
598{
599    write_segment_table_internal(write, segments)
600}
601
602/// Writes a segment table to `write`.
603///
604/// `segments` must contain at least one segment.
605fn write_segment_table_internal<W, R>(write: &mut W, segments: &R) -> Result<()>
606where
607    W: Write,
608    R: message::ReaderSegments + ?Sized,
609{
610    let mut buf: [u8; 8] = [0; 8];
611    let segment_count: u32 = segments.len().try_into().unwrap();
612
613    // write the first Word, which contains segment_count and the 1st segment length
614    buf[0..4].copy_from_slice(&(segment_count - 1).to_le_bytes());
615    buf[4..8].copy_from_slice(
616        &u32::try_from(segments.get_segment(0).unwrap().len() / BYTES_PER_WORD)
617            .unwrap()
618            .to_le_bytes(),
619    );
620    write.write_all(&buf)?;
621
622    if segment_count > 1 {
623        if segment_count < 4 {
624            for idx in 1..segment_count {
625                buf[((idx - 1) * 4) as usize..(idx * 4) as usize].copy_from_slice(
626                    &u32::try_from(segments.get_segment(idx).unwrap().len() / BYTES_PER_WORD)
627                        .unwrap()
628                        .to_le_bytes(),
629                );
630            }
631            if segment_count == 2 {
632                for b in &mut buf[4..8] {
633                    *b = 0
634                }
635            }
636            write.write_all(&buf)?;
637        } else {
638            #[cfg(feature = "alloc")]
639            {
640                let mut buf = vec![0; (segment_count as usize & !1) * 4];
641                for idx in 1..segment_count {
642                    buf[((idx - 1) * 4) as usize..(idx * 4) as usize].copy_from_slice(
643                        &u32::try_from(segments.get_segment(idx).unwrap().len() / BYTES_PER_WORD)
644                            .unwrap()
645                            .to_le_bytes(),
646                    );
647                }
648                if segment_count % 2 == 0 {
649                    let start_idx = buf.len() - 4;
650                    for b in &mut buf[start_idx..] {
651                        *b = 0
652                    }
653                }
654                write.write_all(&buf)?;
655            }
656
657            #[cfg(not(feature = "alloc"))]
658            {
659                unreachable!("multi-segment message builders are not supported in no-alloc mode")
660            }
661        }
662    }
663    Ok(())
664}
665
666/// Writes segments to `write`.
667fn write_segments<W, R: message::ReaderSegments + ?Sized>(write: &mut W, segments: &R) -> Result<()>
668where
669    W: Write,
670{
671    for i in 0.. {
672        if let Some(segment) = segments.get_segment(i) {
673            write.write_all(segment)?;
674        } else {
675            break;
676        }
677    }
678    Ok(())
679}
680
681/// Returns the number of bytes required to serialize the message (including the
682/// segment table).
683fn compute_serialized_size<R: message::ReaderSegments + ?Sized>(segments: &R) -> usize {
684    // Table size
685    let len = segments.len();
686    let mut size = ((len / 2) + 1) * BYTES_PER_WORD;
687    for i in 0..len {
688        let segment = segments.get_segment(i.try_into().unwrap()).unwrap();
689        size += segment.len();
690    }
691    size
692}
693
694/// Returns the number of (8-byte) words required to serialize the message (including the
695/// segment table).
696///
697/// Multiply this by 8 (or `std::mem::size_of::<capnp::Word>()`) to get the number of bytes
698/// that [`write_message()`](fn.write_message.html) will write.
699pub fn compute_serialized_size_in_words<A>(message: &crate::message::Builder<A>) -> usize
700where
701    A: crate::message::Allocator,
702{
703    compute_serialized_size(&message.get_segments_for_output()) / BYTES_PER_WORD
704}
705
706#[cfg(feature = "alloc")]
707#[cfg(test)]
708pub mod test {
709    use crate::io::{Read, Write};
710
711    use quickcheck::{quickcheck, TestResult};
712
713    use super::{
714        flatten_segments, read_message, read_message_from_flat_slice, read_segment_table,
715        try_read_message, write_segment_table, write_segments,
716    };
717    use crate::message;
718    use crate::message::ReaderSegments;
719
720    /// Writes segments as if they were a Capnproto message.
721    pub fn write_message_segments<W>(write: &mut W, segments: &[alloc::vec::Vec<crate::Word>])
722    where
723        W: Write,
724    {
725        let borrowed_segments: &[&[u8]] = &segments
726            .iter()
727            .map(|segment| crate::Word::words_to_bytes(&segment[..]))
728            .collect::<alloc::vec::Vec<_>>()[..];
729        write_segment_table(write, borrowed_segments).unwrap();
730        write_segments(write, borrowed_segments).unwrap();
731    }
732
733    #[test]
734    fn try_read_empty() {
735        let mut buf: &[u8] = &[];
736        assert!(try_read_message(&mut buf, message::ReaderOptions::new())
737            .unwrap()
738            .is_none());
739    }
740
741    #[test]
742    fn test_read_segment_table() {
743        let mut buf = vec![];
744
745        buf.extend(
746            [
747                0, 0, 0, 0, // 1 segments
748                0, 0, 0, 0,
749            ], // 0 length
750        );
751        let segment_lengths_builder =
752            read_segment_table(&mut &buf[..], message::ReaderOptions::new())
753                .unwrap()
754                .unwrap();
755        assert_eq!(0, segment_lengths_builder.total_words());
756        assert_eq!(vec![(0, 0)], segment_lengths_builder.to_segment_indices());
757        buf.clear();
758
759        buf.extend(
760            [
761                0, 0, 0, 0, // 1 segments
762                1, 0, 0, 0,
763            ], // 1 length
764        );
765        let segment_lengths_builder =
766            read_segment_table(&mut &buf[..], message::ReaderOptions::new())
767                .unwrap()
768                .unwrap();
769        assert_eq!(1, segment_lengths_builder.total_words());
770        assert_eq!(vec![(0, 1)], segment_lengths_builder.to_segment_indices());
771        buf.clear();
772
773        buf.extend(
774            [
775                1, 0, 0, 0, // 2 segments
776                1, 0, 0, 0, // 1 length
777                1, 0, 0, 0, // 1 length
778                0, 0, 0, 0,
779            ], // padding
780        );
781        let segment_lengths_builder =
782            read_segment_table(&mut &buf[..], message::ReaderOptions::new())
783                .unwrap()
784                .unwrap();
785        assert_eq!(2, segment_lengths_builder.total_words());
786        assert_eq!(
787            vec![(0, 1), (1, 2)],
788            segment_lengths_builder.to_segment_indices()
789        );
790        buf.clear();
791
792        buf.extend(
793            [
794                2, 0, 0, 0, // 3 segments
795                1, 0, 0, 0, // 1 length
796                1, 0, 0, 0, // 1 length
797                0, 1, 0, 0,
798            ], // 256 length
799        );
800        let segment_lengths_builder =
801            read_segment_table(&mut &buf[..], message::ReaderOptions::new())
802                .unwrap()
803                .unwrap();
804        assert_eq!(258, segment_lengths_builder.total_words());
805        assert_eq!(
806            vec![(0, 1), (1, 2), (2, 258)],
807            segment_lengths_builder.to_segment_indices()
808        );
809        buf.clear();
810
811        buf.extend(
812            [
813                3, 0, 0, 0, // 4 segments
814                77, 0, 0, 0, // 77 length
815                23, 0, 0, 0, // 23 length
816                1, 0, 0, 0, // 1 length
817                99, 0, 0, 0, // 99 length
818                0, 0, 0, 0,
819            ], // padding
820        );
821        let segment_lengths_builder =
822            read_segment_table(&mut &buf[..], message::ReaderOptions::new())
823                .unwrap()
824                .unwrap();
825        assert_eq!(200, segment_lengths_builder.total_words());
826        assert_eq!(
827            vec![(0, 77), (77, 100), (100, 101), (101, 200)],
828            segment_lengths_builder.to_segment_indices()
829        );
830        buf.clear();
831    }
832
833    struct MaxRead<R>
834    where
835        R: Read,
836    {
837        inner: R,
838        max: usize,
839    }
840
841    impl<R> Read for MaxRead<R>
842    where
843        R: Read,
844    {
845        fn read(&mut self, buf: &mut [u8]) -> crate::Result<usize> {
846            if buf.len() <= self.max {
847                self.inner.read(buf)
848            } else {
849                self.inner.read(&mut buf[0..self.max])
850            }
851        }
852    }
853
854    #[test]
855    fn test_read_segment_table_max_read() {
856        // Make sure things still work well when we read less than a word at a time.
857        let mut buf: alloc::vec::Vec<u8> = vec![];
858        buf.extend(
859            [
860                0, 0, 0, 0, // 1 segments
861                1, 0, 0, 0,
862            ], // 1 length
863        );
864        let segment_lengths_builder = read_segment_table(
865            &mut MaxRead {
866                inner: &buf[..],
867                max: 2,
868            },
869            message::ReaderOptions::new(),
870        )
871        .unwrap()
872        .unwrap();
873        assert_eq!(1, segment_lengths_builder.total_words());
874        assert_eq!(vec![(0, 1)], segment_lengths_builder.to_segment_indices());
875    }
876
877    #[test]
878    fn test_try_read_message_no_alloc_max_read() {
879        // A message with multiple segments, so that reading the segment table
880        // requires more than one read() call.
881        let mut msg = message::Builder::new(message::HeapAllocator::new().first_segment_words(1));
882        msg.set_root("hello world!").unwrap();
883        assert!(msg.get_segments_for_output().len() > 1);
884
885        let mut bytes = alloc::vec::Vec::new();
886        super::write_message(&mut bytes, &msg).unwrap();
887
888        let mut buffer = [crate::word(0, 0, 0, 0, 0, 0, 0, 0); 64];
889        let reader = super::try_read_message_no_alloc(
890            MaxRead {
891                inner: &bytes[..],
892                max: 2,
893            },
894            crate::Word::words_to_bytes_mut(&mut buffer),
895            message::ReaderOptions::new(),
896        )
897        .unwrap()
898        .unwrap();
899        let text: crate::text::Reader = reader.get_root().unwrap();
900        assert_eq!("hello world!", text);
901    }
902
903    #[test]
904    fn test_read_invalid_segment_table() {
905        let mut buf = vec![];
906
907        buf.extend([0, 2, 0, 0]); // 513 segments
908        buf.extend([0; 513 * 4]);
909        assert!(read_segment_table(&mut &buf[..], message::ReaderOptions::new()).is_err());
910        buf.clear();
911
912        buf.extend([0, 0, 0, 0]); // 1 segments
913        assert!(read_segment_table(&mut &buf[..], message::ReaderOptions::new()).is_err());
914        buf.clear();
915
916        buf.extend([0, 0, 0, 0]); // 1 segments
917        buf.extend([0; 3]);
918        assert!(read_segment_table(&mut &buf[..], message::ReaderOptions::new()).is_err());
919        buf.clear();
920
921        buf.extend([255, 255, 255, 255]); // 0 segments
922        assert!(read_segment_table(&mut &buf[..], message::ReaderOptions::new()).is_err());
923        buf.clear();
924    }
925
926    #[test]
927    fn test_read_segment_table_overflow() {
928        let mut buf = vec![];
929
930        buf.extend([1, 0, 0, 0]); // 2 segments
931        buf.extend([0xff, 0xff, 0xff, 0xff]); // 2^32 - 1 words
932        buf.extend([2, 0, 0, 0]); // 2 words
933        buf.extend([0, 0, 0, 0]); // padding
934        assert!(read_segment_table(&mut &buf[..], message::ReaderOptions::new()).is_err());
935    }
936
937    #[test]
938    fn test_write_segment_table() {
939        let mut buf = vec![];
940
941        let segment_0 = [0u8; 0];
942        let segment_1 = [1u8, 1, 1, 1, 1, 1, 1, 1];
943        let segment_199 = [201u8; 199 * 8];
944
945        write_segment_table(&mut buf, &[&segment_0]).unwrap();
946        assert_eq!(
947            &[
948                0, 0, 0, 0, // 1 segments
949                0, 0, 0, 0
950            ], // 0 length
951            &buf[..]
952        );
953        buf.clear();
954
955        write_segment_table(&mut buf, &[&segment_1]).unwrap();
956        assert_eq!(
957            &[
958                0, 0, 0, 0, // 1 segments
959                1, 0, 0, 0
960            ], // 1 length
961            &buf[..]
962        );
963        buf.clear();
964
965        write_segment_table(&mut buf, &[&segment_199]).unwrap();
966        assert_eq!(
967            &[
968                0, 0, 0, 0, // 1 segments
969                199, 0, 0, 0
970            ], // 199 length
971            &buf[..]
972        );
973        buf.clear();
974
975        write_segment_table(&mut buf, &[&segment_0, &segment_1]).unwrap();
976        assert_eq!(
977            &[
978                1, 0, 0, 0, // 2 segments
979                0, 0, 0, 0, // 0 length
980                1, 0, 0, 0, // 1 length
981                0, 0, 0, 0
982            ], // padding
983            &buf[..]
984        );
985        buf.clear();
986
987        write_segment_table(
988            &mut buf,
989            &[&segment_199, &segment_1, &segment_199, &segment_0],
990        )
991        .unwrap();
992        assert_eq!(
993            &[
994                3, 0, 0, 0, // 4 segments
995                199, 0, 0, 0, // 199 length
996                1, 0, 0, 0, // 1 length
997                199, 0, 0, 0, // 199 length
998                0, 0, 0, 0, // 0 length
999                0, 0, 0, 0
1000            ], // padding
1001            &buf[..]
1002        );
1003        buf.clear();
1004
1005        write_segment_table(
1006            &mut buf,
1007            &[
1008                &segment_199,
1009                &segment_1,
1010                &segment_199,
1011                &segment_0,
1012                &segment_1,
1013            ],
1014        )
1015        .unwrap();
1016        assert_eq!(
1017            &[
1018                4, 0, 0, 0, // 5 segments
1019                199, 0, 0, 0, // 199 length
1020                1, 0, 0, 0, // 1 length
1021                199, 0, 0, 0, // 199 length
1022                0, 0, 0, 0, // 0 length
1023                1, 0, 0, 0
1024            ], // 1 length
1025            &buf[..]
1026        );
1027        buf.clear();
1028    }
1029
1030    quickcheck! {
1031        #[cfg_attr(miri, ignore)] // miri takes a long time with quickcheck
1032        fn test_round_trip(segments: alloc::vec::Vec<alloc::vec::Vec<crate::Word>>) -> TestResult {
1033            if segments.is_empty() { return TestResult::discard(); }
1034            let mut buf: alloc::vec::Vec<u8> = vec![];
1035
1036            write_message_segments(&mut buf, &segments);
1037            let message = read_message(&mut &buf[..], message::ReaderOptions::new()).unwrap();
1038            let result_segments = message.into_segments();
1039
1040            TestResult::from_bool(segments.iter().enumerate().all(|(i, segment)| {
1041                crate::Word::words_to_bytes(&segment[..]) == result_segments.get_segment(i as u32).unwrap()
1042            }))
1043        }
1044
1045        #[cfg_attr(miri, ignore)] // miri takes a long time with quickcheck
1046        fn test_round_trip_slice_segments(segments: alloc::vec::Vec<alloc::vec::Vec<crate::Word>>) -> TestResult {
1047            if segments.is_empty() { return TestResult::discard(); }
1048            let borrowed_segments: &[&[u8]] = &segments.iter()
1049                .map(|segment| crate::Word::words_to_bytes(&segment[..]))
1050                .collect::<alloc::vec::Vec<_>>()[..];
1051            let words = flatten_segments(borrowed_segments);
1052            let mut word_slice = &words[..];
1053            let message = read_message_from_flat_slice(&mut word_slice, message::ReaderOptions::new()).unwrap();
1054            assert!(word_slice.is_empty());  // no remaining words
1055            let result_segments = message.into_segments();
1056
1057            TestResult::from_bool(segments.iter().enumerate().all(|(i, segment)| {
1058                crate::Word::words_to_bytes(&segment[..]) == result_segments.get_segment(i as u32).unwrap()
1059            }))
1060        }
1061    }
1062
1063    #[test]
1064    fn read_message_from_flat_slice_with_remainder() {
1065        let segments = [
1066            vec![123, 0, 0, 0, 0, 0, 0, 0],
1067            vec![4, 0, 0, 0, 0, 0, 0, 0, 5, 0, 0, 0, 0, 0, 0, 0],
1068        ];
1069
1070        let borrowed_segments: &[&[u8]] = &segments
1071            .iter()
1072            .map(|segment| &segment[..])
1073            .collect::<alloc::vec::Vec<_>>()[..];
1074
1075        let mut bytes = flatten_segments(borrowed_segments);
1076        let extra_bytes: &[u8] = &[9, 9, 9, 9, 9, 9, 9, 9, 8, 7, 6, 5, 4, 3, 2, 1];
1077        for &b in extra_bytes {
1078            bytes.push(b);
1079        }
1080        let mut byte_slice = &bytes[..];
1081        let message =
1082            read_message_from_flat_slice(&mut byte_slice, message::ReaderOptions::new()).unwrap();
1083        assert_eq!(byte_slice, extra_bytes);
1084        let result_segments = message.into_segments();
1085        for (idx, segment) in segments.iter().enumerate() {
1086            assert_eq!(
1087                *segment,
1088                result_segments
1089                    .get_segment(idx as u32)
1090                    .expect("segment should exist")
1091            );
1092        }
1093    }
1094
1095    #[test]
1096    fn read_message_from_flat_slice_too_short() {
1097        let segments = [
1098            vec![1, 0, 0, 0, 0, 0, 0, 0],
1099            vec![2, 0, 0, 0, 0, 0, 0, 0, 3, 0, 0, 0, 0, 0, 0, 0],
1100        ];
1101
1102        let borrowed_segments: &[&[u8]] = &segments
1103            .iter()
1104            .map(|segment| &segment[..])
1105            .collect::<alloc::vec::Vec<_>>()[..];
1106
1107        let mut bytes = flatten_segments(borrowed_segments);
1108        while !bytes.is_empty() {
1109            bytes.pop();
1110            assert!(
1111                read_message_from_flat_slice(&mut &bytes[..], message::ReaderOptions::new())
1112                    .is_err()
1113            );
1114        }
1115    }
1116
1117    #[test]
1118    fn compute_serialized_size() {
1119        const LIST_LENGTH_IN_WORDS: u32 = 5;
1120        let mut m = message::Builder::new_default();
1121        {
1122            let root: crate::any_pointer::Builder = m.init_root();
1123            let _list_builder: crate::primitive_list::Builder<u64> =
1124                root.initn_as(LIST_LENGTH_IN_WORDS);
1125        }
1126
1127        // The message body has a list pointer (one word) and the list (LIST_LENGTH_IN_WORDS words).
1128        // The message has one segment, so the header is one word.
1129        assert_eq!(
1130            super::compute_serialized_size_in_words(&m) as u32,
1131            1 + 1 + LIST_LENGTH_IN_WORDS
1132        )
1133    }
1134}