git.delta.rocks / fleet / refs/commits / 615754ca0747

difftreelog

source

crates/fleet-usb/src/stream.rs4.1 KiBsourcehistory
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}