Skip to main content

mtorrent_core/pe/
io.rs

1use super::{Decryptor, Encryptor};
2use bytes::{Buf, BufMut, BytesMut};
3use mtorrent_utils::split_stream::SplitStream;
4use pin_project_lite::pin_project;
5use std::pin::Pin;
6use std::task::{Context, Poll, ready};
7use std::{cmp, io};
8use tokio::io::{AsyncBufRead, AsyncRead, AsyncReadExt, AsyncWrite, Chain, ReadBuf};
9
10pin_project! {
11    /// A wrapper around an [`AsyncRead`](https://docs.rs/tokio/latest/tokio/io/trait.AsyncRead.html) that decrypts the data read from it using a [`Decryptor`].
12    #[derive(Debug)]
13    pub struct DecryptingReader<R: AsyncRead> {
14        #[pin]
15        inner: R,
16        crypto: Decryptor,
17    }
18}
19
20impl<R: AsyncRead> DecryptingReader<R> {
21    /// Creates a new `DecryptingReader` that wraps the given [`AsyncRead`](https://docs.rs/tokio/latest/tokio/io/trait.AsyncRead.html) and [`Decryptor`].
22    pub fn new(inner: R, crypto: Decryptor) -> Self {
23        Self { inner, crypto }
24    }
25
26    /// Consumes the `DecryptingReader`, returning the wrapped [`AsyncRead`](https://docs.rs/tokio/latest/tokio/io/trait.AsyncRead.html) and [`Decryptor`].
27    pub fn into_parts(self) -> (R, Decryptor) {
28        (self.inner, self.crypto)
29    }
30}
31
32impl<R: AsyncRead> AsyncRead for DecryptingReader<R> {
33    fn poll_read(
34        self: Pin<&mut Self>,
35        cx: &mut Context<'_>,
36        buf: &mut ReadBuf<'_>,
37    ) -> Poll<io::Result<()>> {
38        let this = self.project();
39
40        let old_len = buf.filled().len();
41        ready!(this.inner.poll_read(cx, buf))?;
42        let new_len = buf.filled().len();
43
44        if new_len > old_len {
45            this.crypto.decrypt(&mut buf.filled_mut()[old_len..new_len]);
46        }
47        Poll::Ready(Ok(()))
48    }
49}
50
51const BUFFER_SIZE: usize = 33 * 1024;
52
53pin_project! {
54    /// A buffering reader that wraps an [`AsyncRead`](https://docs.rs/tokio/latest/tokio/io/trait.AsyncRead.html) and decrypts the data read from it using a [`Decryptor`].
55    ///
56    /// Implementation of `DecryptingBufReader` is similar to [`tokio::io::BufReader`](https://docs.rs/tokio/latest/tokio/io/struct.BufReader.html), slightly
57    /// simplified thanks to [`BytesMut`](https://docs.rs/bytes/latest/bytes/struct.BytesMut.html), and with added calls to `Decryptor::decrypt()`.
58    #[derive(Debug)]
59    pub struct DecryptingBufReader<R: AsyncRead> {
60        #[pin]
61        inner: R,
62        crypto: Decryptor,
63        buffer: BytesMut,
64    }
65}
66
67impl<R: AsyncRead> DecryptingBufReader<R> {
68    /// Creates a new `DecryptingBufReader` that wraps the given [`AsyncRead`](https://docs.rs/tokio/latest/tokio/io/trait.AsyncRead.html) and [`Decryptor`].
69    pub fn new(inner: R, crypto: Decryptor) -> Self {
70        Self {
71            inner,
72            crypto,
73            buffer: BytesMut::with_capacity(BUFFER_SIZE),
74        }
75    }
76
77    /// Consumes the `DecryptingBufReader`, returning the wrapped [`AsyncRead`](https://docs.rs/tokio/latest/tokio/io/trait.AsyncRead.html) and [`Decryptor`].
78    pub fn into_parts(self) -> (R, Decryptor) {
79        (self.inner, self.crypto)
80    }
81}
82
83impl<R: AsyncRead> AsyncRead for DecryptingBufReader<R> {
84    fn poll_read(
85        mut self: Pin<&mut Self>,
86        cx: &mut Context<'_>,
87        buf: &mut ReadBuf<'_>,
88    ) -> Poll<io::Result<()>> {
89        let data = ready!(self.as_mut().poll_fill_buf(cx))?;
90        let bytes_to_copy = cmp::min(data.len(), buf.remaining());
91        buf.put_slice(&data[..bytes_to_copy]);
92        self.consume(bytes_to_copy);
93        Poll::Ready(Ok(()))
94    }
95}
96
97impl<R: AsyncRead> AsyncBufRead for DecryptingBufReader<R> {
98    fn poll_fill_buf(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<&[u8]>> {
99        let this = self.project();
100
101        if this.buffer.is_empty() {
102            assert!(this.buffer.try_reclaim(BUFFER_SIZE));
103            let mut rd = ReadBuf::uninit(this.buffer.spare_capacity_mut());
104            ready!(this.inner.poll_read(cx, &mut rd))?;
105
106            this.crypto.decrypt(rd.filled_mut());
107
108            let bytes_read = rd.filled().len();
109            unsafe { this.buffer.advance_mut(bytes_read) }
110        }
111        Poll::Ready(Ok(this.buffer))
112    }
113
114    fn consume(self: Pin<&mut Self>, amt: usize) {
115        let this = self.project();
116        let amt = cmp::min(amt, this.buffer.remaining());
117        this.buffer.advance(amt);
118    }
119}
120
121pin_project! {
122    /// A buffering writer that wraps an [`AsyncWrite`](https://docs.rs/tokio/latest/tokio/io/trait.AsyncWrite.html) and encrypts the data written to it using an [`Encryptor`].
123    ///
124    /// Implementation of `EncryptingWriter` is similar to [`tokio::io::BufWriter`](https://docs.rs/tokio/latest/tokio/io/struct.BufWriter.html), slightly
125    /// simplified thanks to [`BytesMut`](https://docs.rs/bytes/latest/bytes/struct.BytesMut.html), and with added calls to `Encryptor::encrypt()`.
126    #[derive(Debug)]
127    pub struct EncryptingWriter<W: AsyncWrite> {
128        #[pin]
129        inner: W,
130        crypto: Encryptor,
131        buffer: BytesMut,
132    }
133}
134
135impl<W: AsyncWrite> EncryptingWriter<W> {
136    /// Creates a new `EncryptingWriter` that wraps the given [`AsyncWrite`](https://docs.rs/tokio/latest/tokio/io/trait.AsyncWrite.html) and [`Encryptor`].
137    pub fn new(inner: W, crypto: Encryptor) -> Self {
138        Self {
139            inner,
140            crypto,
141            buffer: BytesMut::with_capacity(BUFFER_SIZE),
142        }
143    }
144
145    /// Consumes the `EncryptingWriter`, returning the wrapped [`AsyncWrite`](https://docs.rs/tokio/latest/tokio/io/trait.AsyncWrite.html) and [`Encryptor`].
146    pub fn into_parts(self) -> (W, Encryptor) {
147        (self.inner, self.crypto)
148    }
149
150    fn flush_buf(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
151        let mut this = self.project();
152
153        while !this.buffer.is_empty() {
154            let written = ready!(this.inner.as_mut().poll_write(cx, this.buffer))?;
155            if written == 0 {
156                return Poll::Ready(Err(io::Error::new(
157                    io::ErrorKind::WriteZero,
158                    "failed to write the buffered data",
159                )));
160            }
161            this.buffer.advance(written);
162        }
163        Poll::Ready(Ok(()))
164    }
165
166    fn fill_buf(self: Pin<&mut Self>, data: &[u8]) -> io::Result<usize> {
167        let this = self.project();
168        assert!(this.buffer.capacity() <= BUFFER_SIZE);
169
170        let filled_len = this.buffer.len();
171        let available_cap = {
172            let curr_available = this.buffer.capacity() - filled_len;
173            if data.len() > curr_available {
174                let max_available = BUFFER_SIZE - filled_len;
175                assert!(this.buffer.try_reclaim(max_available));
176                max_available
177            } else {
178                curr_available
179            }
180        };
181
182        let bytes_to_write = cmp::min(data.len(), available_cap);
183
184        if bytes_to_write == 0 {
185            return Err(io::Error::new(io::ErrorKind::WriteZero, "can't write data to buffer"));
186        }
187
188        this.buffer.extend_from_slice(&data[..bytes_to_write]);
189        this.crypto.encrypt(&mut this.buffer[filled_len..]);
190
191        Ok(bytes_to_write)
192    }
193}
194
195impl<W: AsyncWrite> AsyncWrite for EncryptingWriter<W> {
196    fn poll_write(
197        mut self: Pin<&mut Self>,
198        cx: &mut Context<'_>,
199        buf: &[u8],
200    ) -> Poll<io::Result<usize>> {
201        ready!(self.as_mut().flush_buf(cx))?;
202        Poll::Ready(self.fill_buf(buf))
203    }
204
205    fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
206        ready!(self.as_mut().flush_buf(cx))?;
207        self.project().inner.poll_flush(cx)
208    }
209
210    fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
211        ready!(self.as_mut().flush_buf(cx))?;
212        self.project().inner.poll_shutdown(cx)
213    }
214}
215
216pin_project! {
217    /// A wrapper around an [`AsyncRead`](https://docs.rs/tokio/latest/tokio/io/trait.AsyncRead.html) that first returns the bytes from an initial buffer `T`, and then
218    /// delegates to the wrapped stream `S`. Similar to [`Chain`](https://docs.rs/tokio/latest/tokio/io/struct.Chain.html), but also implements `AsyncWrite` and `SplitStream`
219    /// if `S` does.
220    pub struct PrefixedStream<T: Buf, S> {
221        prefix: T,
222        #[pin]
223        stream: S,
224    }
225}
226
227impl<T: Buf, S> PrefixedStream<T, S> {
228    /// Creates a new `PrefixedStream` that wraps the given initial buffer and stream.
229    pub fn new(prefix: T, stream: S) -> Self {
230        Self { prefix, stream }
231    }
232
233    /// Consumes the `PrefixedStream`, returning the prefix buffer and the wrapped stream.
234    pub fn into_parts(self) -> (T, S) {
235        (self.prefix, self.stream)
236    }
237}
238
239impl<T: Buf, S: AsyncRead> AsyncRead for PrefixedStream<T, S> {
240    fn poll_read(
241        self: Pin<&mut Self>,
242        cx: &mut Context<'_>,
243        buf: &mut ReadBuf<'_>,
244    ) -> Poll<io::Result<()>> {
245        if self.prefix.has_remaining() {
246            let bytes_to_copy = cmp::min(buf.remaining(), self.prefix.chunk().len());
247            buf.put_slice(&self.prefix.chunk()[..bytes_to_copy]);
248            self.project().prefix.advance(bytes_to_copy);
249            Poll::Ready(Ok(()))
250        } else {
251            self.project().stream.poll_read(cx, buf)
252        }
253    }
254}
255
256impl<T: Buf, S: AsyncWrite> AsyncWrite for PrefixedStream<T, S> {
257    fn poll_write(
258        self: Pin<&mut Self>,
259        cx: &mut Context<'_>,
260        buf: &[u8],
261    ) -> Poll<io::Result<usize>> {
262        self.project().stream.poll_write(cx, buf)
263    }
264
265    fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
266        self.project().stream.poll_flush(cx)
267    }
268
269    fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
270        self.project().stream.poll_shutdown(cx)
271    }
272
273    fn poll_write_vectored(
274        self: Pin<&mut Self>,
275        cx: &mut Context<'_>,
276        bufs: &[io::IoSlice<'_>],
277    ) -> Poll<io::Result<usize>> {
278        self.project().stream.poll_write_vectored(cx, bufs)
279    }
280
281    fn is_write_vectored(&self) -> bool {
282        self.stream.is_write_vectored()
283    }
284}
285
286impl<T: Buf, S: SplitStream> SplitStream for PrefixedStream<T, S> {
287    type Ingress<'i>
288        = Chain<&'i [u8], <S as SplitStream>::Ingress<'i>>
289    where
290        Self: 'i;
291
292    type Egress<'e>
293        = S::Egress<'e>
294    where
295        Self: 'e;
296
297    fn split(&mut self) -> (Self::Ingress<'_>, Self::Egress<'_>) {
298        let (ingress, egress) = self.stream.split();
299        let ingress = AsyncReadExt::chain(self.prefix.chunk(), ingress);
300        (ingress, egress)
301    }
302}
303
304#[cfg(test)]
305mod tests {
306    use super::super::cipher::crypto_pair;
307    use super::*;
308    use local_async_utils::prelude::*;
309    use tokio::io::{AsyncReadExt, AsyncWriteExt};
310    use tokio::join;
311
312    #[tokio::test]
313    async fn test_pipe_big_data_over_encrypted_stream() {
314        let (enc, dec) = crypto_pair(&rand::random());
315
316        let (reader, writer) = local_pipe::Pipe::new(1024).into_split();
317        let mut reader = DecryptingReader::new(reader, dec);
318        let mut writer = EncryptingWriter::new(writer, enc);
319
320        for _ in 0..3 {
321            let src_data: [u8; BUFFER_SIZE * 2] = rand::random();
322            let mut dest_data = [0u8; BUFFER_SIZE * 2];
323
324            let write_fut = async {
325                writer.write_all(&src_data).await.unwrap();
326                writer.flush().await.unwrap();
327            };
328            let read_fut = async {
329                reader.read_exact(&mut dest_data).await.unwrap();
330            };
331            join!(write_fut, read_fut);
332
333            assert_eq!(src_data, dest_data);
334        }
335    }
336
337    #[tokio::test]
338    async fn test_pipe_small_data_over_encrypted_stream() {
339        let (enc, dec) = crypto_pair(&rand::random());
340
341        let (reader, writer) = local_pipe::Pipe::new(1024).into_split();
342        let mut reader = DecryptingReader::new(reader, dec);
343        let mut writer = EncryptingWriter::new(writer, enc);
344
345        for _ in 0..3 {
346            let src_data: [u8; BUFFER_SIZE / 2] = rand::random();
347            let mut dest_data = [0u8; BUFFER_SIZE / 2];
348
349            let write_fut = async {
350                writer.write_all(&src_data).await.unwrap();
351                writer.flush().await.unwrap();
352            };
353            let read_fut = async {
354                reader.read_exact(&mut dest_data).await.unwrap();
355            };
356            join!(write_fut, read_fut);
357
358            assert_eq!(src_data, dest_data);
359        }
360    }
361
362    #[tokio::test]
363    async fn test_pipe_big_data_over_encrypted_buffered_stream() {
364        let (enc, dec) = crypto_pair(&rand::random());
365
366        let (reader, writer) = local_pipe::Pipe::new(1024).into_split();
367        let mut reader = DecryptingBufReader::new(reader, dec);
368        let mut writer = EncryptingWriter::new(writer, enc);
369
370        for _ in 0..3 {
371            let src_data: [u8; BUFFER_SIZE * 2] = rand::random();
372            let mut dest_data = [0u8; BUFFER_SIZE * 2];
373
374            let write_fut = async {
375                writer.write_all(&src_data).await.unwrap();
376                writer.flush().await.unwrap();
377            };
378            let read_fut = async {
379                reader.read_exact(&mut dest_data).await.unwrap();
380            };
381            join!(write_fut, read_fut);
382
383            assert_eq!(src_data, dest_data);
384        }
385    }
386
387    #[tokio::test]
388    async fn test_pipe_small_data_over_encrypted_buffered_stream() {
389        let (enc, dec) = crypto_pair(&rand::random());
390
391        let (reader, writer) = local_pipe::Pipe::new(1024).into_split();
392        let mut reader = DecryptingBufReader::new(reader, dec);
393        let mut writer = EncryptingWriter::new(writer, enc);
394
395        for _ in 0..3 {
396            let src_data: [u8; BUFFER_SIZE / 2] = rand::random();
397            let mut dest_data = [0u8; BUFFER_SIZE / 2];
398
399            let write_fut = async {
400                writer.write_all(&src_data).await.unwrap();
401                writer.flush().await.unwrap();
402            };
403            let read_fut = async {
404                reader.read_exact(&mut dest_data).await.unwrap();
405            };
406            join!(write_fut, read_fut);
407
408            assert_eq!(src_data, dest_data);
409        }
410    }
411
412    #[tokio::test]
413    async fn test_prefixed_stream() {
414        let prefix: [u8; 10] = rand::random();
415        let data: [u8; 20] = rand::random();
416        let mut stream = PrefixedStream::new(BytesMut::from(&prefix[..]), &data[..]);
417
418        let mut buf = Vec::new();
419        stream.read_to_end(&mut buf).await.unwrap();
420
421        assert_eq!(&buf[..10], &prefix);
422        assert_eq!(&buf[10..], &data);
423    }
424}