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 #[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 pub fn new(inner: R, crypto: Decryptor) -> Self {
23 Self { inner, crypto }
24 }
25
26 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 #[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 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 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 #[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 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 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 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 pub fn new(prefix: T, stream: S) -> Self {
230 Self { prefix, stream }
231 }
232
233 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}