use core::{
fmt::{self, Write},
ops::{Index, IndexMut, Range, RangeBounds},
};
#[cfg(feature = "serde")]
use sequential_storage::map::PostcardValue;
#[cfg(feature = "serde")]
use serde::{
Deserialize, Deserializer, Serialize, Serializer,
de::{SeqAccess, Visitor},
};
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct FixedBuf<const N: usize> {
pub buf: [u8; N],
pub length: usize,
}
impl<const N: usize> FixedBuf<N> {
pub const fn new() -> Self {
Self { buf: [0u8; N], length: 0 }
}
}
impl<const N: usize> Default for FixedBuf<N> {
fn default() -> Self {
Self::new()
}
}
#[allow(unused)]
impl<const N: usize> FixedBuf<N> {
pub fn clear(&mut self) {
self.length = 0;
}
pub fn as_bytes(&self) -> &[u8] {
&self.buf[..self.length]
}
pub fn as_bytes_mut(&mut self) -> &mut [u8] {
&mut self.buf[..self.length]
}
pub fn as_str(&self) -> Result<&str, core::str::Utf8Error> {
core::str::from_utf8(self.as_bytes())
}
pub fn get(&self, index: usize) -> Option<&u8> {
if index < self.length { Some(&self.buf[index]) } else { None }
}
pub fn get_mut(&mut self, index: usize) -> Option<&mut u8> {
if index < self.length { Some(&mut self.buf[index]) } else { None }
}
pub fn try_fill_range<R>(&mut self, range: R, value: u8) -> Result<(), ()>
where
R: RangeBounds<usize>,
{
let start = match range.start_bound() {
core::ops::Bound::Included(&s) => s,
core::ops::Bound::Excluded(&s) => s + 1,
core::ops::Bound::Unbounded => 0,
};
let end = match range.end_bound() {
core::ops::Bound::Included(&e) => e + 1,
core::ops::Bound::Excluded(&e) => e,
core::ops::Bound::Unbounded => self.length,
};
if start > end || end > self.length {
return Err(());
}
self.buf[start..end].fill(value);
Ok(())
}
}
impl<const N: usize> AsRef<[u8]> for FixedBuf<N> {
fn as_ref(&self) -> &[u8] {
self.as_bytes()
}
}
impl<const N: usize> AsMut<[u8]> for FixedBuf<N> {
fn as_mut(&mut self) -> &mut [u8] {
self.as_bytes_mut()
}
}
impl<const N: usize> Index<usize> for FixedBuf<N> {
type Output = u8;
fn index(&self, index: usize) -> &Self::Output {
assert!(index < self.length, "FixedBuf index out of bounds");
&self.buf[index]
}
}
impl<const N: usize> IndexMut<usize> for FixedBuf<N> {
fn index_mut(&mut self, index: usize) -> &mut Self::Output {
assert!(index < self.length, "FixedBuf index out of bounds");
&mut self.buf[index]
}
}
impl<const N: usize> Index<Range<usize>> for FixedBuf<N> {
type Output = [u8];
fn index(&self, range: Range<usize>) -> &Self::Output {
assert!(range.end <= self.length, "FixedBuf slice index out of bounds");
&self.buf[range]
}
}
impl<const N: usize> IndexMut<Range<usize>> for FixedBuf<N> {
fn index_mut(&mut self, range: Range<usize>) -> &mut Self::Output {
assert!(range.end <= self.length, "FixedBuf slice index out of bounds");
&mut self.buf[range]
}
}
impl<const N: usize> Write for FixedBuf<N> {
fn write_str(&mut self, s: &str) -> fmt::Result {
let bytes = s.as_bytes();
if self.length + bytes.len() > N {
return Err(fmt::Error);
}
self.buf[self.length..self.length + bytes.len()].copy_from_slice(bytes);
self.length += bytes.len();
Ok(())
}
}
#[cfg(feature = "serde")]
impl<const N: usize> Serialize for FixedBuf<N> {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
serializer.serialize_bytes(self.as_bytes())
}
}
#[cfg(feature = "serde")]
impl<'de, const N: usize> Deserialize<'de> for FixedBuf<N> {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
struct FixedBufVisitor<const M: usize>;
impl<'de, const M: usize> Visitor<'de> for FixedBufVisitor<M> {
type Value = FixedBuf<M>;
fn expecting(&self, formatter: &mut core::fmt::Formatter) -> core::fmt::Result {
formatter.write_str("a byte array or byte sequence")
}
fn visit_bytes<E>(self, v: &[u8]) -> Result<Self::Value, E>
where
E: serde::de::Error,
{
if v.len() > M {
return Err(E::custom("buffer overflow: input data exceeds capacity"));
}
let mut fixed_buf = FixedBuf::new();
fixed_buf.buf[..v.len()].copy_from_slice(v);
fixed_buf.length = v.len();
Ok(fixed_buf)
}
fn visit_seq<A>(self, mut seq: A) -> Result<Self::Value, A::Error>
where
A: SeqAccess<'de>,
{
let mut fixed_buf = FixedBuf::new();
let mut idx = 0;
while let Some(byte) = seq.next_element()? {
if idx >= M {
return Err(serde::de::Error::custom("buffer overflow"));
}
fixed_buf.buf[idx] = byte;
idx += 1;
}
fixed_buf.length = idx;
Ok(fixed_buf)
}
}
deserializer.deserialize_bytes(FixedBufVisitor::<N>)
}
}
#[cfg(feature = "serde")]
impl<const N: usize> PostcardValue<'_> for FixedBuf<N> {}
#[cfg(test)]
mod tests {
use super::*;
fn _is_normal<T: Sized + Send + Sync + Unpin>() {}
fn is_full<T: Sized + Send + Sync + Unpin + Copy + Clone + Default + PartialEq>() {}
#[cfg(feature = "serde")]
fn is_config<T: Serialize + for<'a> Deserialize<'a> + for<'a> PostcardValue<'a>>() {}
#[test]
fn normal_types() {
is_full::<FixedBuf<16>>();
#[cfg(feature = "serde")]
is_config::<FixedBuf<16>>();
}
}