use core::{marker::PhantomData, ops::BitOr};
pub trait BitFlag: Copy {
const ALL_BITS: u8;
fn bits(self) -> u8;
}
pub struct BitFlags<T: BitFlag> {
bits: u8,
_marker: PhantomData<T>,
}
impl<T> BitFlags<T>
where
T: BitFlag,
{
pub const fn empty() -> Self {
Self {
bits: 0,
_marker: PhantomData,
}
}
pub(crate) const unsafe fn from_bits_unchecked(bits: u8) -> Self {
Self {
bits,
_marker: PhantomData,
}
}
pub fn from_bits(bits: u8) -> Result<Self, UnknownBits> {
let unknown = bits & !T::ALL_BITS;
if unknown != 0 {
return Err(UnknownBits { bits: unknown });
}
Ok(Self {
bits,
_marker: PhantomData,
})
}
pub const fn bits(self) -> u8 {
self.bits
}
pub fn is_empty(self) -> bool {
self.bits == 0
}
pub fn contains(self, other: impl Into<BitFlags<T>>) -> bool {
let other = other.into().bits;
self.bits & other == other
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct UnknownBits {
pub bits: u8,
}
impl core::fmt::Display for UnknownBits {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
write!(f, "unknown flag bits {:#010b}", self.bits)
}
}
impl core::error::Error for UnknownBits {}
impl<T> From<T> for BitFlags<T>
where
T: BitFlag,
{
fn from(flag: T) -> Self {
Self {
bits: flag.bits(),
_marker: PhantomData,
}
}
}
impl<T> BitOr for BitFlags<T>
where
T: BitFlag,
{
type Output = Self;
fn bitor(self, rhs: Self) -> Self {
Self {
bits: self.bits | rhs.bits,
_marker: PhantomData,
}
}
}
impl<T> BitOr<T> for BitFlags<T>
where
T: BitFlag,
{
type Output = Self;
fn bitor(self, rhs: T) -> Self {
self | BitFlags::from(rhs)
}
}
impl<T> Clone for BitFlags<T>
where
T: BitFlag,
{
fn clone(&self) -> Self {
*self
}
}
impl<T> Copy for BitFlags<T> where T: BitFlag {}
impl<T> PartialEq for BitFlags<T>
where
T: BitFlag,
{
fn eq(&self, other: &Self) -> bool {
self.bits == other.bits
}
}
impl<T> Eq for BitFlags<T> where T: BitFlag {}
impl<T> core::fmt::Debug for BitFlags<T>
where
T: BitFlag,
{
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
write!(f, "BitFlags({:#010b})", self.bits)
}
}
macro_rules! bitflag_enum {
(
$(#[$meta:meta])*
$vis:vis enum $name:ident {
$( $(#[$vmeta:meta])* $variant:ident = $value:expr ),+ $(,)?
}
) => {
$(#[$meta])*
#[repr(u8)]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
$vis enum $name {
$( $(#[$vmeta])* $variant = $value ),+
}
impl $crate::flags::BitFlag for $name {
const ALL_BITS: u8 = 0 $( | ($name::$variant as u8) )+;
fn bits(self) -> u8 {
self as u8
}
}
impl ::core::ops::BitOr for $name {
type Output = $crate::flags::BitFlags<$name>;
fn bitor(self, rhs: Self) -> Self::Output {
$crate::flags::BitFlags::from(self) | rhs
}
}
};
}
pub(crate) use bitflag_enum;
#[cfg(test)]
mod tests {
use super::*;
bitflag_enum! {
pub enum Test {
A = 1 << 0,
B = 1 << 1,
C = 1 << 2,
}
}
#[test]
fn all_bits_is_or_of_variants() {
assert_eq!(Test::ALL_BITS, 0b0000_0111);
}
#[test]
fn from_bits_validates_bits() {
assert_eq!(
BitFlags::<Test>::from_bits(0b1101),
Err(UnknownBits { bits: 0b1000 }),
);
let set = BitFlags::<Test>::from_bits(0b0101).unwrap();
assert_eq!(set.bits(), 0b0101);
assert!(set.contains(Test::A));
assert!(set.contains(Test::C));
assert!(!set.contains(Test::B));
assert!(BitFlags::<Test>::from_bits(0).unwrap().is_empty());
}
#[test]
fn empty_set_is_empty() {
let set = BitFlags::<Test>::empty();
assert!(set.is_empty());
assert_eq!(set.bits(), 0);
assert!(!set.contains(Test::A));
}
#[test]
fn from_single_flag() {
let set: BitFlags<Test> = Test::B.into();
assert_eq!(set.bits(), 0b0010);
assert!(!set.is_empty());
}
#[test]
fn combine_with_bitor() {
let set = Test::A | Test::C; assert!(set.contains(Test::A));
assert!(set.contains(Test::C));
assert!(!set.contains(Test::B));
let more = set | Test::B; assert_eq!(more.bits(), Test::ALL_BITS);
let union = (Test::A | Test::B) | (Test::B | Test::C); assert_eq!(union.bits(), Test::ALL_BITS);
}
#[test]
fn contains_superset_semantics() {
let set = Test::A | Test::B;
assert!(set.contains(Test::A | Test::B));
assert!(!set.contains(Test::A | Test::C));
}
#[test]
fn unknown_bits_is_an_error() {
fn assert_error<E>()
where
E: core::error::Error,
{
}
assert_error::<UnknownBits>();
use core::fmt::Write as _;
let mut buf = crate::fmt::FmtBuf::<64>::new();
write!(buf, "{}", UnknownBits { bits: 0b1000 }).unwrap();
assert_eq!(buf.as_str(), "unknown flag bits 0b00001000");
}
}