use alloc::string::{String, ToString};
use alloc::vec::Vec;
use core::fmt;
use crate::field::{BitOrder, Bits, ByteOrder};
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct BitError {
pub kind: ErrorKind,
pub at: usize,
pub field: Option<&'static str>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub enum ErrorKind {
UnexpectedEof {
needed: usize,
remaining: usize,
},
Incomplete {
needed: Option<usize>,
},
TrailingBytes {
remaining: usize,
},
TooWide {
width: usize,
},
#[cfg(feature = "std")]
Io(std::io::ErrorKind),
BadMagic {
expected: u128,
found: u128,
},
Convert {
message: String,
},
NotSeekable,
BufferFull {
cap: usize,
},
}
impl BitError {
#[must_use]
pub fn new(kind: ErrorKind, at: usize) -> Self {
Self {
kind,
at,
field: None,
}
}
#[must_use]
pub fn bad_magic(expected: u128, found: u128, at: usize) -> Self {
Self::new(ErrorKind::BadMagic { expected, found }, at)
}
#[must_use]
pub fn convert(message: String, at: usize) -> Self {
Self::new(ErrorKind::Convert { message }, at)
}
#[must_use]
pub fn in_field(mut self, field: &'static str) -> Self {
if self.field.is_none() {
self.field = Some(field);
}
self
}
#[must_use]
pub fn is_incomplete(&self) -> bool {
matches!(self.kind, ErrorKind::Incomplete { .. })
}
}
#[cfg(feature = "std")]
impl From<std::io::Error> for BitError {
fn from(e: std::io::Error) -> Self {
BitError::new(ErrorKind::Io(e.kind()), 0)
}
}
impl fmt::Display for BitError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match &self.kind {
ErrorKind::UnexpectedEof { needed, remaining } => write!(
f,
"unexpected end of input: needed {needed} bits, {remaining} remain"
)?,
ErrorKind::Incomplete { needed } => match needed {
Some(n) => write!(f, "incomplete: need ~{n} more bytes")?,
None => write!(f, "incomplete: need more bytes")?,
},
ErrorKind::TrailingBytes { remaining } => {
write!(f, "{remaining} trailing bytes after the message")?;
}
ErrorKind::TooWide { width } => {
write!(f, "field width {width} exceeds the 128-bit carrier")?;
}
#[cfg(feature = "std")]
ErrorKind::Io(kind) => write!(f, "I/O error: {kind:?}")?,
ErrorKind::BadMagic { expected, found } => {
write!(f, "bad magic: expected {expected:#x}, found {found:#x}")?;
}
ErrorKind::Convert { message } => {
write!(f, "conversion failed: {message}")?;
}
ErrorKind::NotSeekable => {
write!(f, "a position directive ran on a non-seekable source")?;
}
ErrorKind::BufferFull { cap } => {
write!(f, "buffered source exceeded its {cap}-byte cap")?;
}
}
write!(f, " at bit {}", self.at)?;
if let Some(field) = self.field {
write!(f, " (field `{field}`)")?;
}
Ok(())
}
}
impl core::error::Error for BitError {}
impl From<crate::error::Error> for BitError {
#[inline]
fn from(e: crate::error::Error) -> Self {
BitError::convert(e.to_string(), 0)
}
}
#[doc(hidden)]
#[must_use]
pub const fn bits_of<T: Bits>(_value: &T) -> u32 {
T::BITS
}
#[doc(hidden)]
pub fn verify_magic<T: Bits, S: Source>(r: &mut S, expected: T) -> Result<(), BitError> {
let at = r.bit_pos();
let found: T = r.read()?;
let (e, g) = (expected.into_bits(), found.into_bits());
if e != g {
return Err(BitError::bad_magic(e, g, at));
}
Ok(())
}
#[doc(hidden)]
pub fn read_mapped<W, T, S, F>(r: &mut S, f: F) -> Result<T, BitError>
where
W: Bits,
S: Source,
F: FnOnce(W) -> T,
{
let raw: W = r.read()?;
Ok(f(raw))
}
#[doc(hidden)]
pub fn read_try_mapped<W, T, E, S, F>(r: &mut S, f: F) -> Result<T, BitError>
where
W: Bits,
S: Source,
E: fmt::Display,
F: FnOnce(W) -> Result<T, E>,
{
let at = r.bit_pos();
let raw: W = r.read()?;
f(raw).map_err(|e| BitError::convert(e.to_string(), at))
}
#[doc(hidden)]
pub fn write_mapped<W, T, K, F>(w: &mut K, value: &T, f: F) -> Result<(), BitError>
where
W: Bits,
K: Sink,
F: FnOnce(&T) -> W,
{
w.write(f(value))
}
pub trait BitAmount: Copy {
fn bits(self) -> u32;
fn bytes(self) -> u32;
}
macro_rules! impl_bit_amount {
($($t:ty),*) => {$(
impl BitAmount for $t {
fn bits(self) -> u32 { self as u32 }
fn bytes(self) -> u32 { (self as u32) * 8 }
}
)*};
}
impl_bit_amount!(u8, u16, u32, u64, usize, i32);
#[doc(hidden)]
pub fn skip_read<S: Source>(r: &mut S, bits: u32) -> Result<(), BitError> {
let mut left = bits;
while left > 0 {
let n = left.min(128);
r.read_bits(n)?;
left -= n;
}
Ok(())
}
#[doc(hidden)]
pub fn skip_write<K: Sink>(w: &mut K, bits: u32) -> Result<(), BitError> {
let mut left = bits;
while left > 0 {
let n = left.min(128);
w.write_bits(0, n)?;
left -= n;
}
Ok(())
}
#[doc(hidden)]
pub fn align_read<S: Source>(r: &mut S) -> Result<(), BitError> {
let pad = (8 - (r.bit_pos() % 8)) % 8;
skip_read(r, pad as u32)
}
#[doc(hidden)]
pub fn align_write<K: Sink>(w: &mut K) -> Result<(), BitError> {
let pad = (8 - (w.bit_pos() % 8)) % 8;
skip_write(w, pad as u32)
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct Layout {
pub bit: BitOrder,
pub byte: ByteOrder,
}
#[inline]
fn apply_byte_order(raw: u128, bits: u32, byte: ByteOrder) -> u128 {
if byte == ByteOrder::Big || bits % 8 != 0 {
return raw;
}
let n = (bits / 8) as usize;
let le = raw.to_le_bytes();
let mut out = 0u128;
let mut i = 0;
while i < n {
out |= (le[i] as u128) << (8 * (n - 1 - i));
i += 1;
}
out
}
#[inline]
fn extract_bits(buf: &[u8], pos: usize, n: usize, order: BitOrder) -> u128 {
if pos % 8 == 0 && n % 8 == 0 {
let start = pos / 8;
let nbytes = n / 8;
let mut acc = 0u128;
match order {
BitOrder::Msb => {
for j in 0..nbytes {
acc = (acc << 8) | u128::from(buf[start + j]);
}
}
BitOrder::Lsb => {
for j in 0..nbytes {
acc |= u128::from(buf[start + j]) << (8 * j);
}
}
}
return acc;
}
let mut acc = 0u128;
match order {
BitOrder::Msb => {
for k in 0..n {
let p = pos + k;
acc = (acc << 1) | u128::from((buf[p >> 3] >> (7 - (p & 7))) & 1);
}
}
BitOrder::Lsb => {
for k in 0..n {
let p = pos + k;
acc |= u128::from((buf[p >> 3] >> (p & 7)) & 1) << k;
}
}
}
acc
}
#[inline]
fn emit_bits(out: &mut Vec<u8>, bit_pos: usize, value: u128, n: usize, order: BitOrder) {
if n % 8 == 0 && bit_pos % 8 == 0 && bit_pos / 8 == out.len() {
let nbytes = n / 8;
match order {
BitOrder::Msb => {
for j in 0..nbytes {
out.push((value >> (8 * (nbytes - 1 - j))) as u8);
}
}
BitOrder::Lsb => {
for j in 0..nbytes {
out.push((value >> (8 * j)) as u8);
}
}
}
return;
}
for k in 0..n {
let p = bit_pos + k;
let (i, shift) = match order {
BitOrder::Msb => (n - 1 - k, 7 - (p & 7)),
BitOrder::Lsb => (k, p & 7),
};
let byte_idx = p >> 3;
if byte_idx == out.len() {
out.push(0);
}
if (value >> i) & 1 != 0 {
out[byte_idx] |= 1 << shift;
}
}
}
#[derive(Clone, Debug)]
pub struct BitReader<'a> {
bytes: &'a [u8],
bit_pos: usize,
order: BitOrder,
byte: ByteOrder,
}
impl<'a> BitReader<'a> {
#[must_use]
pub fn new(bytes: &'a [u8]) -> Self {
Self::with_order(bytes, BitOrder::Msb)
}
#[must_use]
pub fn with_order(bytes: &'a [u8], order: BitOrder) -> Self {
Self::with_layout(
bytes,
Layout {
bit: order,
byte: ByteOrder::Big,
},
)
}
#[must_use]
pub fn with_layout(bytes: &'a [u8], layout: Layout) -> Self {
Self {
bytes,
bit_pos: 0,
order: layout.bit,
byte: layout.byte,
}
}
#[must_use]
pub fn bit_pos(&self) -> usize {
self.bit_pos
}
#[must_use]
pub fn remaining_bits(&self) -> usize {
self.bytes.len() * 8 - self.bit_pos
}
#[inline]
pub fn read_bits(&mut self, n: u32) -> Result<u128, BitError> {
let n = n as usize;
if n > 128 {
return Err(BitError::new(ErrorKind::TooWide { width: n }, self.bit_pos));
}
if n > self.remaining_bits() {
return Err(BitError::new(
ErrorKind::UnexpectedEof {
needed: n,
remaining: self.remaining_bits(),
},
self.bit_pos,
));
}
let acc = extract_bits(self.bytes, self.bit_pos, n, self.order);
self.bit_pos += n;
Ok(acc)
}
#[inline]
pub fn read<T: Bits>(&mut self) -> Result<T, BitError> {
let raw = self.read_bits(T::BITS)?;
Ok(T::from_bits(apply_byte_order(raw, T::BITS, self.byte)))
}
pub fn seek_to_bit(&mut self, pos: usize) -> Result<(), BitError> {
let end = self.bytes.len() * 8;
if pos > end {
return Err(BitError::new(
ErrorKind::UnexpectedEof {
needed: pos,
remaining: end,
},
self.bit_pos,
));
}
self.bit_pos = pos;
Ok(())
}
pub fn align_to_byte(&mut self) {
self.bit_pos = (self.bit_pos + 7) & !7;
}
}
#[derive(Clone, Debug, Default)]
pub struct BitWriter {
bytes: Vec<u8>,
bit_pos: usize,
order: BitOrder,
byte: ByteOrder,
}
impl BitWriter {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn with_order(order: BitOrder) -> Self {
Self::with_layout(Layout {
bit: order,
byte: ByteOrder::Big,
})
}
#[must_use]
pub fn with_layout(layout: Layout) -> Self {
Self {
bytes: Vec::new(),
bit_pos: 0,
order: layout.bit,
byte: layout.byte,
}
}
#[must_use]
pub fn bit_len(&self) -> usize {
self.bit_pos
}
#[inline]
pub fn write_bits(&mut self, value: u128, n: u32) -> Result<(), BitError> {
let n = n as usize;
if n > 128 {
return Err(BitError::new(ErrorKind::TooWide { width: n }, self.bit_pos));
}
emit_bits(&mut self.bytes, self.bit_pos, value, n, self.order);
self.bit_pos += n;
Ok(())
}
#[inline]
pub fn write<T: Bits>(&mut self, value: T) -> Result<(), BitError> {
let raw = apply_byte_order(value.into_bits(), T::BITS, self.byte);
self.write_bits(raw, T::BITS)
}
#[must_use]
pub fn into_bytes(self) -> Vec<u8> {
self.bytes
}
}
pub trait Source {
fn read_bits(&mut self, n: u32) -> Result<u128, BitError>;
fn bit_pos(&self) -> usize;
fn byte_order(&self) -> ByteOrder {
ByteOrder::Big
}
fn seek_to_bit(&mut self, _pos: usize) -> Result<(), BitError> {
Err(BitError::new(ErrorKind::NotSeekable, self.bit_pos()))
}
#[inline]
fn read<T: Bits>(&mut self) -> Result<T, BitError> {
let raw = self.read_bits(T::BITS)?;
Ok(T::from_bits(apply_byte_order(
raw,
T::BITS,
self.byte_order(),
)))
}
#[cfg(feature = "std")]
fn as_read(&mut self) -> SourceReader<'_, Self>
where
Self: Sized,
{
SourceReader(self)
}
}
#[cfg(feature = "std")]
pub struct SourceReader<'a, S: Source>(&'a mut S);
#[cfg(feature = "std")]
impl<S: Source> std::io::Read for SourceReader<'_, S> {
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
for (i, slot) in buf.iter_mut().enumerate() {
match self.0.read_bits(8) {
Ok(b) => *slot = b as u8,
Err(e) if i == 0 => {
return Err(std::io::Error::new(
std::io::ErrorKind::UnexpectedEof,
e.to_string(),
));
}
Err(_) => return Ok(i),
}
}
Ok(buf.len())
}
}
pub trait SeekSource: Source {}
impl SeekSource for BitReader<'_> {}
pub trait Sink {
fn write_bits(&mut self, value: u128, n: u32) -> Result<(), BitError>;
fn bit_pos(&self) -> usize;
fn byte_order(&self) -> ByteOrder {
ByteOrder::Big
}
#[inline]
fn write<T: Bits>(&mut self, value: T) -> Result<(), BitError> {
let raw = apply_byte_order(value.into_bits(), T::BITS, self.byte_order());
self.write_bits(raw, T::BITS)
}
#[cfg(feature = "std")]
fn as_write(&mut self) -> SinkWriter<'_, Self>
where
Self: Sized,
{
SinkWriter(self)
}
}
#[cfg(feature = "std")]
pub struct SinkWriter<'a, K: Sink>(&'a mut K);
#[cfg(feature = "std")]
impl<K: Sink> std::io::Write for SinkWriter<'_, K> {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
for &b in buf {
self.0
.write_bits(u128::from(b), 8)
.map_err(|e| std::io::Error::other(e.to_string()))?;
}
Ok(buf.len())
}
fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}
impl Source for BitReader<'_> {
#[inline]
fn read_bits(&mut self, n: u32) -> Result<u128, BitError> {
BitReader::read_bits(self, n)
}
#[inline]
fn bit_pos(&self) -> usize {
self.bit_pos
}
#[inline]
fn byte_order(&self) -> ByteOrder {
self.byte
}
#[inline]
fn seek_to_bit(&mut self, pos: usize) -> Result<(), BitError> {
BitReader::seek_to_bit(self, pos)
}
}
impl Sink for BitWriter {
#[inline]
fn write_bits(&mut self, value: u128, n: u32) -> Result<(), BitError> {
BitWriter::write_bits(self, value, n)
}
#[inline]
fn bit_pos(&self) -> usize {
self.bit_pos
}
#[inline]
fn byte_order(&self) -> ByteOrder {
self.byte
}
}
pub trait BitDecode: Sized {
fn bit_decode<S: Source>(r: &mut S) -> Result<Self, BitError>;
}
pub trait FixedBitLen {
const BIT_LEN: u32;
}
pub trait BitEncode {
const LAYOUT: Layout = Layout {
bit: BitOrder::Msb,
byte: ByteOrder::Big,
};
fn bit_encode<K: Sink>(&self, w: &mut K) -> Result<(), BitError>;
}
#[cfg(feature = "std")]
pub trait EncodeExt: BitEncode {
fn encode<W: std::io::Write>(&self, w: &mut W) -> Result<(), BitError>
where
Self: Sized,
{
encode_to_writer(self, w, Self::LAYOUT)
}
}
#[cfg(feature = "std")]
impl<T: BitEncode> EncodeExt for T {}
pub trait SpecEncode {
const SPEC_LAYOUT: Layout;
fn spec_bit_encode<K: Sink>(&self, w: &mut K) -> Result<(), BitError>;
}
#[cfg(feature = "std")]
pub trait SpecEncodeExt: SpecEncode {
fn spec_encode<W: std::io::Write>(&self, w: &mut W) -> Result<(), BitError>
where
Self: Sized,
{
encode_to_writer_with(w, Self::SPEC_LAYOUT, |bw| self.spec_bit_encode(bw))
}
}
#[cfg(feature = "std")]
impl<T: SpecEncode> SpecEncodeExt for T {}
pub trait DecodeWith<A>: Sized {
fn decode_with<S: Source>(r: &mut S, args: A) -> Result<Self, BitError>;
}
pub trait EncodeWith<A> {
fn encode_with<K: Sink>(&self, w: &mut K, args: A) -> Result<(), BitError>;
}
impl<T: BitDecode> DecodeWith<()> for T {
fn decode_with<S: Source>(r: &mut S, _args: ()) -> Result<Self, BitError> {
T::bit_decode(r)
}
}
impl<T: BitEncode> EncodeWith<()> for T {
fn encode_with<K: Sink>(&self, w: &mut K, _args: ()) -> Result<(), BitError> {
self.bit_encode(w)
}
}
#[doc(hidden)]
pub fn decode_consume<T: BitDecode>(buf: &mut &[u8], layout: Layout) -> Result<T, BitError> {
let input = core::mem::take(buf);
let mut r = BitReader::with_layout(input, layout);
match T::bit_decode(&mut r) {
Ok(v) => {
*buf = &input[r.bit_pos().div_ceil(8)..];
Ok(v)
}
Err(e) => {
*buf = input;
Err(e)
}
}
}
#[doc(hidden)]
pub fn decode_peek<T: BitDecode>(bytes: &[u8], layout: Layout) -> Result<T, BitError> {
T::bit_decode(&mut BitReader::with_layout(bytes, layout))
}
#[doc(hidden)]
pub fn decode_peek_with<T, F>(bytes: &[u8], layout: Layout, f: F) -> Result<T, BitError>
where
F: FnOnce(&mut BitReader) -> Result<T, BitError>,
{
f(&mut BitReader::with_layout(bytes, layout))
}
#[doc(hidden)]
pub fn decode_exact_with<T, F>(bytes: &[u8], layout: Layout, f: F) -> Result<T, BitError>
where
F: FnOnce(&mut BitReader) -> Result<T, BitError>,
{
let mut r = BitReader::with_layout(bytes, layout);
let v = f(&mut r)?;
let consumed = r.bit_pos().div_ceil(8);
if consumed < bytes.len() {
return Err(BitError::new(
ErrorKind::TrailingBytes {
remaining: bytes.len() - consumed,
},
r.bit_pos(),
));
}
Ok(v)
}
#[doc(hidden)]
pub fn encode_to_vec_with<F>(layout: Layout, f: F) -> Result<Vec<u8>, BitError>
where
F: FnOnce(&mut BitWriter) -> Result<(), BitError>,
{
let mut w = BitWriter::with_layout(layout);
f(&mut w)?;
Ok(w.into_bytes())
}
#[doc(hidden)]
pub fn decode_exact<T: BitDecode>(bytes: &[u8], layout: Layout) -> Result<T, BitError> {
let mut r = BitReader::with_layout(bytes, layout);
let v = T::bit_decode(&mut r)?;
let consumed = r.bit_pos().div_ceil(8);
if consumed < bytes.len() {
return Err(BitError::new(
ErrorKind::TrailingBytes {
remaining: bytes.len() - consumed,
},
r.bit_pos(),
));
}
Ok(v)
}
#[doc(hidden)]
pub fn encode_to_vec<T: BitEncode>(value: &T, layout: Layout) -> Result<Vec<u8>, BitError> {
let mut w = BitWriter::with_layout(layout);
value.bit_encode(&mut w)?;
Ok(w.into_bytes())
}
#[cfg(feature = "std")]
#[doc(hidden)]
pub fn encode_to_writer<T: BitEncode, W: std::io::Write>(
value: &T,
w: &mut W,
layout: Layout,
) -> Result<(), BitError> {
let mut bw = BitWriter::with_layout(layout);
value.bit_encode(&mut bw)?;
let at = bw.bit_len();
w.write_all(&bw.into_bytes())
.map_err(|e| BitError::new(ErrorKind::Io(e.kind()), at))
}
#[cfg(feature = "std")]
#[doc(hidden)]
pub fn encode_to_writer_with<W, F>(w: &mut W, layout: Layout, f: F) -> Result<(), BitError>
where
W: std::io::Write,
F: FnOnce(&mut BitWriter) -> Result<(), BitError>,
{
let mut bw = BitWriter::with_layout(layout);
f(&mut bw)?;
let at = bw.bit_len();
w.write_all(&bw.into_bytes())
.map_err(|e| BitError::new(ErrorKind::Io(e.kind()), at))
}
#[doc(hidden)]
pub fn read_byte_array<const N: usize, S: Source>(r: &mut S) -> Result<[u8; N], BitError> {
let mut arr = [0u8; N];
for b in &mut arr {
*b = r.read_bits(8)? as u8;
}
Ok(arr)
}
#[doc(hidden)]
pub fn peek_bytes<S: Source>(r: &mut S, max: usize) -> Result<Vec<u8>, BitError> {
let start = r.bit_pos();
let mut out = Vec::with_capacity(max);
for _ in 0..max {
match r.read_bits(8) {
Ok(b) => out.push(b as u8),
Err(_) => break, }
}
r.seek_to_bit(start)?;
Ok(out)
}
#[doc(hidden)]
pub fn write_byte_array<const N: usize, K: Sink>(arr: &[u8; N], w: &mut K) -> Result<(), BitError> {
for &b in arr {
w.write_bits(u128::from(b), 8)?;
}
Ok(())
}
#[cfg(feature = "std")]
#[derive(Debug)]
pub struct StreamBitReader<R> {
inner: R,
lead: u32,
lead_bits: u32,
pos: usize,
}
#[cfg(feature = "std")]
impl<R: std::io::Read> StreamBitReader<R> {
pub fn new(inner: R) -> Self {
Self {
inner,
lead: 0,
lead_bits: 0,
pos: 0,
}
}
#[must_use]
pub fn bit_pos(&self) -> usize {
self.pos
}
pub fn read_bits(&mut self, n: u32) -> Result<u128, BitError> {
if n > 128 {
return Err(BitError::new(
ErrorKind::TooWide { width: n as usize },
self.pos,
));
}
let at = self.pos;
let mut result: u128 = 0;
let mut need = n;
while need > 0 {
if self.lead_bits == 0 {
let mut b = [0u8; 1];
if self.inner.read_exact(&mut b).is_err() {
return Err(BitError::new(ErrorKind::Incomplete { needed: None }, at));
}
self.lead = u32::from(b[0]);
self.lead_bits = 8;
}
let take = need.min(self.lead_bits);
let shift = self.lead_bits - take;
let chunk = (self.lead >> shift) & ((1u32 << take) - 1);
result = (result << take) | u128::from(chunk);
self.lead_bits -= take;
self.lead &= (1u32 << self.lead_bits) - 1; need -= take;
}
self.pos += n as usize;
Ok(result)
}
pub fn read<T: Bits>(&mut self) -> Result<T, BitError> {
Ok(T::from_bits(self.read_bits(T::BITS)?))
}
}
#[cfg(feature = "std")]
impl<R: std::io::Read> Source for StreamBitReader<R> {
fn read_bits(&mut self, n: u32) -> Result<u128, BitError> {
StreamBitReader::read_bits(self, n)
}
fn bit_pos(&self) -> usize {
self.pos
}
}
#[cfg(feature = "std")]
#[derive(Clone, Debug)]
pub struct BufSource<R> {
inner: R,
buf: Vec<u8>,
bit_pos: usize,
cap: usize,
layout: Layout,
eof: bool,
}
#[cfg(feature = "std")]
impl<R: std::io::Read> BufSource<R> {
#[must_use]
pub fn new(inner: R) -> Self {
Self::with_cap(inner, 64 * 1024)
}
#[must_use]
pub fn with_cap(inner: R, cap: usize) -> Self {
Self::with_cap_and_layout(inner, cap, Layout::default())
}
#[must_use]
pub fn with_cap_and_layout(inner: R, cap: usize, layout: Layout) -> Self {
Self {
inner,
buf: Vec::new(),
bit_pos: 0,
cap,
layout,
eof: false,
}
}
fn fill_to(&mut self, byte_end: usize) -> Result<(), BitError> {
while self.buf.len() < byte_end && !self.eof {
if self.buf.len() >= self.cap {
return Err(BitError::new(
ErrorKind::BufferFull { cap: self.cap },
self.bit_pos,
));
}
let want = (byte_end - self.buf.len()).min(self.cap - self.buf.len());
let start = self.buf.len();
self.buf.resize(start + want, 0);
match self.inner.read(&mut self.buf[start..]) {
Ok(0) => {
self.buf.truncate(start);
self.eof = true;
}
Ok(got) => self.buf.truncate(start + got),
Err(e) => {
self.buf.truncate(start);
return Err(BitError::new(ErrorKind::Io(e.kind()), self.bit_pos));
}
}
}
Ok(())
}
}
#[cfg(feature = "std")]
impl<R: std::io::Read> Source for BufSource<R> {
fn read_bits(&mut self, n: u32) -> Result<u128, BitError> {
if n > 128 {
return Err(BitError::new(
ErrorKind::TooWide { width: n as usize },
self.bit_pos,
));
}
let byte_end = (self.bit_pos + n as usize).div_ceil(8);
self.fill_to(byte_end)?;
if self.buf.len() < byte_end {
return Err(BitError::new(
ErrorKind::Incomplete {
needed: Some(byte_end - self.buf.len()),
},
self.bit_pos,
));
}
let acc = extract_bits(&self.buf, self.bit_pos, n as usize, self.layout.bit);
self.bit_pos += n as usize;
Ok(acc)
}
fn bit_pos(&self) -> usize {
self.bit_pos
}
fn byte_order(&self) -> ByteOrder {
self.layout.byte
}
fn seek_to_bit(&mut self, pos: usize) -> Result<(), BitError> {
self.bit_pos = pos;
Ok(())
}
}
#[cfg(feature = "std")]
impl<R: std::io::Read> SeekSource for BufSource<R> {}
#[cfg(feature = "std")]
#[derive(Clone, Debug)]
pub struct SeekReader<R> {
inner: R,
bit_pos: usize,
layout: Layout,
}
#[cfg(feature = "std")]
impl<R: std::io::Read + std::io::Seek> SeekReader<R> {
#[must_use]
pub fn new(inner: R) -> Self {
Self::with_layout(inner, Layout::default())
}
#[must_use]
pub fn with_layout(inner: R, layout: Layout) -> Self {
Self {
inner,
bit_pos: 0,
layout,
}
}
}
#[cfg(feature = "std")]
impl<R: std::io::Read + std::io::Seek> Source for SeekReader<R> {
fn read_bits(&mut self, n: u32) -> Result<u128, BitError> {
if n > 128 {
return Err(BitError::new(
ErrorKind::TooWide { width: n as usize },
self.bit_pos,
));
}
let bit_off = self.bit_pos % 8;
let byte_start = (self.bit_pos / 8) as u64;
let nbytes = (bit_off + n as usize).div_ceil(8);
self.inner
.seek(std::io::SeekFrom::Start(byte_start))
.map_err(|e| BitError::new(ErrorKind::Io(e.kind()), self.bit_pos))?;
let mut buf = vec![0u8; nbytes];
self.inner.read_exact(&mut buf).map_err(|e| {
let kind = if e.kind() == std::io::ErrorKind::UnexpectedEof {
ErrorKind::UnexpectedEof {
needed: n as usize,
remaining: 0,
}
} else {
ErrorKind::Io(e.kind())
};
BitError::new(kind, self.bit_pos)
})?;
let acc = extract_bits(&buf, bit_off, n as usize, self.layout.bit);
self.bit_pos += n as usize;
Ok(acc)
}
fn bit_pos(&self) -> usize {
self.bit_pos
}
fn byte_order(&self) -> ByteOrder {
self.layout.byte
}
fn seek_to_bit(&mut self, pos: usize) -> Result<(), BitError> {
self.bit_pos = pos; Ok(())
}
}
#[cfg(feature = "std")]
impl<R: std::io::Read + std::io::Seek> SeekSource for SeekReader<R> {}
#[cfg(feature = "bytes")]
mod bytes_io {
use super::{BitError, BitReader, BitWriter, ByteOrder, Layout, SeekSource, Sink, Source};
#[derive(Clone, Debug)]
pub struct BytesReader {
data: bytes::Bytes,
bit_pos: usize,
layout: Layout,
}
impl BytesReader {
#[must_use]
pub fn new(data: bytes::Bytes) -> Self {
Self::with_layout(data, Layout::default())
}
#[must_use]
pub fn with_layout(data: bytes::Bytes, layout: Layout) -> Self {
Self {
data,
bit_pos: 0,
layout,
}
}
}
impl Source for BytesReader {
fn read_bits(&mut self, n: u32) -> Result<u128, BitError> {
let mut br = BitReader::with_layout(&self.data, self.layout);
br.seek_to_bit(self.bit_pos)?;
let v = br.read_bits(n)?;
self.bit_pos = Source::bit_pos(&br);
Ok(v)
}
fn bit_pos(&self) -> usize {
self.bit_pos
}
fn byte_order(&self) -> ByteOrder {
self.layout.byte
}
fn seek_to_bit(&mut self, pos: usize) -> Result<(), BitError> {
self.bit_pos = pos;
Ok(())
}
}
impl SeekSource for BytesReader {}
#[derive(Clone, Debug, Default)]
pub struct BytesWriter {
inner: BitWriter,
}
impl BytesWriter {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn with_layout(layout: Layout) -> Self {
Self {
inner: BitWriter::with_layout(layout),
}
}
#[must_use]
pub fn freeze(self) -> bytes::Bytes {
bytes::Bytes::from(self.inner.into_bytes())
}
}
impl Sink for BytesWriter {
fn write_bits(&mut self, value: u128, n: u32) -> Result<(), BitError> {
self.inner.write_bits(value, n)
}
fn bit_pos(&self) -> usize {
Sink::bit_pos(&self.inner)
}
fn byte_order(&self) -> ByteOrder {
Sink::byte_order(&self.inner)
}
}
}
#[cfg(feature = "bytes")]
pub use bytes_io::{BytesReader, BytesWriter};
#[cfg(test)]
mod unit {
use super::*;
use crate::{u4, u12};
#[test]
fn unaligned_round_trip() {
let mut w = BitWriter::new();
w.write(u4::new(0xA)).unwrap();
w.write(u12::new(0xBCD)).unwrap();
assert_eq!(w.bit_len(), 16);
let bytes = w.into_bytes();
assert_eq!(bytes, [0xAB, 0xCD]);
let mut r = BitReader::new(&bytes);
assert_eq!(r.read::<u4>().unwrap(), u4::new(0xA));
assert_eq!(r.read::<u12>().unwrap(), u12::new(0xBCD));
assert_eq!(r.remaining_bits(), 0);
}
#[test]
fn eof_is_an_error_not_a_panic() {
let mut r = BitReader::new(&[0xFF]);
assert_eq!(r.read::<u4>().unwrap(), u4::new(0xF));
let err = r.read_bits(8).unwrap_err();
assert_eq!(
err.kind,
ErrorKind::UnexpectedEof {
needed: 8,
remaining: 4
}
);
assert_eq!(err.at, 4, "error records the bit offset");
assert!(err.field.is_none(), "no field context at the reader level");
}
#[test]
fn too_wide_is_rejected() {
let mut r = BitReader::new(&[0u8; 32]);
let err = r.read_bits(129).unwrap_err();
assert_eq!(err.kind, ErrorKind::TooWide { width: 129 });
}
#[test]
fn stream_reader_matches_slice_up_to_128_bits() {
let bytes: Vec<u8> = (0u8..16).collect();
let mut s = StreamBitReader::new(&bytes[..]);
let mut r = BitReader::new(&bytes);
assert_eq!(s.read_bits(128).unwrap(), r.read_bits(128).unwrap());
let mut s = StreamBitReader::new(&bytes[..]);
let mut r = BitReader::new(&bytes);
assert_eq!(s.read_bits(100).unwrap(), r.read_bits(100).unwrap());
assert_eq!(s.read_bits(28).unwrap(), r.read_bits(28).unwrap());
let mut s = StreamBitReader::new(&bytes[..]);
assert_eq!(
s.read_bits(65).unwrap(),
BitReader::new(&bytes).read_bits(65).unwrap(),
"a 65-bit read used to be rejected"
);
let mut s = StreamBitReader::new(&bytes[..]);
assert_eq!(
s.read_bits(129).unwrap_err().kind,
ErrorKind::TooWide { width: 129 }
);
}
}