use core::fmt;
use core::ops::{BitAnd, BitOr, BitOrAssign, Not};
use core::sync::atomic::{AtomicU32, Ordering};
#[derive(Clone, Copy, PartialEq, Eq, Hash, Default)]
pub struct SymbolFlags(u32);
impl SymbolFlags {
pub const EMPTY: Self = Self(0);
pub const REFERENCED: Self = Self(1 << 0);
pub const WEAK_REFERENCED: Self = Self(1 << 1);
pub const ADDRESS_TAKEN: Self = Self(1 << 2);
pub const EXPORTED: Self = Self(1 << 3);
pub const NEEDS_GOT: Self = Self(1 << 4);
pub const NEEDS_PLT: Self = Self(1 << 5);
pub const NEEDS_COPY_RELOC: Self = Self(1 << 6);
pub const NEEDS_TLSGD: Self = Self(1 << 7);
pub const NEEDS_TLSDESC: Self = Self(1 << 8);
pub const NEEDS_GOTTPOFF: Self = Self(1 << 9);
pub const NEEDS_DYNSYM: Self = Self(1 << 10);
pub const NEEDS_CANONICAL_PLT: Self = Self(1 << 11);
pub const FIRST_BACKEND_BIT: u32 = 16;
#[inline]
#[must_use]
pub const fn backend(n: u32) -> Self {
assert!(n < 16, "backend flag index out of range");
Self(1 << (Self::FIRST_BACKEND_BIT + n))
}
#[inline]
#[must_use]
pub const fn from_bits(bits: u32) -> Self {
Self(bits)
}
#[inline]
#[must_use]
pub const fn bits(self) -> u32 {
self.0
}
#[inline]
#[must_use]
pub const fn contains(self, other: Self) -> bool {
self.0 & other.0 == other.0
}
#[inline]
#[must_use]
pub const fn intersects(self, other: Self) -> bool {
self.0 & other.0 != 0
}
#[inline]
#[must_use]
pub const fn is_empty(self) -> bool {
self.0 == 0
}
}
impl BitOr for SymbolFlags {
type Output = Self;
#[inline]
fn bitor(self, rhs: Self) -> Self {
Self(self.0 | rhs.0)
}
}
impl BitOrAssign for SymbolFlags {
#[inline]
fn bitor_assign(&mut self, rhs: Self) {
self.0 |= rhs.0;
}
}
impl BitAnd for SymbolFlags {
type Output = Self;
#[inline]
fn bitand(self, rhs: Self) -> Self {
Self(self.0 & rhs.0)
}
}
impl Not for SymbolFlags {
type Output = Self;
#[inline]
fn not(self) -> Self {
Self(!self.0)
}
}
impl fmt::Debug for SymbolFlags {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
const NAMES: [&str; 12] = [
"REFERENCED",
"WEAK_REFERENCED",
"ADDRESS_TAKEN",
"EXPORTED",
"NEEDS_GOT",
"NEEDS_PLT",
"NEEDS_COPY_RELOC",
"NEEDS_TLSGD",
"NEEDS_TLSDESC",
"NEEDS_GOTTPOFF",
"NEEDS_DYNSYM",
"NEEDS_CANONICAL_PLT",
];
f.write_str("SymbolFlags(")?;
let mut first = true;
for bit in 0..32u32 {
if self.0 & (1 << bit) == 0 {
continue;
}
if !first {
f.write_str(" | ")?;
}
first = false;
match NAMES.get(bit as usize) {
Some(name) => f.write_str(name)?,
None if bit >= Self::FIRST_BACKEND_BIT => {
write!(f, "BACKEND_{}", bit - Self::FIRST_BACKEND_BIT)?;
}
None => write!(f, "BIT_{bit}")?,
}
}
f.write_str(")")
}
}
#[inline]
pub(crate) fn set(cell: &AtomicU32, flags: SymbolFlags) -> SymbolFlags {
let current = cell.load(Ordering::Relaxed);
if current & flags.0 == flags.0 {
return SymbolFlags(current);
}
SymbolFlags(cell.fetch_or(flags.0, Ordering::Relaxed))
}
#[inline]
pub(crate) fn clear(cell: &AtomicU32, flags: SymbolFlags) -> SymbolFlags {
SymbolFlags(cell.fetch_and(!flags.0, Ordering::Relaxed))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn set_reports_previous_bits() {
let cell = AtomicU32::new(0);
let before = set(&cell, SymbolFlags::NEEDS_GOT);
assert!(before.is_empty());
let before = set(&cell, SymbolFlags::NEEDS_GOT | SymbolFlags::NEEDS_PLT);
assert_eq!(before, SymbolFlags::NEEDS_GOT);
let before = set(&cell, SymbolFlags::NEEDS_PLT);
assert!(before.contains(SymbolFlags::NEEDS_GOT | SymbolFlags::NEEDS_PLT));
let before = clear(&cell, SymbolFlags::NEEDS_GOT);
assert!(before.contains(SymbolFlags::NEEDS_GOT));
assert_eq!(
SymbolFlags::from_bits(cell.load(Ordering::Relaxed)),
SymbolFlags::NEEDS_PLT
);
}
#[test]
fn backend_bits_do_not_overlap_generic_bits() {
let generic = SymbolFlags::from_bits(0xffff);
for n in 0..16 {
assert!(!generic.intersects(SymbolFlags::backend(n)));
}
assert_eq!(
format!("{:?}", SymbolFlags::EXPORTED | SymbolFlags::backend(2)),
"SymbolFlags(EXPORTED | BACKEND_2)"
);
}
#[test]
fn concurrent_sets_are_all_kept() {
let cell = AtomicU32::new(0);
std::thread::scope(|scope| {
for bit in 0..16 {
let cell = &cell;
scope.spawn(move || {
for _ in 0..1000 {
set(cell, SymbolFlags::from_bits(1 << bit));
}
});
}
});
assert_eq!(cell.load(Ordering::Relaxed), 0xffff);
}
}