use crate::MemoryError;
use crate::cell::Cell;
use crate::mem_safe::{MemSafe, MemSafeRead, MemSafeWrite};
pub struct Secret<const N: usize> {
inner: MemSafe<[u8; N]>,
}
impl<const N: usize> Secret<N> {
pub fn new_with<F>(init: F) -> Result<Self, MemoryError>
where
F: FnOnce(&mut [u8; N]),
{
Cell::<[u8; N]>::new_with(init).map(|cell| Secret {
inner: MemSafe { cell },
})
}
pub fn from_bytes<T: AsMut<[u8]>>(bytes: T) -> Result<Self, (T, MemoryError)> {
Cell::<[u8; N]>::from_bytes(bytes).map(|cell| Secret {
inner: MemSafe { cell },
})
}
pub fn read(&mut self) -> Result<MemSafeRead<'_, [u8; N]>, MemoryError> {
self.inner.read()
}
pub fn write(&mut self) -> Result<MemSafeWrite<'_, [u8; N]>, MemoryError> {
self.inner.write()
}
}
impl<const N: usize> TryFrom<&str> for Secret<N> {
type Error = MemoryError;
fn try_from(s: &str) -> Result<Self, Self::Error> {
if s.len() > N {
return Err(MemoryError::from(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"string length exceeds buffer size",
)));
}
let bytes = s.as_bytes();
let len = bytes.len();
Self::new_with(|page| unsafe {
std::ptr::copy_nonoverlapping(bytes.as_ptr(), page.as_mut_ptr(), len);
})
}
}
impl<const N: usize> TryFrom<String> for Secret<N> {
type Error = (String, MemoryError);
fn try_from(s: String) -> Result<Self, Self::Error> {
Self::from_bytes(s.into_bytes()).map_err(|(v, e)| {
(
String::from_utf8(v).expect("zeroed or original UTF-8 bytes"),
e,
)
})
}
}