Skip to main content

pb_mapper_protocol/
buffer.rs

1//! Define the buffer interface for reading data in different situations
2
3use snafu::ResultExt;
4use tokio::io::AsyncReadExt;
5
6use pb_mapper_core::error::MsgNetworkReadBufferdRawDataSnafu;
7
8const INIT_BUF_SIZE: usize = 8 * 1024;
9const MAX_BUF_SIZE: usize = 8 * 1024 * 1024;
10/// Buffer for situations where the length of the data to be read is not known
11pub trait DynamicSizeBuffer {
12    fn need_resize(&self) -> bool;
13
14    /// Dynamic capacity adjustment based on `need_size`
15    fn dyn_resize(&mut self);
16
17    /// If the `buffer` is filled, it is expanded; if the `buffer` is not filled and the filled
18    /// content is less than `INIT_BUF_SIZE` bytes and the `need_size` is greater than
19    /// `INIT_BUF_SIZE`, it is shrunk.
20    ///
21    /// # Arguments `n` : how many bytes were read into the
22    /// buffer this time
23    fn update_need_size(&mut self, n: usize);
24}
25
26/// Buffer for situations where the length of the data to be read is already known
27pub trait FixedSizeBuffer {
28    /// The length of the buffer must reach `size` after resize, as opposed to
29    /// [`DynamicSizeBuffer`]'s internal state-based resize.
30    fn fixed_resize(&mut self, size: usize);
31}
32
33pub trait BufferGetter {
34    fn buffer(&self) -> &'_ [u8];
35
36    fn buffer_mut(&mut self) -> &'_ mut [u8];
37}
38
39pub struct CommonBuffer {
40    buffer: Vec<u8>,
41    need_size: usize,
42}
43
44impl Default for CommonBuffer {
45    fn default() -> Self {
46        CommonBuffer::new()
47    }
48}
49
50impl CommonBuffer {
51    pub fn new() -> Self {
52        Self {
53            buffer: vec![0; INIT_BUF_SIZE],
54            need_size: INIT_BUF_SIZE,
55        }
56    }
57}
58
59impl DynamicSizeBuffer for CommonBuffer {
60    #[inline]
61    fn need_resize(&self) -> bool {
62        self.buffer.len() != self.need_size
63    }
64
65    #[inline]
66    fn dyn_resize(&mut self) {
67        if self.need_size >= MAX_BUF_SIZE {
68            self.need_size = MAX_BUF_SIZE;
69        }
70        self.buffer.resize(self.need_size, 0);
71    }
72
73    #[inline]
74    fn update_need_size(&mut self, n: usize) {
75        // update `need_size` for expand
76        if n == self.buffer.len() {
77            self.need_size = n * 2;
78        }
79        // update `need_size` for shrink
80        else if n != 0 && n < INIT_BUF_SIZE && self.need_size > INIT_BUF_SIZE {
81            self.need_size = INIT_BUF_SIZE;
82        }
83    }
84}
85
86impl FixedSizeBuffer for CommonBuffer {
87    #[inline]
88    fn fixed_resize(&mut self, size: usize) {
89        self.buffer.resize(size, 0)
90    }
91}
92
93impl BufferGetter for CommonBuffer {
94    #[inline]
95    fn buffer(&self) -> &'_ [u8] {
96        &self.buffer
97    }
98
99    #[inline]
100    fn buffer_mut(&mut self) -> &'_ mut [u8] {
101        &mut self.buffer
102    }
103}
104
105/// This trait is used for buffered reads where the packet length is not known
106pub trait BufferedReader {
107    async fn read(&mut self) -> pb_mapper_core::error::Result<&'_ [u8]>;
108}
109
110pub struct BufferReader<'a, T> {
111    reader: &'a mut T,
112    buffer: CommonBuffer,
113}
114impl<'reader, T: AsyncReadExt + Unpin> BufferReader<'reader, T> {
115    pub fn new(reader: &'reader mut T) -> Self {
116        Self {
117            reader,
118            buffer: CommonBuffer::new(),
119        }
120    }
121
122    async fn read_inner(&mut self) -> pb_mapper_core::error::Result<&[u8]> {
123        if self.buffer.need_resize() {
124            self.buffer.dyn_resize()
125        }
126        let n = self
127            .reader
128            .read(self.buffer.buffer_mut())
129            .await
130            .context(MsgNetworkReadBufferdRawDataSnafu)?;
131        self.buffer.update_need_size(n);
132        Ok(&self.buffer.buffer()[0..n])
133    }
134}
135
136impl<'reader, T: AsyncReadExt + Unpin> BufferedReader for BufferReader<'reader, T> {
137    async fn read(&mut self) -> pb_mapper_core::error::Result<&'_ [u8]> {
138        self.read_inner().await
139    }
140}