use std::io::{self, Write};
use crate::SecureBytes;
pub struct SecureBytesWriter<'a> {
bytes: &'a mut SecureBytes,
}
impl<'a> SecureBytesWriter<'a> {
pub fn new(bytes: &'a mut SecureBytes) -> Self {
Self { bytes }
}
}
impl Write for SecureBytesWriter<'_> {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
self.bytes.extend_from_slice(buf);
Ok(buf.len())
}
fn write_all(&mut self, buf: &[u8]) -> io::Result<()> {
self.bytes.extend_from_slice(buf);
Ok(())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_writer_appends_to_secure_bytes() {
let mut buffer = SecureBytes::new_with_capacity(8).unwrap();
{
let mut writer = SecureBytesWriter::new(&mut buffer);
writer.write_all(b"hello ").unwrap();
writer.write_all(b"world").unwrap();
assert_eq!(writer.write(b"!").unwrap(), 1);
writer.flush().unwrap();
}
buffer.unlock_slice(|bytes| assert_eq!(bytes, b"hello world!"));
}
#[test]
fn test_writer_grows_without_losing_data() {
let mut buffer = SecureBytes::new().unwrap();
{
let mut writer = SecureBytesWriter::new(&mut buffer);
for _ in 0..64 {
writer.write_all(&[0xAB; 32]).unwrap();
}
}
buffer.unlock_slice(|bytes| {
assert_eq!(bytes.len(), 64 * 32);
assert!(bytes.iter().all(|byte| *byte == 0xAB));
});
}
}