1use std::io::{self, Read, Write};23use chacha20poly1305::{KeyInit as _, XChaCha20Poly1305, XNonce, aead::Aead as _};45pub const FRAME_SIZE: usize = 256 * 1024;6pub const TAG_SIZE: usize = 16;78fn nonce(counter: u64, last: bool) -> XNonce {9 let mut n = [0u8; 24];10 n[..8].copy_from_slice(&counter.to_le_bytes());11 n[23] = last as u8;12 n.into()13}1415pub struct EncryptWriter<W: Write> {16 inner: W,17 aead: XChaCha20Poly1305,18 buf: Vec<u8>,19 counter: u64,20}2122impl<W: Write> EncryptWriter<W> {23 pub fn new(key: &[u8; 32], inner: W) -> Self {24 Self {25 inner,26 aead: XChaCha20Poly1305::new(key.into()),27 buf: Vec::with_capacity(FRAME_SIZE),28 counter: 0,29 }30 }3132 fn emit(&mut self, last: bool) -> io::Result<()> {33 let ct = self34 .aead35 .encrypt(&nonce(self.counter, last), self.buf.as_slice())36 .map_err(|_| io::Error::other("aead encrypt failed"))?;37 self.inner.write_all(&ct)?;38 self.counter += 1;39 self.buf.clear();40 Ok(())41 }4243 pub fn finish(mut self) -> io::Result<W> {44 self.emit(true)?;45 self.inner.flush()?;46 Ok(self.inner)47 }48}4950impl<W: Write> Write for EncryptWriter<W> {51 fn write(&mut self, data: &[u8]) -> io::Result<usize> {52 let take = data.len().min(FRAME_SIZE - self.buf.len());53 self.buf.extend_from_slice(&data[..take]);54 if self.buf.len() == FRAME_SIZE {55 self.emit(false)?;56 }57 Ok(take)58 }59 fn flush(&mut self) -> io::Result<()> {60 Ok(())61 }62}6364pub struct DecryptReader<R: Read> {65 inner: R,66 aead: XChaCha20Poly1305,67 counter: u64,68 plain: Vec<u8>,69 plain_pos: usize,70 lookahead: Option<u8>,71 done: bool,72}7374impl<R: Read> DecryptReader<R> {75 pub fn new(key: &[u8; 32], inner: R) -> Self {76 Self {77 inner,78 aead: XChaCha20Poly1305::new(key.into()),79 counter: 0,80 plain: Vec::new(),81 plain_pos: 0,82 lookahead: None,83 done: false,84 }85 }8687 fn read_frame(&mut self) -> io::Result<()> {88 let mut frame = Vec::with_capacity(FRAME_SIZE + TAG_SIZE + 1);89 if let Some(b) = self.lookahead.take() {90 frame.push(b);91 }92 (&mut self.inner)93 .take((FRAME_SIZE + TAG_SIZE + 1 - frame.len()) as u64)94 .read_to_end(&mut frame)?;95 let last = if frame.len() == FRAME_SIZE + TAG_SIZE + 1 {96 self.lookahead = frame.pop();97 false98 } else {99 true100 };101 if frame.len() < TAG_SIZE {102 return Err(io::Error::other("truncated stream: incomplete frame"));103 }104 self.plain = self105 .aead106 .decrypt(&nonce(self.counter, last), frame.as_slice())107 .map_err(|_| io::Error::other("aead decrypt failed: corrupted or truncated stream"))?;108 self.counter += 1;109 self.plain_pos = 0;110 self.done = last;111 Ok(())112 }113}114115impl<R: Read> Read for DecryptReader<R> {116 fn read(&mut self, out: &mut [u8]) -> io::Result<usize> {117 while self.plain_pos == self.plain.len() {118 if self.done {119 return Ok(0);120 }121 self.read_frame()?;122 }123 let take = out.len().min(self.plain.len() - self.plain_pos);124 out[..take].copy_from_slice(&self.plain[self.plain_pos..self.plain_pos + take]);125 self.plain_pos += take;126 Ok(take)127 }128}129130#[cfg(test)]131mod tests {132 use super::*;133134 fn roundtrip(len: usize) {135 let key = [7u8; 32];136 let data = (0..len).map(|i| (i * 31 % 256) as u8).collect::<Vec<_>>();137 let mut enc = EncryptWriter::new(&key, Vec::new());138 enc.write_all(&data).unwrap();139 let ct = enc.finish().unwrap();140 let mut out = Vec::new();141 DecryptReader::new(&key, ct.as_slice())142 .read_to_end(&mut out)143 .unwrap();144 assert_eq!(out, data);145 }146147 #[test]148 fn roundtrips() {149 roundtrip(0);150 roundtrip(1);151 roundtrip(FRAME_SIZE - 1);152 roundtrip(FRAME_SIZE);153 roundtrip(FRAME_SIZE + 1);154 roundtrip(FRAME_SIZE * 3 + 17);155 }156157 #[test]158 fn truncation_detected() {159 let key = [7u8; 32];160 let mut enc = EncryptWriter::new(&key, Vec::new());161 enc.write_all(&vec![1u8; FRAME_SIZE * 2]).unwrap();162 let ct = enc.finish().unwrap();163 let mut out = Vec::new();164 let res = DecryptReader::new(&key, &ct[..ct.len() - TAG_SIZE - 1]).read_to_end(&mut out);165 assert!(res.is_err());166 }167168 #[test]169 fn deterministic() {170 let key = [7u8; 32];171 let data = vec![3u8; FRAME_SIZE + 100];172 let ct = |d: &[u8]| {173 let mut enc = EncryptWriter::new(&key, Vec::new());174 enc.write_all(d).unwrap();175 enc.finish().unwrap()176 };177 assert_eq!(ct(&data), ct(&data));178 }179}